Add registration for external model configurations (#39452)
This commit is contained in:
@@ -20,7 +20,7 @@ import os
|
|||||||
from enum import Enum, IntEnum, auto
|
from enum import Enum, IntEnum, auto
|
||||||
from functools import cached_property
|
from functools import cached_property
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, List, Optional, Set, Union
|
from typing import Any, Callable, List, Optional, Set, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from transformers import PretrainedConfig
|
from transformers import PretrainedConfig
|
||||||
@@ -46,6 +46,29 @@ from sglang.utils import is_in_ci
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_MODEL_CONFIG_FACTORIES: dict[type[ServerArgs], Callable[..., "ModelConfig"]] = {}
|
||||||
|
|
||||||
|
|
||||||
|
def register_model_config_factory(
|
||||||
|
server_args_type: type[ServerArgs], factory: Callable[..., "ModelConfig"]
|
||||||
|
) -> None:
|
||||||
|
"""Register a factory with the same signature as ModelConfig.from_server_args.
|
||||||
|
|
||||||
|
Register before resolution or model construction in each process. The nearest
|
||||||
|
registered type in the argument record's MRO wins, so registrations cover
|
||||||
|
subclasses too. Repeating the same registration is harmless; replacing a
|
||||||
|
different factory for the same type is an error.
|
||||||
|
"""
|
||||||
|
if not issubclass(server_args_type, ServerArgs):
|
||||||
|
raise TypeError("model-config factories require a ServerArgs subclass")
|
||||||
|
previous = _MODEL_CONFIG_FACTORIES.get(server_args_type)
|
||||||
|
if previous is not None and previous is not factory:
|
||||||
|
raise ValueError(
|
||||||
|
f"A model-config factory is already registered for {server_args_type.__qualname__}"
|
||||||
|
)
|
||||||
|
_MODEL_CONFIG_FACTORIES[server_args_type] = factory
|
||||||
|
|
||||||
|
|
||||||
MIMO_V2_MODEL_ARCHS = (
|
MIMO_V2_MODEL_ARCHS = (
|
||||||
"MiMoV2ForCausalLM",
|
"MiMoV2ForCausalLM",
|
||||||
"MiMoV2FlashForCausalLM",
|
"MiMoV2FlashForCausalLM",
|
||||||
@@ -692,27 +715,7 @@ class ModelConfig:
|
|||||||
self.hf_config.ngram_embedding_n if self.use_ngram_embedding else 0
|
self.hf_config.ngram_embedding_n if self.use_ngram_embedding else 0
|
||||||
)
|
)
|
||||||
self.use_engram = bool(getattr(self.hf_config, "engram_layer_ids", ()))
|
self.use_engram = bool(getattr(self.hf_config, "engram_layer_ids", ()))
|
||||||
# A multimodal arch is piecewise-incompatible until its LM prefill is validated.
|
self._derive_multimodal_cuda_graph_support(enable_multimodal)
|
||||||
self.is_piecewise_cuda_graph_disabled_model = (
|
|
||||||
is_piecewise_cuda_graph_disabled_model(self.hf_config.architectures)
|
|
||||||
or (
|
|
||||||
self.is_multimodal
|
|
||||||
and not is_multimodal_piecewise_cuda_graph_supported(
|
|
||||||
self.hf_config.architectures
|
|
||||||
)
|
|
||||||
)
|
|
||||||
)
|
|
||||||
# Multimodal archs whose language-model prefill is verified safe to capture
|
|
||||||
# under piecewise CUDA graph. ServerArgs otherwise disables prefill piecewise
|
|
||||||
# CG for every multimodal model; this opt-in re-enables it for listed archs
|
|
||||||
# (the vision encoder still runs eagerly via general_mm_embed_routine, only the
|
|
||||||
# LM forward is captured).
|
|
||||||
self.is_multimodal_piecewise_cuda_graph_supported = enable_multimodal and (
|
|
||||||
is_multimodal_piecewise_cuda_graph_supported(self.hf_config.architectures)
|
|
||||||
)
|
|
||||||
self.is_multimodal_breakable_cuda_graph_supported = enable_multimodal and (
|
|
||||||
is_multimodal_breakable_cuda_graph_supported(self.hf_config.architectures)
|
|
||||||
)
|
|
||||||
self.dtype = _get_and_verify_dtype(self.hf_text_config, dtype)
|
self.dtype = _get_and_verify_dtype(self.hf_text_config, dtype)
|
||||||
|
|
||||||
# Derive context length and model shapes
|
# Derive context length and model shapes
|
||||||
@@ -763,6 +766,18 @@ class ModelConfig:
|
|||||||
context_length: Optional[int] = None,
|
context_length: Optional[int] = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
|
for record_type in type(server_args).__mro__:
|
||||||
|
factory = _MODEL_CONFIG_FACTORIES.get(record_type)
|
||||||
|
if factory is not None:
|
||||||
|
return factory(
|
||||||
|
server_args,
|
||||||
|
model_path=model_path,
|
||||||
|
model_revision=model_revision,
|
||||||
|
is_draft_model=is_draft_model,
|
||||||
|
context_length=context_length,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
cfg = resolving_view(server_args)
|
cfg = resolving_view(server_args)
|
||||||
quantization = (
|
quantization = (
|
||||||
cfg.speculative_draft_model_quantization
|
cfg.speculative_draft_model_quantization
|
||||||
@@ -804,6 +819,31 @@ class ModelConfig:
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _derive_multimodal_cuda_graph_support(self, enable_multimodal: bool) -> None:
|
||||||
|
"""Declare graph capabilities before deriving shapes and validating config.
|
||||||
|
|
||||||
|
External ModelConfig subclasses can extend this without modifying global
|
||||||
|
architecture tables or repairing the config after construction.
|
||||||
|
"""
|
||||||
|
# A multimodal arch is piecewise-incompatible until its LM prefill is validated.
|
||||||
|
self.is_piecewise_cuda_graph_disabled_model = (
|
||||||
|
is_piecewise_cuda_graph_disabled_model(self.hf_config.architectures)
|
||||||
|
or (
|
||||||
|
self.is_multimodal
|
||||||
|
and not is_multimodal_piecewise_cuda_graph_supported(
|
||||||
|
self.hf_config.architectures
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
# The vision encoder still runs eagerly; this opt-in captures only the
|
||||||
|
# language-model prefill of architectures validated for piecewise graphs.
|
||||||
|
self.is_multimodal_piecewise_cuda_graph_supported = enable_multimodal and (
|
||||||
|
is_multimodal_piecewise_cuda_graph_supported(self.hf_config.architectures)
|
||||||
|
)
|
||||||
|
self.is_multimodal_breakable_cuda_graph_supported = enable_multimodal and (
|
||||||
|
is_multimodal_breakable_cuda_graph_supported(self.hf_config.architectures)
|
||||||
|
)
|
||||||
|
|
||||||
def _config_draft_model(self):
|
def _config_draft_model(self):
|
||||||
is_draft_model = self.is_draft_model
|
is_draft_model = self.is_draft_model
|
||||||
|
|
||||||
|
|||||||
@@ -1,16 +1,23 @@
|
|||||||
"""Unit tests for hybrid attention model configuration."""
|
"""Unit tests for model configuration."""
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
|
from tempfile import TemporaryDirectory
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
from unittest import mock
|
||||||
|
|
||||||
|
from transformers import LlamaConfig
|
||||||
|
|
||||||
|
from sglang.srt.arg_groups.overrides import model_config_of
|
||||||
from sglang.srt.configs.model_config import (
|
from sglang.srt.configs.model_config import (
|
||||||
ModelConfig,
|
ModelConfig,
|
||||||
get_hybrid_layer_ids,
|
get_hybrid_layer_ids,
|
||||||
is_embedding_gemma,
|
is_embedding_gemma,
|
||||||
is_multimodal_model,
|
is_multimodal_model,
|
||||||
|
register_model_config_factory,
|
||||||
resolve_spec_hidden_size,
|
resolve_spec_hidden_size,
|
||||||
)
|
)
|
||||||
from sglang.srt.configs.qwen4_exp import Qwen4ExpTextConfig
|
from sglang.srt.configs.qwen4_exp import Qwen4ExpTextConfig
|
||||||
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
@@ -112,5 +119,99 @@ class TestDraftModelConfig(CustomTestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestExternalModelConfig(CustomTestCase):
|
||||||
|
def setUp(self):
|
||||||
|
registry = mock.patch.dict(
|
||||||
|
"sglang.srt.configs.model_config._MODEL_CONFIG_FACTORIES", clear=True
|
||||||
|
)
|
||||||
|
registry.start()
|
||||||
|
self.addCleanup(registry.stop)
|
||||||
|
|
||||||
|
def test_factory_preserves_arguments_and_uses_the_shared_cache(self):
|
||||||
|
class ExternalArgs(ServerArgs):
|
||||||
|
pass
|
||||||
|
|
||||||
|
class DerivedArgs(ExternalArgs):
|
||||||
|
pass
|
||||||
|
|
||||||
|
result = SimpleNamespace(is_hybrid_swa=False)
|
||||||
|
factory = mock.Mock(return_value=result)
|
||||||
|
register_model_config_factory(ExternalArgs, factory)
|
||||||
|
args = DerivedArgs(model_path="dummy")
|
||||||
|
overrides = dict(
|
||||||
|
model_path="draft",
|
||||||
|
model_revision="draft-revision",
|
||||||
|
is_draft_model=True,
|
||||||
|
context_length=512,
|
||||||
|
dtype="bfloat16",
|
||||||
|
)
|
||||||
|
self.assertIs(ModelConfig.from_server_args(args, **overrides), result)
|
||||||
|
factory.assert_called_once_with(args, **overrides)
|
||||||
|
|
||||||
|
factory.reset_mock()
|
||||||
|
self.assertIs(model_config_of(args), result)
|
||||||
|
self.assertIs(model_config_of(args), result)
|
||||||
|
factory.assert_called_once_with(
|
||||||
|
args,
|
||||||
|
model_path=None,
|
||||||
|
model_revision=None,
|
||||||
|
is_draft_model=False,
|
||||||
|
context_length=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_registration_is_idempotent_and_the_nearest_type_wins(self):
|
||||||
|
class ExternalArgs(ServerArgs):
|
||||||
|
pass
|
||||||
|
|
||||||
|
class DerivedArgs(ExternalArgs):
|
||||||
|
pass
|
||||||
|
|
||||||
|
parent_factory = mock.Mock()
|
||||||
|
child_factory = mock.Mock()
|
||||||
|
register_model_config_factory(ExternalArgs, parent_factory)
|
||||||
|
register_model_config_factory(ExternalArgs, parent_factory)
|
||||||
|
with self.assertRaisesRegex(ValueError, "already registered"):
|
||||||
|
register_model_config_factory(ExternalArgs, child_factory)
|
||||||
|
with self.assertRaisesRegex(TypeError, "ServerArgs subclass"):
|
||||||
|
register_model_config_factory(SimpleNamespace, child_factory)
|
||||||
|
register_model_config_factory(DerivedArgs, child_factory)
|
||||||
|
args = DerivedArgs(model_path="dummy")
|
||||||
|
self.assertIs(ModelConfig.from_server_args(args), child_factory.return_value)
|
||||||
|
parent_factory.assert_not_called()
|
||||||
|
|
||||||
|
def test_capabilities_are_available_during_construction(self):
|
||||||
|
class ExternalArgs(ServerArgs):
|
||||||
|
pass
|
||||||
|
|
||||||
|
class ExternalConfig(ModelConfig):
|
||||||
|
def _derive_multimodal_cuda_graph_support(self, enable_multimodal):
|
||||||
|
super()._derive_multimodal_cuda_graph_support(enable_multimodal)
|
||||||
|
self.is_multimodal_breakable_cuda_graph_supported = True
|
||||||
|
|
||||||
|
def _derive_model_shapes(self):
|
||||||
|
assert self.is_multimodal_breakable_cuda_graph_supported
|
||||||
|
super()._derive_model_shapes()
|
||||||
|
|
||||||
|
register_model_config_factory(
|
||||||
|
ExternalArgs,
|
||||||
|
lambda args, **kwargs: ExternalConfig(model_path=args.model_path),
|
||||||
|
)
|
||||||
|
with TemporaryDirectory() as checkpoint:
|
||||||
|
LlamaConfig(
|
||||||
|
architectures=["LlamaForCausalLM"],
|
||||||
|
hidden_size=16,
|
||||||
|
intermediate_size=32,
|
||||||
|
num_attention_heads=2,
|
||||||
|
num_hidden_layers=2,
|
||||||
|
vocab_size=128,
|
||||||
|
).save_pretrained(checkpoint)
|
||||||
|
ordinary = ModelConfig.from_server_args(ServerArgs(model_path=checkpoint))
|
||||||
|
external = ModelConfig.from_server_args(ExternalArgs(model_path=checkpoint))
|
||||||
|
self.assertIs(type(ordinary), ModelConfig)
|
||||||
|
self.assertFalse(ordinary.is_multimodal_breakable_cuda_graph_supported)
|
||||||
|
self.assertIs(type(external), ExternalConfig)
|
||||||
|
self.assertTrue(external.is_multimodal_breakable_cuda_graph_supported)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user