From 6d7c0d8ecb3b8a6a930d8710c1be0378badc0afa Mon Sep 17 00:00:00 2001 From: Andrewxu313 Date: Tue, 9 Jun 2026 12:19:21 +0000 Subject: [PATCH] refactor: consolidate model dispatch into a single registry Replace the duplicated if/elif name-matching in get_initializer.py and get_parallel_strategy_manager.py with one MODEL_REGISTRY + resolve_model() in batchgen/model_dispatch.py. Behavior-preserving: identical substring/exact matching, same order, same classes; per-model imports stay lazy. Adds tests/test_model_dispatch.py (14 cases). Co-Authored-By: Claude Opus 4.8 (1M context) Signed-off-by: Andrewxu313 --- batchgen/get_initializer.py | 49 ++------- batchgen/get_parallel_strategy_manager.py | 50 ++------- batchgen/model_dispatch.py | 125 ++++++++++++++++++++++ tests/test_model_dispatch.py | 44 ++++++++ 4 files changed, 187 insertions(+), 81 deletions(-) create mode 100644 batchgen/model_dispatch.py create mode 100644 tests/test_model_dispatch.py diff --git a/batchgen/get_initializer.py b/batchgen/get_initializer.py index e304cf38a..5baf61324 100644 --- a/batchgen/get_initializer.py +++ b/batchgen/get_initializer.py @@ -1,43 +1,12 @@ -KIMI_K25_BACKEND_NAME_PATTERNS = ( - "moonshotai/kimi-k2.5", - "moonshotai/kimi-k2.6", - "kimi-k2.5", - "kimi_k2.5", - "kimi-k25", - "kimi_k25", - "kimi-k2.6", - "kimi_k2.6", - "kimi-k26", - "kimi_k26", -) +"""Resolve a model name to its Initializer class. +Thin wrapper over the single dispatch registry in `batchgen.model_dispatch` +(see batchgen_design/model_architecture_spec.md section 2.1 -- model->implementation +dispatch lives only in the registry layer; the runtime core must not branch on +model names). +""" +from batchgen.model_dispatch import resolve_model -def _is_kimi_k25_backend_model(model_name: str) -> bool: - model_lower = model_name.strip().lower() - return any(pattern in model_lower for pattern in KIMI_K25_BACKEND_NAME_PATTERNS) -def get_initializer(model_name:str): - model_lower = model_name.lower() - if "minimax" in model_lower or "minimax-m2.5" in model_lower: - from batchgen.models.minimax.minimax_m25.minimax_m25_initializer import MiniMaxM25Initializer - return MiniMaxM25Initializer - elif _is_kimi_k25_backend_model(model_name): - from batchgen.models.moonshotai.kimi_k25.kimi_initializer import KimiK25Initializer - return KimiK25Initializer - elif "deepseek-v4" in model_lower: - from batchgen.models.deepseek.deepseekv4_flash.deepseekv4_flash_initializer import DeepSeekV4FlashInitializer - return DeepSeekV4FlashInitializer - elif model_lower in [ - "deepseek-ai/deepseek-r1", - "deepseek-ai/deepseek-v3", - ]: - from batchgen.models.deepseek.deepseekv3.deepseekv3_initializer import DeepseekV3Initializer - return DeepseekV3Initializer - elif "gpt-oss-120b" in model_lower: - from batchgen.models.openai.gpt_oss_120b.gpt_oss_initializer import GptOssInitializer - return GptOssInitializer - elif "glm-5" in model_lower: - from batchgen.models.glm.glm5.glm5_initializer import GLM5Initializer - return GLM5Initializer - else: - raise ValueError(f"Unsupported model name: {model_name}") +def get_initializer(model_name: str): + return resolve_model(model_name).initializer_loader() diff --git a/batchgen/get_parallel_strategy_manager.py b/batchgen/get_parallel_strategy_manager.py index 3b3c61bc0..05caff88f 100644 --- a/batchgen/get_parallel_strategy_manager.py +++ b/batchgen/get_parallel_strategy_manager.py @@ -1,44 +1,12 @@ -KIMI_K25_BACKEND_NAME_PATTERNS = ( - "moonshotai/kimi-k2.5", - "moonshotai/kimi-k2.6", - "kimi-k2.5", - "kimi_k2.5", - "kimi-k25", - "kimi_k25", - "kimi-k2.6", - "kimi_k2.6", - "kimi-k26", - "kimi_k26", -) +"""Resolve a model name to its Parallel Strategy Manager (PSM) class. +Thin wrapper over the single dispatch registry in `batchgen.model_dispatch` +(see batchgen_design/model_architecture_spec.md section 2.1 -- model->implementation +dispatch lives only in the registry layer; the runtime core must not branch on +model names). +""" +from batchgen.model_dispatch import resolve_model -def _is_kimi_k25_backend_model(model_name: str) -> bool: - model_lower = model_name.strip().lower() - return any(pattern in model_lower for pattern in KIMI_K25_BACKEND_NAME_PATTERNS) - -def get_parallel_strategy_manager(model_name:str): - model_lower = model_name.lower() - if "minimax" in model_lower or "minimax-m2.5" in model_lower: - from batchgen.models.minimax.minimax_m25.Parallel_Strategy_Manager import MiniMaxM25ParallelStrategyManager - return MiniMaxM25ParallelStrategyManager - elif _is_kimi_k25_backend_model(model_name): - from batchgen.models.moonshotai.kimi_k25.Parallel_Strategy_Manager import KimiK25ParallelStrategyManager - return KimiK25ParallelStrategyManager - elif "deepseek-v4" in model_lower: - from batchgen.models.deepseek.deepseekv4_flash.Parallel_Strategy_Manager import DeepSeekV4FlashParallelStrategyManager - return DeepSeekV4FlashParallelStrategyManager - elif model_lower in [ - "deepseek-ai/deepseek-r1", - "deepseek-ai/deepseek-v3", - ]: - from batchgen.models.deepseek.deepseekv3.Parallel_Strategy_Manager import DeepseekV3ParallelStrategyManager - return DeepseekV3ParallelStrategyManager - elif "gpt-oss-120b" in model_lower: - from batchgen.models.openai.gpt_oss_120b.Parallel_Strategy_Manager import GptOssParallelStrategyManager - return GptOssParallelStrategyManager - elif "glm-5" in model_lower: - from batchgen.models.glm.glm5.Parallel_Strategy_Manager import GLM5ParallelStrategyManager - return GLM5ParallelStrategyManager - else: - raise ValueError(f"Unsupported model name: {model_name}") +def get_parallel_strategy_manager(model_name: str): + return resolve_model(model_name).psm_loader() diff --git a/batchgen/model_dispatch.py b/batchgen/model_dispatch.py new file mode 100644 index 000000000..91bffdd28 --- /dev/null +++ b/batchgen/model_dispatch.py @@ -0,0 +1,125 @@ +"""Single source of truth for model-name -> (Initializer, PSM) dispatch. + +The runtime core must never branch on a model name or import a model package; +all such mapping lives here. (Design: batchgen_design/model_architecture_spec.md +section 2.1 -- model->implementation dispatch lives only in the registry layer.) + +This is a behavior-preserving consolidation of the former duplicated if/elif +chains in `get_initializer.py` and `get_parallel_strategy_manager.py`: the same +matching semantics (substring / exact, same order) resolve every name to the +same classes as before. Migrating the key to the exact canonical HuggingFace id +is a follow-up that depends on standardizing launch identifiers. + +Class imports stay lazy (per-entry loader closures) so importing this module +does not pull in every model package. +""" +from __future__ import annotations + +from dataclasses import dataclass +from typing import Callable, Optional, Tuple + +from batchgen.config.model_name_utils import KIMI_K25_BACKEND_NAME_PATTERNS + + +@dataclass(frozen=True) +class ModelEntry: + """One registered model: how to match its name and how to load its classes.""" + + key: str + substrings: Tuple[str, ...] = () + exact_names: Tuple[str, ...] = () + initializer_loader: Optional[Callable[[], type]] = None + psm_loader: Optional[Callable[[], type]] = None + + def matches(self, name_lower: str) -> bool: + if name_lower in self.exact_names: + return True + return any(pattern in name_lower for pattern in self.substrings) + + +# --- lazy loaders: keep model-package imports out of module import time ------- +def _minimax_initializer(): + from batchgen.models.minimax.minimax_m25.minimax_m25_initializer import MiniMaxM25Initializer + return MiniMaxM25Initializer + + +def _minimax_psm(): + from batchgen.models.minimax.minimax_m25.Parallel_Strategy_Manager import MiniMaxM25ParallelStrategyManager + return MiniMaxM25ParallelStrategyManager + + +def _kimi_k25_initializer(): + from batchgen.models.moonshotai.kimi_k25.kimi_initializer import KimiK25Initializer + return KimiK25Initializer + + +def _kimi_k25_psm(): + from batchgen.models.moonshotai.kimi_k25.Parallel_Strategy_Manager import KimiK25ParallelStrategyManager + return KimiK25ParallelStrategyManager + + +def _deepseek_v4_flash_initializer(): + from batchgen.models.deepseek.deepseekv4_flash.deepseekv4_flash_initializer import DeepSeekV4FlashInitializer + return DeepSeekV4FlashInitializer + + +def _deepseek_v4_flash_psm(): + from batchgen.models.deepseek.deepseekv4_flash.Parallel_Strategy_Manager import DeepSeekV4FlashParallelStrategyManager + return DeepSeekV4FlashParallelStrategyManager + + +def _deepseek_v3_initializer(): + from batchgen.models.deepseek.deepseekv3.deepseekv3_initializer import DeepseekV3Initializer + return DeepseekV3Initializer + + +def _deepseek_v3_psm(): + from batchgen.models.deepseek.deepseekv3.Parallel_Strategy_Manager import DeepseekV3ParallelStrategyManager + return DeepseekV3ParallelStrategyManager + + +def _gpt_oss_initializer(): + from batchgen.models.openai.gpt_oss_120b.gpt_oss_initializer import GptOssInitializer + return GptOssInitializer + + +def _gpt_oss_psm(): + from batchgen.models.openai.gpt_oss_120b.Parallel_Strategy_Manager import GptOssParallelStrategyManager + return GptOssParallelStrategyManager + + +def _glm5_initializer(): + from batchgen.models.glm.glm5.glm5_initializer import GLM5Initializer + return GLM5Initializer + + +def _glm5_psm(): + from batchgen.models.glm.glm5.Parallel_Strategy_Manager import GLM5ParallelStrategyManager + return GLM5ParallelStrategyManager + + +# Order matters: first match wins. This reproduces the original if/elif order in +# get_initializer.py / get_parallel_strategy_manager.py exactly. +MODEL_REGISTRY: Tuple[ModelEntry, ...] = ( + ModelEntry("minimax_m25", substrings=("minimax",), + initializer_loader=_minimax_initializer, psm_loader=_minimax_psm), + ModelEntry("kimi_k25", substrings=KIMI_K25_BACKEND_NAME_PATTERNS, + initializer_loader=_kimi_k25_initializer, psm_loader=_kimi_k25_psm), + ModelEntry("deepseek_v4_flash", substrings=("deepseek-v4",), + initializer_loader=_deepseek_v4_flash_initializer, psm_loader=_deepseek_v4_flash_psm), + ModelEntry("deepseek_v3", exact_names=("deepseek-ai/deepseek-r1", "deepseek-ai/deepseek-v3"), + initializer_loader=_deepseek_v3_initializer, psm_loader=_deepseek_v3_psm), + ModelEntry("gpt_oss", substrings=("gpt-oss-120b",), + initializer_loader=_gpt_oss_initializer, psm_loader=_gpt_oss_psm), + ModelEntry("glm5", substrings=("glm-5",), + initializer_loader=_glm5_initializer, psm_loader=_glm5_psm), +) + + +def resolve_model(model_name: str) -> ModelEntry: + """Return the registry entry for `model_name`, or raise ValueError if unsupported.""" + name_lower = (model_name or "").strip().lower() + for entry in MODEL_REGISTRY: + if entry.matches(name_lower): + return entry + raise ValueError(f"Unsupported model name: {model_name}") diff --git a/tests/test_model_dispatch.py b/tests/test_model_dispatch.py new file mode 100644 index 000000000..158c89d22 --- /dev/null +++ b/tests/test_model_dispatch.py @@ -0,0 +1,44 @@ +"""M1 regression: the consolidated registry resolves every model name to the +same class the former get_initializer / get_parallel_strategy_manager if/elif +chains did. + +GPU-free: asserts on the entry key, so it does not import any model package. +""" +import pytest + +from batchgen.model_dispatch import resolve_model + + +# (input model name, expected entry key) -- one per branch of the old if/elif, +# incl. the canonical HF id the GLM-5 server launches with. +CASES = [ + ("zai-org/GLM-5.1-FP8", "glm5"), + ("GLM-5", "glm5"), + ("glm-5.1-fp8", "glm5"), + ("moonshotai/Kimi-K2.5", "kimi_k25"), + ("moonshotai/Kimi-K2.6", "kimi_k25"), + ("kimi_k25", "kimi_k25"), + ("deepseek-ai/DeepSeek-V4-Flash", "deepseek_v4_flash"), + ("deepseek-ai/DeepSeek-R1", "deepseek_v3"), + ("deepseek-ai/DeepSeek-V3", "deepseek_v3"), + ("openai/gpt-oss-120b", "gpt_oss"), + ("MiniMax-M2.5", "minimax_m25"), + ("minimax", "minimax_m25"), +] + + +@pytest.mark.parametrize("name,expected_key", CASES) +def test_resolve_model_key(name, expected_key): + assert resolve_model(name).key == expected_key + + +def test_unknown_model_raises(): + with pytest.raises(ValueError): + resolve_model("not-a-real-model") + + +def test_every_entry_has_loaders(): + for name, _ in CASES: + entry = resolve_model(name) + assert entry.initializer_loader is not None + assert entry.psm_loader is not None