[diffusion] model: support Hunyuan3D-2 (#18170)
Co-authored-by: yingluosanqian <yingluosanqian@gmail.com> Co-authored-by: daiweitao <dwti614707404@163.com> Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
co-authored by
yingluosanqian
daiweitao
Mick
parent
f6ee6dc8c3
commit
57c5c343d7
@@ -45,6 +45,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.flux import (
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.glm_image import (
|
||||
GlmImagePipelineConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.hunyuan3d import (
|
||||
Hunyuan3D2PipelineConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import LTX2PipelineConfig
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.mova import (
|
||||
MOVA360PConfig,
|
||||
@@ -75,6 +78,7 @@ from sglang.multimodal_gen.configs.sample.hunyuan import (
|
||||
FastHunyuanSamplingParam,
|
||||
HunyuanSamplingParams,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.sample.hunyuan3d import Hunyuan3DSamplingParams
|
||||
from sglang.multimodal_gen.configs.sample.ltx_2 import LTX2SamplingParams
|
||||
from sglang.multimodal_gen.configs.sample.mova import (
|
||||
MOVA_360P_SamplingParams,
|
||||
@@ -367,26 +371,36 @@ def get_model_info(
|
||||
# 1. Discover all available pipeline classes and cache them
|
||||
_discover_and_register_pipelines()
|
||||
|
||||
# 2. Get pipeline class from model's 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)
|
||||
except Exception as e:
|
||||
logger.error(f"Could not read model config for '{model_path}': {e}")
|
||||
if backend == Backend.AUTO:
|
||||
logger.info("Falling back to diffusers backend")
|
||||
return _get_diffusers_model_info(model_path)
|
||||
return None
|
||||
# 2. Get pipeline class - check non-diffusers models first
|
||||
pipeline_class_name = get_non_diffusers_pipeline_name(model_path)
|
||||
if pipeline_class_name:
|
||||
# Known non-diffusers model, skip model_index.json download
|
||||
logger.debug(
|
||||
f"Using registered pipeline '{pipeline_class_name}' for non-diffusers model '{model_path}'"
|
||||
)
|
||||
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)
|
||||
except Exception as e:
|
||||
logger.error(f"Could not read model config for '{model_path}': {e}")
|
||||
if backend == Backend.AUTO:
|
||||
logger.info("Falling back to diffusers backend")
|
||||
return _get_diffusers_model_info(model_path)
|
||||
return None
|
||||
|
||||
pipeline_class_name = config.get("_class_name")
|
||||
if not pipeline_class_name:
|
||||
logger.error(f"'_class_name' not found in model_index.json for '{model_path}'")
|
||||
if backend == Backend.AUTO:
|
||||
logger.info("Falling back to diffusers backend")
|
||||
return _get_diffusers_model_info(model_path)
|
||||
return None
|
||||
pipeline_class_name = config.get("_class_name")
|
||||
if not pipeline_class_name:
|
||||
logger.error(
|
||||
f"'_class_name' not found in model_index.json for '{model_path}'"
|
||||
)
|
||||
if backend == Backend.AUTO:
|
||||
logger.info("Falling back to diffusers backend")
|
||||
return _get_diffusers_model_info(model_path)
|
||||
return None
|
||||
|
||||
pipeline_cls = _PIPELINE_REGISTRY.get(pipeline_class_name)
|
||||
if not pipeline_cls:
|
||||
@@ -453,7 +467,7 @@ def _register_configs():
|
||||
hf_model_paths=[
|
||||
"hunyuanvideo-community/HunyuanVideo",
|
||||
],
|
||||
model_detectors=[lambda hf_id: "hunyuan" in hf_id.lower()],
|
||||
model_detectors=[lambda hf_id: "hunyuanvideo" in hf_id.lower()],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=FastHunyuanSamplingParam,
|
||||
@@ -658,6 +672,37 @@ def _register_configs():
|
||||
pipeline_config_cls=GlmImagePipelineConfig,
|
||||
model_detectors=[lambda hf_id: "glm-image" in hf_id.lower()],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=Hunyuan3DSamplingParams,
|
||||
pipeline_config_cls=Hunyuan3D2PipelineConfig,
|
||||
hf_model_paths=[
|
||||
"tencent/Hunyuan3D-2",
|
||||
],
|
||||
model_detectors=[lambda hf_id: "hunyuan3d" in hf_id.lower()],
|
||||
)
|
||||
|
||||
|
||||
_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",
|
||||
}
|
||||
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
|
||||
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():
|
||||
if pattern in model_path_lower:
|
||||
return pipeline_name
|
||||
return None
|
||||
|
||||
Reference in New Issue
Block a user