Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
49 changes: 9 additions & 40 deletions batchgen/get_initializer.py
Original file line number Diff line number Diff line change
@@ -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()
50 changes: 9 additions & 41 deletions batchgen/get_parallel_strategy_manager.py
Original file line number Diff line number Diff line change
@@ -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()
125 changes: 125 additions & 0 deletions batchgen/model_dispatch.py
Original file line number Diff line number Diff line change
@@ -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}")
44 changes: 44 additions & 0 deletions tests/test_model_dispatch.py
Original file line number Diff line number Diff line change
@@ -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