Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 6 additions & 6 deletions tests/benchmarks/test_accuracy_bench_utils.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,11 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project

# ruff: noqa: E402, I001
import argparse
import math
import os
import sys
import types
from pathlib import Path

import pytest
Expand Down Expand Up @@ -106,11 +108,9 @@ def fake_snapshot_download(*, repo_id, repo_type, allow_patterns):
captured["allow_patterns"] = allow_patterns
return str(downloaded_root)

monkeypatch.setitem(
sys.modules,
"huggingface_hub",
types.SimpleNamespace(snapshot_download=fake_snapshot_download),
)
from vllm_omni.transformers_utils import repo_utils

monkeypatch.setattr(repo_utils.hf_api(), "snapshot_download", fake_snapshot_download)

resolved = resolve_seed_tts_root(
"zhaochenyang20/seed-tts-eval",
Expand Down
13 changes: 7 additions & 6 deletions tests/dfx/perf/scripts/run_diffusion_benchmark.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project

"""
Performance benchmark CI runner for diffusion models.

Expand Down Expand Up @@ -356,6 +359,8 @@ def _resolve_offline_model(model: str) -> str:
"""
import huggingface_hub

from vllm_omni.transformers_utils.repo_utils import hf_api

if not model or os.path.isdir(model):
return model

Expand All @@ -380,9 +385,7 @@ def _resolve_offline_model(model: str) -> str:
if len(parts) >= 3:
repo_id = "/".join(parts[:2])
subfolder = "/".join(parts[2:])
from huggingface_hub import snapshot_download

snapshot_root = snapshot_download(
snapshot_root = hf_api().snapshot_download(
repo_id,
allow_patterns=[f"{subfolder}/**"],
local_files_only=huggingface_hub.constants.HF_HUB_OFFLINE,
Expand All @@ -391,9 +394,7 @@ def _resolve_offline_model(model: str) -> str:

if not huggingface_hub.constants.HF_HUB_OFFLINE:
return model
from huggingface_hub import snapshot_download

return snapshot_download(model, local_files_only=True)
return hf_api().snapshot_download(model, local_files_only=True)


class DiffusionServer:
Expand Down
4 changes: 2 additions & 2 deletions tests/diffusion/model_loader/test_diffusers_loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,6 @@
import pytest
import torch
import torch.nn as nn
from huggingface_hub import snapshot_download
from safetensors.torch import save_file
from vllm.config.load import LoadConfig

Expand All @@ -32,6 +31,7 @@
from vllm_omni.diffusion.models.host_weight_contract import FinalLayoutModelContract
from vllm_omni.diffusion.registry import initialize_model
from vllm_omni.quantization.component_config import ComponentQuantizationConfig
from vllm_omni.transformers_utils.repo_utils import hf_api

pytestmark = [pytest.mark.core_model, pytest.mark.diffusion, pytest.mark.cpu]

Expand All @@ -41,7 +41,7 @@
@pytest.fixture(scope="module")
def prefetch_helios_model():
"""Downloads the tiny helios model prior to running a test."""
snapshot_download(model_path)
hf_api().snapshot_download(model_path)


@pytest.fixture(scope="function")
Expand Down
6 changes: 3 additions & 3 deletions tests/diffusion/model_loader/test_hub_prefetch.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project

import contextlib

Expand All @@ -14,10 +14,10 @@
def test_prefetch_subfolders_propagates_revision(monkeypatch):
calls = []

def fake_snapshot_download(**kwargs):
def fake_snapshot_download(self, **kwargs):
calls.append(kwargs)

monkeypatch.setattr(huggingface_hub, "snapshot_download", fake_snapshot_download)
monkeypatch.setattr(huggingface_hub.HfApi, "snapshot_download", fake_snapshot_download)
monkeypatch.setattr(hub_prefetch, "_repo_prefetch_lock", lambda _model: contextlib.nullcontext())

hub_prefetch.prefetch_subfolders(
Expand Down
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project

from __future__ import annotations

Expand Down Expand Up @@ -81,9 +81,8 @@ def test_from_config_loads_local_diffusers_component(tmp_path, monkeypatch: pyte


def test_from_config_downloads_component_from_hf_repo(tmp_path, monkeypatch: pytest.MonkeyPatch) -> None:
import huggingface_hub

from vllm_omni.diffusion.models.cosmos3 import sound_tokenizer
from vllm_omni.transformers_utils import repo_utils

cache_dir = tmp_path / "hf"
_write_component(cache_dir, checkpoint_name=DIFFUSERS_SOUND_TOKENIZER_CHECKPOINT_NAME)
Expand All @@ -95,7 +94,7 @@ def fake_snapshot_download(repo_id: str, *, revision: str | None, allow_patterns
calls.append((repo_id, revision, allow_patterns))
return str(cache_dir)

monkeypatch.setattr(huggingface_hub, "snapshot_download", fake_snapshot_download)
monkeypatch.setattr(repo_utils.hf_api(), "snapshot_download", fake_snapshot_download)

sound_tokenizer.Cosmos3SoundTokenizer.from_config(
SimpleNamespace(
Expand Down
13 changes: 8 additions & 5 deletions tests/diffusion/models/ltx2/test_ltx2_pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
from types import SimpleNamespace
from typing import Any

import huggingface_hub
import numpy as np
import pytest
import torch
Expand Down Expand Up @@ -302,11 +303,11 @@ def test_ltx_artifact_uses_source_revision_and_hub_fallback(
filename = "ltx-sidecar.safetensors"
calls = []

def fake_download(**kwargs):
def fake_download(self, **kwargs):
calls.append(kwargs)
return "/cache/ltx-sidecar.safetensors"

monkeypatch.setattr(ltx2_components, "hf_hub_download", fake_download)
monkeypatch.setattr(huggingface_hub.HfApi, "hf_hub_download", fake_download)

assert (
resolve_ltx_artifact(
Expand All @@ -331,7 +332,9 @@ def test_ltx_artifact_prefers_model_root(tmp_path, monkeypatch):
filename = "ltx-sidecar.safetensors"
expected = tmp_path / filename
expected.write_bytes(b"sidecar")
monkeypatch.setattr(ltx2_components, "hf_hub_download", lambda **_kwargs: pytest.fail("unexpected Hub lookup"))
monkeypatch.setattr(
huggingface_hub.HfApi, "hf_hub_download", lambda *_args, **_kwargs: pytest.fail("unexpected Hub lookup")
)

assert resolve_ltx_artifact(
str(tmp_path),
Expand All @@ -345,11 +348,11 @@ def test_ltx_artifact_prefers_model_root(tmp_path, monkeypatch):
def test_ltx_artifact_local_model_missing_sidecar_falls_back_to_hub(tmp_path, monkeypatch):
calls = []

def fake_download(**kwargs):
def fake_download(self, **kwargs):
calls.append(kwargs)
return "/cache/ltx-sidecar.safetensors"

monkeypatch.setattr(ltx2_components, "hf_hub_download", fake_download)
monkeypatch.setattr(huggingface_hub.HfApi, "hf_hub_download", fake_download)

assert (
resolve_ltx_artifact(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -42,10 +42,10 @@


def _resolve_fl2va_model_ref() -> str:
from huggingface_hub import snapshot_download
from vllm_omni.transformers_utils.repo_utils import hf_api

repo_root = Path(
snapshot_download(
hf_api().snapshot_download(
repo_id=_MINIMAX_H3_REPO,
revision=_MINIMAX_H3_REVISION,
allow_patterns=["FL2VA/**"],
Expand Down Expand Up @@ -253,13 +253,13 @@ def test_resolve_fl2va_model_ref(tmp_path, monkeypatch):
fl2va_root.mkdir()
(fl2va_root / "model_index.json").write_text("{}", encoding="utf-8")

def fake_snapshot_download(*, repo_id, revision, allow_patterns):
def fake_snapshot_download(self, *, repo_id, revision, allow_patterns):
assert repo_id == _MINIMAX_H3_REPO
assert revision == _MINIMAX_H3_REVISION
assert allow_patterns == ["FL2VA/**"]
return str(tmp_path)

monkeypatch.setattr("huggingface_hub.snapshot_download", fake_snapshot_download)
monkeypatch.setattr("huggingface_hub.HfApi.snapshot_download", fake_snapshot_download)
assert _resolve_fl2va_model_ref() == str(fl2va_root)


Expand Down
10 changes: 5 additions & 5 deletions tests/e2e/accuracy/ltx/test_ltx_official_similarity.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project

"""E2E accuracy guard against a pinned Lightricks LTX pipeline revision.

Expand All @@ -22,11 +22,11 @@
import numpy as np
import pytest
import torch
from huggingface_hub import hf_hub_download, snapshot_download
from torchmetrics.image import PeakSignalNoiseRatio, StructuralSimilarityIndexMeasure

from tests.e2e.accuracy.helpers import reset_artifact_dir
from tests.helpers.mark import hardware_test
from vllm_omni.transformers_utils.repo_utils import hf_api

OFFICIAL_REPOSITORY = "https://github.com/Lightricks/LTX-2.git"
OFFICIAL_REVISION = "9377758131b1ffde4b7f766804590a6617bf2ab9"
Expand Down Expand Up @@ -326,7 +326,7 @@ def _resolve_model(case: LTXAccuracyCase) -> Path:
if revision is None and model_id == case.model_id:
revision = case.model_revision
return Path(
snapshot_download(
hf_api().snapshot_download(
repo_id=model_id,
revision=revision,
allow_patterns=[
Expand Down Expand Up @@ -359,7 +359,7 @@ def _resolve_gemma_root(case: LTXAccuracyCase, model: Path) -> Path:
if configured_model and Path(configured_model).is_dir():
return Path(configured_model)
return Path(
snapshot_download(
hf_api().snapshot_download(
repo_id=case.gemma_model_id,
revision=case.gemma_model_revision,
allow_patterns=[
Expand All @@ -383,7 +383,7 @@ def _resolve_artifact(artifact: LTXArtifact, model: Path | None = None) -> Path:
if model_path.is_file():
return model_path
return Path(
hf_hub_download(
hf_api().hf_hub_download(
repo_id=artifact.repo_id,
repo_type=None if artifact.repo_type == "model" else artifact.repo_type,
filename=artifact.filename,
Expand Down
6 changes: 3 additions & 3 deletions tests/e2e/features/rlhf_test/test_verl_omni_e2e.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project

"""E2E test that follows the EXACT same flow as
``verl-omni/tests/workers/rollout/rollout_vllm/test_vllm_omni_generate.py``
Expand Down Expand Up @@ -35,7 +35,6 @@
import pytest
import ray
import torch
from huggingface_hub import snapshot_download
from omegaconf import OmegaConf
from transformers import AutoTokenizer

Expand All @@ -46,6 +45,7 @@
vLLMOmniHttpServerLocal,
)
from tests.helpers.mark import hardware_test
from vllm_omni.transformers_utils.repo_utils import hf_api

MODEL = "tiny-random/Qwen-Image"
TOKENIZER_MODEL = "Qwen/Qwen2-1.5B-Instruct"
Expand All @@ -55,7 +55,7 @@
def _resolve_model_path(repo_id: str) -> str:
if os.path.isdir(repo_id):
return repo_id
return snapshot_download(repo_id=repo_id)
return hf_api().snapshot_download(repo_id=repo_id)


@lru_cache(maxsize=1)
Expand Down
6 changes: 3 additions & 3 deletions tests/e2e/offline_inference/test_cosyvoice3_expansion.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project
"""
Offline E2E smoke test for CosyVoice3 zero-shot reference inference.

Expand All @@ -20,7 +20,6 @@
import numpy as np
import pytest
import soundfile as sf
from huggingface_hub import snapshot_download
from vllm.sampling_params import SamplingParams

from tests.helpers.mark import hardware_test
Expand All @@ -29,6 +28,7 @@
from tests.helpers.stage_config import get_deploy_config_path
from vllm_omni.model_executor.models.cosyvoice3.tokenizer import get_qwen_tokenizer
from vllm_omni.transformers_utils.configs.cosyvoice3 import CosyVoice3Config
from vllm_omni.transformers_utils.repo_utils import hf_api

MODEL = "FunAudioLLM/Fun-CosyVoice3-0.5B-2512"
MODEL_DIR_ENV = "VLLM_OMNI_COSYVOICE3_MODEL_DIR"
Expand Down Expand Up @@ -58,7 +58,7 @@ def _resolve_model_dir() -> Path:
override = os.environ.get(MODEL_DIR_ENV)
if override:
return Path(override).expanduser().resolve()
return Path(snapshot_download(MODEL, allow_patterns=["*"]))
return Path(hf_api().snapshot_download(MODEL, allow_patterns=["*"]))


def _reference_zero_shot_stage0_sampling(*, text: str) -> SamplingParams:
Expand Down
7 changes: 5 additions & 2 deletions tests/e2e/offline_inference/test_mammoth_moda2_expansion.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project

"""
End-to-end test for MammothModa2 text-to-image generation.

Expand All @@ -19,12 +22,12 @@

import pytest
import torch
from huggingface_hub import snapshot_download
from vllm.sampling_params import SamplingParams

from tests.helpers.mark import hardware_test
from tests.helpers.runtime import OmniRunner
from tests.helpers.stage_config import get_deploy_config_path
from vllm_omni.transformers_utils.repo_utils import hf_api

# ---------------------------------------------------------------------------
# Constants
Expand Down Expand Up @@ -64,7 +67,7 @@
# Helpers
# ---------------------------------------------------------------------------
def _load_t2i_gen_config(repo_id: str) -> dict:
weights_dir = Path(snapshot_download(repo_id))
weights_dir = Path(hf_api().snapshot_download(repo_id))
cfg_path = weights_dir / "t2i_generation_config.json"
if not cfg_path.exists():
pytest.skip(f"t2i_generation_config.json not found at {cfg_path}")
Expand Down
7 changes: 4 additions & 3 deletions tests/e2e/offline_inference/test_moss_tts_realtime.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project
"""E2E offline inference tests for MOSS-TTS-Realtime (MossTTSRealtime, 1.7B).

Uses the standard omni_runner + pytestmark pattern (one module-scoped engine
Expand Down Expand Up @@ -160,12 +160,13 @@ def _build_request(ref_audio_path: str, text: str) -> dict:
Frees the codec model before returning to avoid competing with the
running vllm engine for GPU/CPU memory.
"""
from huggingface_hub import snapshot_download
from transformers import AutoModel, AutoTokenizer

from vllm_omni.transformers_utils.repo_utils import hf_api

# Step 1: locate the realtime processor module in the snapshot.
try:
snap_dir = Path(snapshot_download(repo_id=MODEL))
snap_dir = Path(hf_api().snapshot_download(repo_id=MODEL))
except Exception as exc:
msg = f"Cannot locate snapshot for {MODEL}: {exc}"
if os.environ.get("MOSS_TTS_SKIP_ON_NET_FAIL"):
Expand Down
Loading
Loading