[diffusion] chore: improve model info registration and searching strategy (#14281)

Co-authored-by: Jinyan Chen <jinyanc@nvidia.com>
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
Jinyan Chen
2025-12-02 18:28:59 +08:00
committed by GitHub
co-authored by Jinyan Chen Mick
parent e0ec42c710
commit 3067b3f050
4 changed files with 65 additions and 95 deletions
@@ -489,18 +489,9 @@ class PipelineConfig:
# 1. Get the pipeline config class from the registry # 1. Get the pipeline config class from the registry
model_info = get_model_info(model_path) model_info = get_model_info(model_path)
# 2. Instantiate PipelineConfig
if model_info is None:
# The error is already logged in get_model_info.
# We raise an exception here to stop the execution.
raise ValueError(
f"Failed to get model info for '{model_path}'. "
"Please check the model path and ensure it is registered correctly."
)
pipeline_config = model_info.pipeline_config_cls() pipeline_config = model_info.pipeline_config_cls()
# 3. Load PipelineConfig from a json file or a PipelineConfig object if provided # 2. Load PipelineConfig from a json file or a PipelineConfig object if provided
if isinstance(pipeline_config_or_path, str): if isinstance(pipeline_config_or_path, str):
pipeline_config.load_from_json(pipeline_config_or_path) pipeline_config.load_from_json(pipeline_config_or_path)
kwargs[prefix_with_dot + "pipeline_config_path"] = pipeline_config_or_path kwargs[prefix_with_dot + "pipeline_config_path"] = pipeline_config_or_path
@@ -509,7 +500,7 @@ class PipelineConfig:
elif isinstance(pipeline_config_or_path, dict): elif isinstance(pipeline_config_or_path, dict):
pipeline_config.update_pipeline_config(pipeline_config_or_path) pipeline_config.update_pipeline_config(pipeline_config_or_path)
# 4. Update PipelineConfig from CLI arguments if provided # 3. Update PipelineConfig from CLI arguments if provided
kwargs[prefix_with_dot + "model_path"] = model_path kwargs[prefix_with_dot + "model_path"] = model_path
pipeline_config.update_config_from_dict(kwargs, config_cli_prefix) pipeline_config.update_config_from_dict(kwargs, config_cli_prefix)
return pipeline_config return pipeline_config
@@ -299,14 +299,7 @@ class SamplingParams:
from sglang.multimodal_gen.registry import get_model_info from sglang.multimodal_gen.registry import get_model_info
model_info = get_model_info(model_path) model_info = get_model_info(model_path)
if model_info is not None:
sampling_params: SamplingParams = model_info.sampling_param_cls(**kwargs) sampling_params: SamplingParams = model_info.sampling_param_cls(**kwargs)
else:
logger.warning(
"Couldn't find an optimal sampling param for %s. Using the default sampling param.",
model_path,
)
sampling_params = cls(**kwargs)
return sampling_params return sampling_params
@staticmethod @staticmethod
+62 -74
View File
@@ -124,42 +124,38 @@ class ConfigInfo:
_CONFIG_REGISTRY: Dict[str, ConfigInfo] = {} _CONFIG_REGISTRY: Dict[str, ConfigInfo] = {}
# Mappings from Hugging Face model paths to our internal model names # Mappings from Hugging Face model paths to our internal model names
_MODEL_PATH_TO_NAME: Dict[str, str] = {} _MODEL_HF_PATH_TO_NAME: Dict[str, str] = {}
# Detectors to identify model families from paths or class names # Detectors to identify model families from paths or class names
_MODEL_NAME_DETECTORS: List[Tuple[str, Callable[[str], bool]]] = [] _MODEL_NAME_DETECTORS: List[Tuple[str, Callable[[str], bool]]] = []
def register_configs( def register_configs(
model_name: str,
sampling_param_cls: Any, sampling_param_cls: Any,
pipeline_config_cls: Type[PipelineConfig], pipeline_config_cls: Type[PipelineConfig],
model_paths: Optional[List[str]] = None, hf_model_paths: Optional[List[str]] = None,
model_detectors: Optional[List[Callable[[str], bool]]] = None, model_detectors: Optional[List[Callable[[str], bool]]] = None,
): ):
""" """
Registers configuration classes for a new model family. Registers configuration classes for a new model family.
""" """
if model_name in _CONFIG_REGISTRY: model_id = str(len(_CONFIG_REGISTRY))
logger.warning(
f"Config for model '{model_name}' is already registered and will be overwritten."
)
_CONFIG_REGISTRY[model_name] = ConfigInfo( _CONFIG_REGISTRY[model_id] = ConfigInfo(
sampling_param_cls=sampling_param_cls, sampling_param_cls=sampling_param_cls,
pipeline_config_cls=pipeline_config_cls, pipeline_config_cls=pipeline_config_cls,
) )
if model_paths: if hf_model_paths:
for path in model_paths: for path in hf_model_paths:
if path in _MODEL_PATH_TO_NAME: if path in _MODEL_HF_PATH_TO_NAME:
logger.warning( logger.warning(
f"Model path '{path}' is already mapped to '{_MODEL_PATH_TO_NAME[path]}' and will be overwritten by '{model_name}'." f"Model path '{path}' is already mapped to '{_MODEL_HF_PATH_TO_NAME[path]}' and will be overwritten by '{model_id}'."
) )
_MODEL_PATH_TO_NAME[path] = model_name _MODEL_HF_PATH_TO_NAME[path] = model_id
if model_detectors: if model_detectors:
for detector in model_detectors: for detector in model_detectors:
_MODEL_NAME_DETECTORS.append((model_name, detector)) _MODEL_NAME_DETECTORS.append((model_id, detector))
def _get_config_info(model_path: str) -> Optional[ConfigInfo]: def _get_config_info(model_path: str) -> Optional[ConfigInfo]:
@@ -167,18 +163,21 @@ def _get_config_info(model_path: str) -> Optional[ConfigInfo]:
Gets the ConfigInfo for a given model path using mappings and detectors. Gets the ConfigInfo for a given model path using mappings and detectors.
""" """
# 1. Exact match # 1. Exact match
if model_path in _MODEL_PATH_TO_NAME: if model_path in _MODEL_HF_PATH_TO_NAME:
model_name = _MODEL_PATH_TO_NAME[model_path] model_id = _MODEL_HF_PATH_TO_NAME[model_path]
logger.debug(f"Resolved model name '{model_name}' from exact path match.") logger.debug(f"Resolved model path '{model_path}' from exact path match.")
return _CONFIG_REGISTRY.get(model_name) return _CONFIG_REGISTRY.get(model_id)
# 2. Partial match: find the best (longest) match against all registered model names. # 2. Partial match: find the best (longest) match against all registered model hf paths.
cleaned_model_path = re.sub(r"--", "/", model_path.lower()) cleaned_model_path = re.sub(r"--", "/", model_path.lower())
all_model_names = sorted(_CONFIG_REGISTRY.keys(), key=len, reverse=True) all_model_hf_paths = sorted(_MODEL_HF_PATH_TO_NAME.keys(), key=len, reverse=True)
for model_name in all_model_names: for model_hf_path in all_model_hf_paths:
if model_name in cleaned_model_path: if model_hf_path.lower() in cleaned_model_path:
logger.debug(f"Resolved model name '{model_name}' from partial path match.") logger.debug(
return _CONFIG_REGISTRY.get(model_name) f"Resolved model name '{model_hf_path}' from partial path match."
)
model_id = _MODEL_HF_PATH_TO_NAME[model_hf_path]
return _CONFIG_REGISTRY.get(model_id)
# 3. Use detectors # 3. Use detectors
if os.path.exists(model_path): if os.path.exists(model_path):
@@ -188,14 +187,23 @@ def _get_config_info(model_path: str) -> Optional[ConfigInfo]:
pipeline_name = config.get("_class_name", "").lower() pipeline_name = config.get("_class_name", "").lower()
for model_name, detector in _MODEL_NAME_DETECTORS: matched_model_names = []
for model_id, detector in _MODEL_NAME_DETECTORS:
if detector(model_path.lower()) or detector(pipeline_name): if detector(model_path.lower()) or detector(pipeline_name):
logger.debug( logger.debug(
f"Resolved model name '{model_name}' using a registered detector." f"Matched model name '{model_id}' using a registered detector."
) )
return _CONFIG_REGISTRY.get(model_name) matched_model_names += [model_id]
return None if len(matched_model_names) >= 1:
if len(matched_model_names) > 1:
logger.warning(
f"More than one model name is matched, using the first matched"
)
model_id = matched_model_names[0]
return _CONFIG_REGISTRY.get(model_id)
else:
raise RuntimeError(f"No model info found for model path: {model_path}")
# --- Part 3: Main Resolver --- # --- Part 3: Main Resolver ---
@@ -256,7 +264,7 @@ def get_model_info(model_path: str) -> Optional[ModelInfo]:
logger.error( logger.error(
f"Could not resolve configuration for model '{model_path}'. " f"Could not resolve configuration for model '{model_path}'. "
"It is not a registered model path or detected by any registered model family detectors. " "It is not a registered model path or detected by any registered model family detectors. "
f"Known model paths: {list(_MODEL_PATH_TO_NAME.keys())}" f"Known model paths: {list(_MODEL_HF_PATH_TO_NAME.keys())}"
) )
return None return None
@@ -275,160 +283,140 @@ def get_model_info(model_path: str) -> Optional[ModelInfo]:
def _register_configs(): def _register_configs():
# Hunyuan # Hunyuan
register_configs( register_configs(
model_name="hunyuan",
sampling_param_cls=HunyuanSamplingParams, sampling_param_cls=HunyuanSamplingParams,
pipeline_config_cls=HunyuanConfig, pipeline_config_cls=HunyuanConfig,
model_paths=[ hf_model_paths=[
"hunyuanvideo-community/HunyuanVideo", "hunyuanvideo-community/HunyuanVideo",
], ],
model_detectors=[lambda id: "hunyuan" in id.lower()], model_detectors=[lambda hf_id: "hunyuan" in hf_id.lower()],
) )
register_configs( register_configs(
model_name="fasthunyuan",
sampling_param_cls=FastHunyuanSamplingParam, sampling_param_cls=FastHunyuanSamplingParam,
pipeline_config_cls=FastHunyuanConfig, pipeline_config_cls=FastHunyuanConfig,
model_paths=[ hf_model_paths=[
"FastVideo/FastHunyuan-diffusers", "FastVideo/FastHunyuan-diffusers",
], ],
) )
# StepVideo # StepVideo
register_configs( register_configs(
model_name="stepvideo",
sampling_param_cls=StepVideoT2VSamplingParams, sampling_param_cls=StepVideoT2VSamplingParams,
pipeline_config_cls=StepVideoT2VConfig, pipeline_config_cls=StepVideoT2VConfig,
model_paths=[ hf_model_paths=[
"FastVideo/stepvideo-t2v-diffusers", "FastVideo/stepvideo-t2v-diffusers",
], ],
model_detectors=[lambda id: "stepvideo" in id.lower()], model_detectors=[lambda hf_id: "stepvideo" in hf_id.lower()],
) )
# Wan # Wan
register_configs( register_configs(
model_name="wan-t2v-1.3b",
sampling_param_cls=WanT2V_1_3B_SamplingParams, sampling_param_cls=WanT2V_1_3B_SamplingParams,
pipeline_config_cls=WanT2V480PConfig, pipeline_config_cls=WanT2V480PConfig,
model_paths=[ hf_model_paths=[
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
], ],
model_detectors=[lambda id: "wanpipeline" in id.lower()], model_detectors=[lambda hf_id: "wanpipeline" in hf_id.lower()],
) )
register_configs( register_configs(
model_name="wan-t2v-14b",
sampling_param_cls=WanT2V_14B_SamplingParams, sampling_param_cls=WanT2V_14B_SamplingParams,
pipeline_config_cls=WanT2V720PConfig, pipeline_config_cls=WanT2V720PConfig,
model_paths=[ hf_model_paths=[
"Wan-AI/Wan2.1-T2V-14B-Diffusers", "Wan-AI/Wan2.1-T2V-14B-Diffusers",
], ],
) )
register_configs( register_configs(
model_name="wan-i2v-14b-480p",
sampling_param_cls=WanI2V_14B_480P_SamplingParam, sampling_param_cls=WanI2V_14B_480P_SamplingParam,
pipeline_config_cls=WanI2V480PConfig, pipeline_config_cls=WanI2V480PConfig,
model_paths=[ hf_model_paths=[
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers", "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
], ],
model_detectors=[lambda id: "wanimagetovideo" in id.lower()], model_detectors=[lambda hf_id: "wanimagetovideo" in hf_id.lower()],
) )
register_configs( register_configs(
model_name="wan-i2v-14b-720p",
sampling_param_cls=WanI2V_14B_720P_SamplingParam, sampling_param_cls=WanI2V_14B_720P_SamplingParam,
pipeline_config_cls=WanI2V720PConfig, pipeline_config_cls=WanI2V720PConfig,
model_paths=[ hf_model_paths=[
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers", "Wan-AI/Wan2.1-I2V-14B-720P-Diffusers",
], ],
) )
register_configs( register_configs(
model_name="wan-fun-1.3b-inp",
sampling_param_cls=Wan2_1_Fun_1_3B_InP_SamplingParams, sampling_param_cls=Wan2_1_Fun_1_3B_InP_SamplingParams,
pipeline_config_cls=WanI2V480PConfig, pipeline_config_cls=WanI2V480PConfig,
model_paths=[ hf_model_paths=[
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers", "weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers",
], ],
) )
register_configs( register_configs(
model_name="wan-ti2v-5b",
sampling_param_cls=Wan2_2_TI2V_5B_SamplingParam, sampling_param_cls=Wan2_2_TI2V_5B_SamplingParam,
pipeline_config_cls=Wan2_2_TI2V_5B_Config, pipeline_config_cls=Wan2_2_TI2V_5B_Config,
model_paths=[ hf_model_paths=[
"Wan-AI/Wan2.2-TI2V-5B-Diffusers", "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
], ],
) )
register_configs( register_configs(
model_name="fastwan-ti2v-5b",
sampling_param_cls=Wan2_2_TI2V_5B_SamplingParam, sampling_param_cls=Wan2_2_TI2V_5B_SamplingParam,
pipeline_config_cls=FastWan2_2_TI2V_5B_Config, pipeline_config_cls=FastWan2_2_TI2V_5B_Config,
model_paths=[ hf_model_paths=[
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers", "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
"FastVideo/FastWan2.2-TI2V-5B-Diffusers", "FastVideo/FastWan2.2-TI2V-5B-Diffusers",
], ],
) )
register_configs( register_configs(
model_name="wan-t2v-a14b",
sampling_param_cls=Wan2_2_T2V_A14B_SamplingParam, sampling_param_cls=Wan2_2_T2V_A14B_SamplingParam,
pipeline_config_cls=Wan2_2_T2V_A14B_Config, pipeline_config_cls=Wan2_2_T2V_A14B_Config,
model_paths=[ hf_model_paths=["Wan-AI/Wan2.2-T2V-A14B-Diffusers"],
"Wan-AI/Wan2.2-T2V-A14B-Diffusers",
],
) )
register_configs( register_configs(
model_name="wan-i2v-a14b",
sampling_param_cls=Wan2_2_I2V_A14B_SamplingParam, sampling_param_cls=Wan2_2_I2V_A14B_SamplingParam,
pipeline_config_cls=Wan2_2_I2V_A14B_Config, pipeline_config_cls=Wan2_2_I2V_A14B_Config,
model_paths=[ hf_model_paths=["Wan-AI/Wan2.2-I2V-A14B-Diffusers"],
"Wan-AI/Wan2.2-I2V-A14B-Diffusers",
],
) )
register_configs( register_configs(
model_name="fast-wan-t2v-1.3b",
sampling_param_cls=FastWanT2V480PConfig, sampling_param_cls=FastWanT2V480PConfig,
pipeline_config_cls=FastWan2_1_T2V_480P_Config, pipeline_config_cls=FastWan2_1_T2V_480P_Config,
model_paths=[ hf_model_paths=[
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers", "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
], ],
) )
# FLUX # FLUX
register_configs( register_configs(
model_name="flux",
sampling_param_cls=FluxSamplingParams, sampling_param_cls=FluxSamplingParams,
pipeline_config_cls=FluxPipelineConfig, pipeline_config_cls=FluxPipelineConfig,
model_paths=[ hf_model_paths=[
"black-forest-labs/FLUX.1-dev", "black-forest-labs/FLUX.1-dev",
], ],
model_detectors=[lambda id: "flux" in id.lower()], model_detectors=[lambda hf_id: "flux.1" in hf_id.lower()],
) )
register_configs( register_configs(
model_name="flux-2",
sampling_param_cls=FluxSamplingParams, sampling_param_cls=FluxSamplingParams,
pipeline_config_cls=Flux2PipelineConfig, pipeline_config_cls=Flux2PipelineConfig,
model_paths=[ hf_model_paths=[
"black-forest-labs/FLUX.2-dev", "black-forest-labs/FLUX.2-dev",
], ],
model_detectors=[lambda id: "flux.2" in id.lower()], model_detectors=[lambda hf_id: "flux.2" in hf_id.lower()],
) )
register_configs( register_configs(
model_name="Z-image",
sampling_param_cls=ZImageSamplingParams, sampling_param_cls=ZImageSamplingParams,
pipeline_config_cls=ZImagePipelineConfig, pipeline_config_cls=ZImagePipelineConfig,
model_paths=[ hf_model_paths=[
"Tongyi-MAI/Z-Image-Turbo", "Tongyi-MAI/Z-Image-Turbo",
], ],
model_detectors=[lambda id: "z-image" in id.lower()], model_detectors=[lambda hf_id: "z-image" in hf_id.lower()],
) )
# Qwen-Image # Qwen-Image
register_configs( register_configs(
model_name="qwen-image",
sampling_param_cls=QwenImageSamplingParams, sampling_param_cls=QwenImageSamplingParams,
pipeline_config_cls=QwenImagePipelineConfig, pipeline_config_cls=QwenImagePipelineConfig,
hf_model_paths=["Qwen/Qwen-Image"],
) )
register_configs( register_configs(
model_name="qwen-image-edit",
sampling_param_cls=QwenImageSamplingParams, sampling_param_cls=QwenImageSamplingParams,
pipeline_config_cls=QwenImageEditPipelineConfig, pipeline_config_cls=QwenImageEditPipelineConfig,
hf_model_paths=["Qwen/Qwen-Image-Edit"],
) )
@@ -43,8 +43,6 @@ def build_pipeline(
""" """
model_path = server_args.model_path model_path = server_args.model_path
model_info = get_model_info(model_path) model_info = get_model_info(model_path)
if model_info is None:
raise ValueError(f"Unsupported model: {model_path}")
pipeline_cls = model_info.pipeline_cls pipeline_cls = model_info.pipeline_cls