[diffusion] fix: fix gated repo failing the generate cmd (#22040)

Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Mick
2026-04-04 00:43:11 +08:00
committed by GitHub
co-authored by gemini-code-assist[bot] Claude Opus 4.6
parent 896ea75820
commit 151f727163
+26 -1
View File
@@ -4,6 +4,8 @@ import os
import subprocess import subprocess
from functools import lru_cache from functools import lru_cache
from huggingface_hub import HfApi
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.utils import ( from sglang.utils import (
has_diffusion_overlay_registry_match, has_diffusion_overlay_registry_match,
@@ -23,6 +25,16 @@ def _is_overlay_diffusion_model(model_path: str) -> bool:
return has_diffusion_overlay_registry_match(model_path, _load_overlay_registry()) return has_diffusion_overlay_registry_match(model_path, _load_overlay_registry())
def _is_registered_diffusion_model(model_path: str) -> bool:
try:
# if diffusion dependencies are not installed
from sglang.multimodal_gen.registry import get_model_info
except ImportError:
return False
return get_model_info(model_path, backend="sglang") is not None
def _is_diffusers_model_dir(model_dir: str) -> bool: def _is_diffusers_model_dir(model_dir: str) -> bool:
"""Check if a local directory contains a valid diffusers model_index.json.""" """Check if a local directory contains a valid diffusers model_index.json."""
config_path = os.path.join(model_dir, "model_index.json") config_path = os.path.join(model_dir, "model_index.json")
@@ -35,11 +47,22 @@ def _is_diffusers_model_dir(model_dir: str) -> bool:
return "_diffusers_version" in config return "_diffusers_version" in config
def _is_gated_diffusion_repo(repo_id: str) -> bool:
"""Query HF model card metadata to check if a gated repo is a diffusers model."""
try:
info = HfApi().model_info(repo_id)
return getattr(info, "library_name", None) == "diffusers"
except Exception:
return False
def get_is_diffusion_model(model_path: str) -> bool: def get_is_diffusion_model(model_path: str) -> bool:
"""Detect whether model_path points to a diffusion model. """Detect whether model_path points to a diffusion model.
For local directories, checks the filesystem directly. For local directories, checks the filesystem directly.
For HF/ModelScope model IDs, attempts to fetch only model_index.json. For HF/ModelScope model IDs, attempts to fetch only model_index.json.
For gated repos where file download fails, falls back to HF model card
metadata (library_name == "diffusers").
Returns False on any failure (network error, 404, offline mode, etc.) Returns False on any failure (network error, 404, offline mode, etc.)
so that the caller falls through to the standard LLM server path. so that the caller falls through to the standard LLM server path.
""" """
@@ -70,7 +93,9 @@ def get_is_diffusion_model(model_path: str) -> bool:
return _is_diffusers_model_dir(os.path.dirname(file_path)) return _is_diffusers_model_dir(os.path.dirname(file_path))
except Exception as e: except Exception as e:
logger.debug("Failed to auto-detect diffusion model for %s: %s", model_path, e) logger.debug("Failed to auto-detect diffusion model for %s: %s", model_path, e)
return False # For gated repos, file download fails but model card is still accessible.
# Check library_name from HF metadata as a fallback.
return _is_gated_diffusion_repo(model_path)
def get_model_path(extra_argv): def get_model_path(extra_argv):