[diffusion] feat: add ModelScope support (#17924)
This commit is contained in:
@@ -8,7 +8,8 @@ from functools import lru_cache
|
|||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
import filelock
|
import filelock
|
||||||
from huggingface_hub import hf_hub_download
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -43,6 +44,8 @@ def _maybe_download_model(
|
|||||||
Local directory path that contains the downloaded config file, or the original local directory.
|
Local directory path that contains the downloaded config file, or the original local directory.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import hf_hub_download
|
||||||
|
|
||||||
if os.path.exists(model_name_or_path):
|
if os.path.exists(model_name_or_path):
|
||||||
logger.info("Model already exists locally")
|
logger.info("Model already exists locally")
|
||||||
return model_name_or_path
|
return model_name_or_path
|
||||||
@@ -52,9 +55,11 @@ def _maybe_download_model(
|
|||||||
|
|
||||||
with _get_lock(model_name_or_path):
|
with _get_lock(model_name_or_path):
|
||||||
# Try `model_index.json` first (diffusers models)
|
# Try `model_index.json` first (diffusers models)
|
||||||
|
source_hub = "MS Hub" if envs.SGLANG_USE_MODELSCOPE.get() else "HF Hub"
|
||||||
try:
|
try:
|
||||||
logger.info(
|
logger.info(
|
||||||
"Downloading model_index.json from HF Hub for %s...",
|
"Downloading model_index.json from %s for %s...",
|
||||||
|
source_hub,
|
||||||
model_name_or_path,
|
model_name_or_path,
|
||||||
)
|
)
|
||||||
file_path = hf_hub_download(
|
file_path = hf_hub_download(
|
||||||
@@ -70,7 +75,9 @@ def _maybe_download_model(
|
|||||||
# Fallback to `config.json`
|
# Fallback to `config.json`
|
||||||
try:
|
try:
|
||||||
logger.info(
|
logger.info(
|
||||||
"Downloading config.json from HF Hub for %s...", model_name_or_path
|
"Downloading config.json from %s for %s...",
|
||||||
|
source_hub,
|
||||||
|
model_name_or_path,
|
||||||
)
|
)
|
||||||
file_path = hf_hub_download(
|
file_path = hf_hub_download(
|
||||||
repo_id=model_name_or_path,
|
repo_id=model_name_or_path,
|
||||||
@@ -83,9 +90,9 @@ def _maybe_download_model(
|
|||||||
raise ValueError(
|
raise ValueError(
|
||||||
(
|
(
|
||||||
"Could not find model locally at %s and failed to download "
|
"Could not find model locally at %s and failed to download "
|
||||||
"model_index.json/config.json from HF Hub: %s"
|
"model_index.json/config.json from %s: %s"
|
||||||
)
|
)
|
||||||
% (model_name_or_path, e_config)
|
% (model_name_or_path, source_hub, e_config)
|
||||||
) from e_config
|
) from e_config
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
|||||||
from sglang.multimodal_gen.runtime.scheduler_client import sync_scheduler_client
|
from sglang.multimodal_gen.runtime.scheduler_client import sync_scheduler_client
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
@@ -27,13 +28,21 @@ def run_sgl_diffusion_webui(server_args: ServerArgs):
|
|||||||
# import gradio in function to avoid CI crash
|
# import gradio in function to avoid CI crash
|
||||||
|
|
||||||
import gradio as gr
|
import gradio as gr
|
||||||
from huggingface_hub import model_info
|
|
||||||
|
if envs.SGLANG_USE_MODELSCOPE.get():
|
||||||
|
from modelscope.hub.api import HubApi
|
||||||
|
|
||||||
|
api = HubApi()
|
||||||
|
model_info_obj = api.model_info(server_args.model_path)
|
||||||
|
task_name = model_info_obj.tasks[0]["Name"].replace("-synthesis", "")
|
||||||
|
else:
|
||||||
|
from huggingface_hub import model_info
|
||||||
|
|
||||||
|
task_name = model_info(server_args.model_path).pipeline_tag
|
||||||
|
|
||||||
# init client
|
# init client
|
||||||
sync_scheduler_client.initialize(server_args)
|
sync_scheduler_client.initialize(server_args)
|
||||||
|
|
||||||
task_name = model_info(server_args.model_path).pipeline_tag
|
|
||||||
|
|
||||||
if task_name in ("text-to-video", "image-to-video", "video-to-video"):
|
if task_name in ("text-to-video", "image-to-video", "video-to-video"):
|
||||||
task_type = "video"
|
task_type = "video"
|
||||||
elif task_name in ["text-to-image", "image-to-image"]:
|
elif task_name in ["text-to-image", "image-to-image"]:
|
||||||
|
|||||||
@@ -26,12 +26,11 @@ import shutil
|
|||||||
import time
|
import time
|
||||||
from functools import reduce
|
from functools import reduce
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Optional, cast
|
from typing import Any, Optional, Union, cast
|
||||||
|
|
||||||
from diffusers.loaders.lora_base import (
|
from diffusers.loaders.lora_base import (
|
||||||
_best_guess_weight_name, # watch out for potetential removal from diffusers
|
_best_guess_weight_name, # watch out for potetential removal from diffusers
|
||||||
)
|
)
|
||||||
from huggingface_hub import snapshot_download
|
|
||||||
from huggingface_hub.errors import (
|
from huggingface_hub.errors import (
|
||||||
LocalEntryNotFoundError,
|
LocalEntryNotFoundError,
|
||||||
RepositoryNotFoundError,
|
RepositoryNotFoundError,
|
||||||
@@ -45,6 +44,7 @@ from transformers.models.auto.modeling_auto import MODEL_FOR_CAUSAL_LM_MAPPING_N
|
|||||||
from sglang.multimodal_gen.runtime.loader.weight_utils import get_lock
|
from sglang.multimodal_gen.runtime.loader.weight_utils import get_lock
|
||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.utils import is_in_ci
|
from sglang.utils import is_in_ci
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
@@ -467,7 +467,6 @@ def maybe_download_model_index(model_name_or_path: str) -> dict[str, Any]:
|
|||||||
"""
|
"""
|
||||||
import tempfile
|
import tempfile
|
||||||
|
|
||||||
from huggingface_hub import hf_hub_download
|
|
||||||
from huggingface_hub.errors import EntryNotFoundError
|
from huggingface_hub.errors import EntryNotFoundError
|
||||||
|
|
||||||
# If it's a local path, verify it directly
|
# If it's a local path, verify it directly
|
||||||
@@ -638,9 +637,7 @@ def maybe_download_model(
|
|||||||
ignore_patterns=["*.onnx", "*.msgpack"],
|
ignore_patterns=["*.onnx", "*.msgpack"],
|
||||||
local_dir=local_dir,
|
local_dir=local_dir,
|
||||||
local_files_only=True,
|
local_files_only=True,
|
||||||
resume_download=True,
|
|
||||||
max_workers=8,
|
max_workers=8,
|
||||||
etag_timeout=60,
|
|
||||||
)
|
)
|
||||||
if is_lora or _verify_model_complete(local_path):
|
if is_lora or _verify_model_complete(local_path):
|
||||||
# CI validation: check all subdirectories for missing shards
|
# CI validation: check all subdirectories for missing shards
|
||||||
@@ -709,9 +706,7 @@ def maybe_download_model(
|
|||||||
ignore_patterns=["*.onnx", "*.msgpack"],
|
ignore_patterns=["*.onnx", "*.msgpack"],
|
||||||
allow_patterns=allow_patterns,
|
allow_patterns=allow_patterns,
|
||||||
local_dir=local_dir,
|
local_dir=local_dir,
|
||||||
resume_download=True,
|
|
||||||
max_workers=8,
|
max_workers=8,
|
||||||
etag_timeout=120,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# Verify downloaded model is complete (skip for LoRA)
|
# Verify downloaded model is complete (skip for LoRA)
|
||||||
@@ -725,9 +720,7 @@ def maybe_download_model(
|
|||||||
repo_id=model_name_or_path,
|
repo_id=model_name_or_path,
|
||||||
ignore_patterns=["*.onnx", "*.msgpack"],
|
ignore_patterns=["*.onnx", "*.msgpack"],
|
||||||
local_dir=local_dir,
|
local_dir=local_dir,
|
||||||
resume_download=True,
|
|
||||||
max_workers=8,
|
max_workers=8,
|
||||||
etag_timeout=60,
|
|
||||||
force_download=True,
|
force_download=True,
|
||||||
)
|
)
|
||||||
if not _verify_model_complete(local_path):
|
if not _verify_model_complete(local_path):
|
||||||
@@ -775,3 +768,71 @@ def maybe_download_model(
|
|||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Could not find model at {model_name_or_path} and failed to download from HF Hub: {e}"
|
f"Could not find model at {model_name_or_path} and failed to download from HF Hub: {e}"
|
||||||
) from e
|
) from e
|
||||||
|
|
||||||
|
|
||||||
|
# Unified download functions with Hugging Face-compatible names
|
||||||
|
def hf_hub_download(
|
||||||
|
repo_id: str,
|
||||||
|
filename: str,
|
||||||
|
local_dir: Optional[Union[str, Path]] = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> str:
|
||||||
|
"""Unified hf_hub_download that supports both Hugging Face Hub and ModelScope."""
|
||||||
|
if envs.SGLANG_USE_MODELSCOPE.get():
|
||||||
|
from modelscope import model_file_download
|
||||||
|
|
||||||
|
return model_file_download(
|
||||||
|
model_id=repo_id,
|
||||||
|
file_path=filename,
|
||||||
|
cache_dir=local_dir,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
from huggingface_hub import hf_hub_download as _hf_hub_download
|
||||||
|
|
||||||
|
return _hf_hub_download(
|
||||||
|
repo_id=repo_id,
|
||||||
|
filename=filename,
|
||||||
|
local_dir=local_dir,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def snapshot_download(
|
||||||
|
repo_id: str,
|
||||||
|
local_dir: Optional[Union[str, Path]] = None,
|
||||||
|
ignore_patterns: Optional[Union[list[str], str]] = None,
|
||||||
|
allow_patterns: Optional[Union[list[str], str]] = None,
|
||||||
|
local_files_only: bool = False,
|
||||||
|
max_workers: int = 8,
|
||||||
|
**kwargs,
|
||||||
|
) -> str:
|
||||||
|
"""Unified snapshot_download that supports both Hugging Face Hub and ModelScope."""
|
||||||
|
if envs.SGLANG_USE_MODELSCOPE.get():
|
||||||
|
from modelscope import snapshot_download as _ms_snapshot_download
|
||||||
|
|
||||||
|
ms_kwargs = {
|
||||||
|
"model_id": repo_id,
|
||||||
|
"local_dir": local_dir,
|
||||||
|
"ignore_patterns": ignore_patterns,
|
||||||
|
"allow_patterns": allow_patterns,
|
||||||
|
"local_files_only": local_files_only,
|
||||||
|
"max_workers": max_workers,
|
||||||
|
}
|
||||||
|
ms_kwargs.update(kwargs)
|
||||||
|
return _ms_snapshot_download(**ms_kwargs)
|
||||||
|
else:
|
||||||
|
from huggingface_hub import snapshot_download as _hf_snapshot_download
|
||||||
|
|
||||||
|
hf_kwargs = {
|
||||||
|
"repo_id": repo_id,
|
||||||
|
"local_dir": local_dir,
|
||||||
|
"ignore_patterns": ignore_patterns,
|
||||||
|
"allow_patterns": allow_patterns,
|
||||||
|
"local_files_only": local_files_only,
|
||||||
|
"max_workers": max_workers,
|
||||||
|
"resume_download": True,
|
||||||
|
"etag_timeout": 60,
|
||||||
|
}
|
||||||
|
hf_kwargs.update(kwargs)
|
||||||
|
return _hf_snapshot_download(**hf_kwargs)
|
||||||
|
|||||||
Reference in New Issue
Block a user