[diffusion] feat: enhance overlay mechanism (#21648)

This commit is contained in:
Mick
2026-03-30 19:45:34 +08:00
committed by GitHub
parent 1d6424d5ad
commit b76730701b
6 changed files with 112 additions and 63 deletions
+19 -8
View File
@@ -5,10 +5,24 @@ import subprocess
from functools import lru_cache
from sglang.srt.environ import envs
from sglang.utils import (
has_diffusion_overlay_registry_match,
is_known_non_diffusers_diffusion_model,
load_diffusion_overlay_registry_from_env,
)
logger = logging.getLogger(__name__)
@lru_cache(maxsize=1)
def _load_overlay_registry() -> dict:
return load_diffusion_overlay_registry_from_env()
def _is_overlay_diffusion_model(model_path: str) -> bool:
return has_diffusion_overlay_registry_match(model_path, _load_overlay_registry())
def _is_diffusers_model_dir(model_dir: str) -> bool:
"""Check if a local directory contains a valid diffusers model_index.json."""
config_path = os.path.join(model_dir, "model_index.json")
@@ -29,19 +43,16 @@ def get_is_diffusion_model(model_path: str) -> bool:
Returns False on any failure (network error, 404, offline mode, etc.)
so that the caller falls through to the standard LLM server path.
"""
try:
from sglang.multimodal_gen.registry import (
is_known_non_diffusers_multimodal_model,
)
except ImportError:
is_known_non_diffusers_multimodal_model = lambda _: False
if _is_overlay_diffusion_model(model_path):
# short-circuit, if applicable for the overlay mechanism (diffusion-only)
return True
if os.path.isdir(model_path):
if _is_diffusers_model_dir(model_path):
return True
return is_known_non_diffusers_multimodal_model(model_path)
return is_known_non_diffusers_diffusion_model(model_path)
if is_known_non_diffusers_multimodal_model(model_path):
if is_known_non_diffusers_diffusion_model(model_path):
return True
try:
+6 -19
View File
@@ -122,9 +122,9 @@ from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import
)
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
maybe_download_model_index,
verify_model_config_and_directory,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.utils import KNOWN_NON_DIFFUSERS_DIFFUSION_MODEL_PATTERNS
logger = init_logger(__name__)
@@ -327,10 +327,7 @@ def _get_config_info(
return _CONFIG_REGISTRY.get(model_id)
# 3. Use detectors
if os.path.exists(model_path):
config = verify_model_config_and_directory(model_path)
else:
config = maybe_download_model_index(model_path)
config = maybe_download_model_index(model_path)
pipeline_name = config.get("_class_name", "").lower()
matched_model_names = []
@@ -499,10 +496,7 @@ def get_model_info(
else:
# Try to get from model_index.json
try:
if os.path.exists(model_path):
config = verify_model_config_and_directory(model_path)
else:
config = maybe_download_model_index(model_path)
config = maybe_download_model_index(model_path)
except Exception as e:
logger.error(f"Could not read model config for '{model_path}': {e}")
if backend == Backend.AUTO:
@@ -876,25 +870,18 @@ def _register_configs():
_register_configs()
# Known non-diffusers multimodal model patterns
# Maps pattern -> pipeline_name for models that don't have model_index.json
_NON_DIFFUSERS_MULTIMODAL_PATTERNS: Dict[str, str] = {
"hunyuan3d": "Hunyuan3D2Pipeline",
"flux.2-dev-nvfp4": "Flux2NvfpPipeline",
}
def is_known_non_diffusers_multimodal_model(model_path: str) -> bool:
model_path_lower = model_path.lower()
return any(
pattern in model_path_lower for pattern in _NON_DIFFUSERS_MULTIMODAL_PATTERNS
pattern in model_path_lower
for pattern in KNOWN_NON_DIFFUSERS_DIFFUSION_MODEL_PATTERNS
)
def get_non_diffusers_pipeline_name(model_path: str) -> Optional[str]:
"""Get the pipeline name for a known non-diffusers model."""
model_path_lower = model_path.lower()
for pattern, pipeline_name in _NON_DIFFUSERS_MULTIMODAL_PATTERNS.items():
for pattern, pipeline_name in KNOWN_NON_DIFFUSERS_DIFFUSION_MODEL_PATTERNS.items():
if pattern in model_path_lower:
return pipeline_name
return None
@@ -489,15 +489,16 @@ def maybe_download_model_index(model_name_or_path: str) -> dict[str, Any]:
from huggingface_hub.errors import EntryNotFoundError
# If it's a local path, verify it directly
overlay_config = maybe_load_overlay_model_index(
model_name_or_path,
snapshot_download_fn=snapshot_download,
hf_hub_download_fn=hf_hub_download,
)
if overlay_config is not None:
return overlay_config
# If it's a local path, verify it directly.
if os.path.exists(model_name_or_path):
overlay_config = maybe_load_overlay_model_index(
model_name_or_path,
snapshot_download_fn=snapshot_download,
hf_hub_download_fn=hf_hub_download,
)
if overlay_config is not None:
return overlay_config
try:
return verify_model_config_and_directory(model_name_or_path)
except ValueError:
@@ -509,15 +510,6 @@ def maybe_download_model_index(model_name_or_path: str) -> dict[str, Any]:
return config
raise
# return resolved overlay config if applicable
overlay_config = maybe_load_overlay_model_index(
model_name_or_path,
snapshot_download_fn=snapshot_download,
hf_hub_download_fn=hf_hub_download,
)
if overlay_config is not None:
return overlay_config
# For remote models, download just the model_index.json
try:
with tempfile.TemporaryDirectory() as tmp_dir:
@@ -19,6 +19,7 @@ from requests.exceptions import RequestException
from sglang.multimodal_gen.runtime.loader.weight_utils import get_lock
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.utils import load_diffusion_overlay_registry_from_env
logger = init_logger(__name__)
@@ -60,28 +61,12 @@ def _load_model_overlay_registry() -> dict[str, dict[str, Any]]:
# Built-in registry is the stable default path; env only overrides it.
normalized = _normalize_model_overlay_registry(BUILTIN_MODEL_OVERLAY_REGISTRY)
raw_value = os.getenv("SGLANG_DIFFUSION_MODEL_OVERLAY_REGISTRY", "").strip()
if not raw_value:
env_registry = load_diffusion_overlay_registry_from_env()
if not env_registry:
_MODEL_OVERLAY_REGISTRY_CACHE = normalized
return _MODEL_OVERLAY_REGISTRY_CACHE
try:
if raw_value.startswith("{"):
payload = json.loads(raw_value)
else:
with open(os.path.expanduser(raw_value), encoding="utf-8") as f:
payload = json.load(f)
except Exception as exc:
raise ValueError(
"Failed to parse SGLANG_DIFFUSION_MODEL_OVERLAY_REGISTRY"
) from exc
if not isinstance(payload, dict):
raise ValueError(
"SGLANG_DIFFUSION_MODEL_OVERLAY_REGISTRY must be a JSON object"
)
normalized.update(_normalize_model_overlay_registry(payload))
normalized.update(_normalize_model_overlay_registry(env_registry))
_MODEL_OVERLAY_REGISTRY_CACHE = normalized
return _MODEL_OVERLAY_REGISTRY_CACHE
@@ -462,11 +447,23 @@ def materialize_overlay_model(
):
return final_dir
logger.info(
"Materializing overlay model for %s into %s",
source_model_id,
final_dir,
)
logger.info(
"Overlay source repo: %s, overlay repo: %s@%s",
source_model_id,
overlay_repo_id,
overlay_revision,
)
tmp_dir = final_dir + ".tmp"
if os.path.exists(tmp_dir):
shutil.rmtree(tmp_dir)
if os.path.exists(final_dir):
shutil.rmtree(final_dir)
logger.info("Copying overlay metadata into temporary materialized directory")
shutil.copytree(
overlay_dir,
tmp_dir,
@@ -479,11 +476,17 @@ def materialize_overlay_model(
file_mappings = manifest.get("file_mappings", [])
if file_mappings:
logger.info("Applying %d overlay file mappings", len(file_mappings))
_apply_overlay_file_mappings(
source_dir=source_dir,
output_dir=tmp_dir,
file_mappings=cast(list[dict[str, Any]], file_mappings),
)
if manifest.get("custom_materializer"):
logger.info(
"Running custom overlay materializer: %s",
manifest["custom_materializer"],
)
_run_overlay_custom_materializer(
overlay_dir=overlay_dir,
source_dir=source_dir,
@@ -506,6 +509,7 @@ def materialize_overlay_model(
)
os.replace(tmp_dir, final_dir)
logger.info("Overlay materialization finished: %s", final_dir)
return final_dir
+50
View File
@@ -31,6 +31,56 @@ from sglang.srt.environ import envs
logger = logging.getLogger(__name__)
KNOWN_NON_DIFFUSERS_DIFFUSION_MODEL_PATTERNS: dict[str, str] = {
"hunyuan3d": "Hunyuan3D2Pipeline",
"flux.2-dev-nvfp4": "Flux2NvfpPipeline",
}
def load_diffusion_overlay_registry_from_env() -> dict[str, dict[str, Any]]:
raw_value = os.getenv("SGLANG_DIFFUSION_MODEL_OVERLAY_REGISTRY", "").strip()
if not raw_value:
return {}
if raw_value.startswith("{"):
payload = json.loads(raw_value)
else:
with open(os.path.expanduser(raw_value), encoding="utf-8") as f:
payload = json.load(f)
if not isinstance(payload, dict):
return {}
normalized: dict[str, dict[str, Any]] = {}
for source_model_id, spec in payload.items():
if isinstance(spec, str):
normalized[source_model_id] = {"overlay_repo_id": spec}
elif isinstance(spec, dict) and spec.get("overlay_repo_id"):
normalized[source_model_id] = dict(spec)
return normalized
def has_diffusion_overlay_registry_match(
model_path: str, registry: dict[str, dict[str, Any]] | None = None
) -> bool:
registry = (
load_diffusion_overlay_registry_from_env() if registry is None else registry
)
if model_path in registry:
return True
if not os.path.exists(model_path):
return False
base_name = os.path.basename(os.path.normpath(model_path))
return any(base_name == key.rsplit("/", 1)[-1] for key in registry)
def is_known_non_diffusers_diffusion_model(model_path: str) -> bool:
model_path_lower = model_path.lower()
return any(
pattern in model_path_lower
for pattern in KNOWN_NON_DIFFUSERS_DIFFUSION_MODEL_PATTERNS
)
def execute_once(func):
has_run = None