Compare commits

...
Author SHA1 Message Date
alexzms af01289b30 [bugfix] registry: infer workload_type from model registration instead of defaulting to T2V
get_model_info() hard-defaulted workload_type to T2V when the caller
didn't pass one. Models registered for a single non-T2V workload (e.g.
Matrix-Game-2.0, registered I2V-only) then failed pipeline resolution:

  ValueError: Pipeline 'MatrixGame2I2VPipeline' is not supported for
  pipeline type 'basic' and workload type 't2v'

This deterministically breaks the Matrix-Game-2.0 SSIM full-suite test.

When workload_type is None, derive it from the model's registered
ConfigInfo.workload_types: use it if the model registers exactly one
workload, otherwise keep the T2V fallback. Behavior is unchanged for
T2V models and for multi/empty registrations; only single-non-T2V
models (the broken case) change. config_info is resolved once, earlier.
2026-05-17 19:34:52 +00:00
+9 -4
View File
@@ -781,8 +781,16 @@ def get_model_info(
elif isinstance(pipeline_type, str):
pipeline_type = PipelineType.from_string(pipeline_type)
config_info = _get_config_info(model_path, raise_on_missing=True)
assert config_info is not None, "config_info must be resolved"
if workload_type is None:
workload_type = WorkloadType.T2V
# Derive the workload from the model's registration instead of
# assuming T2V. A model registered for a single workload (e.g.
# Matrix-Game-2.0, which is I2V-only) otherwise fails pipeline
# resolution under the blanket T2V default.
registered = config_info.workload_types
workload_type = (registered[0] if len(registered) == 1 else WorkloadType.T2V)
if os.path.exists(model_path):
config = verify_model_config_and_directory(model_path)
@@ -801,9 +809,6 @@ def get_model_info(
pipeline_registry = get_pipeline_registry(pipeline_type)
pipeline_cls = pipeline_registry.resolve_pipeline_cls(pipeline_name, pipeline_type, workload_type)
config_info = _get_config_info(model_path, raise_on_missing=True)
assert config_info is not None, "config_info must be resolved"
sampling_param_cls = config_info.sampling_param_cls or SamplingParam
return ModelInfo(