Introduce ModelRunner.ps ParallelState (#31161)

This commit is contained in:
fzyzcjy
2026-07-14 16:01:14 +08:00
committed by GitHub
parent 1dc48c2c3b
commit 725920915f
28 changed files with 202 additions and 276 deletions
@@ -7,6 +7,7 @@ import torch.nn.functional as F
from torch import nn
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, ReqToTokenPool
@@ -286,7 +287,9 @@ class TinyModelConfig:
num_key_value_heads=num_kv_heads,
head_dim=head_dim,
)
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
def get_num_attention_heads(self, tp_size: int) -> int:
assert self.num_attention_heads % tp_size == 0
@@ -322,6 +325,7 @@ class MockModelRunner(ModelRunner):
self.tp_size = 1
self.dp_size = 1
self.pp_size = 1
self.ps = ParallelState.trivial()
self.is_draft_worker = False
self.spec_algorithm = SpeculativeAlgorithm.NONE
# The runner lifecycle warms up kernels in capture() / first execute()
@@ -5,6 +5,7 @@ from typing import Any
import torch
from torch import nn
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool, ReqToTokenPool
@@ -258,7 +259,9 @@ class TinyDSAModelConfig:
index_topk=index_topk,
num_hidden_layers=1,
)
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
class DSAMockModelRunner(ModelRunner):
@@ -308,6 +311,7 @@ class DSAMockModelRunner(ModelRunner):
self._kernel_warmed_up = True
self.dp_size = 1
self.pp_size = 1
self.ps = ParallelState.trivial()
self._server_args_override = get_context().override_server_args(
attention_backend=case.backend,
chunked_prefill_size=-1,
@@ -19,6 +19,7 @@ from typing import Any
import torch
from torch import nn
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.dsv4.quant_k_cache import (
@@ -282,7 +283,9 @@ class TinyDSV4ModelConfig:
num_hidden_layers=len(compression_ratios),
compress_ratios=list(compression_ratios),
)
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
class MockDSV4ModelRunner:
@@ -332,6 +335,7 @@ class MockDSV4ModelRunner:
self.tp_size = 1
self.dp_size = 1
self.pp_size = 1
self.ps = ParallelState.trivial()
self._server_args_override = get_context().override_server_args(
attention_backend=case.backend,
chunked_prefill_size=-1,
@@ -4,6 +4,7 @@ from types import SimpleNamespace
import torch
from torch import nn
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention import (
dual_chunk_flashattention_backend as _dual_chunk_backend,
)
@@ -286,7 +287,9 @@ class TinyDualChunkModelConfig:
dual_chunk_attention_config or DUAL_CHUNK_CONFIG
),
)
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
def get_num_attention_heads(self, tp_size: int) -> int:
assert self.num_attention_heads % tp_size == 0
@@ -323,6 +326,7 @@ class DualChunkMockModelRunner(ModelRunner):
self._kernel_warmed_up = True
self.dp_size = 1
self.pp_size = 1
self.ps = ParallelState.trivial()
self._server_args_override = get_context().override_server_args(
attention_backend=case.backend,
chunked_prefill_size=-1,
@@ -10,6 +10,7 @@ from sglang.srt.configs.mamba_utils import (
Mamba2StateShape,
)
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
HybridLinearAttnBackend,
@@ -182,7 +183,9 @@ class TinyGDNModelConfig:
self.attention_chunk_size = None
self.sliding_window_size = None
self.hf_config = SimpleNamespace(architectures=["TinyGDNForCausalLM"])
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
def get_num_kv_heads(self, tp_size: int) -> int:
assert self.num_key_value_heads % tp_size == 0
@@ -210,6 +213,7 @@ class MockGDNModelRunner(ModelRunner):
self.dtype = dtype
self.kv_cache_dtype = dtype
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.canary_manager = None
self.page_size = case.page_size
self.model_config = model_config
@@ -10,6 +10,7 @@ from sglang.srt.configs.mamba_utils import (
Mamba2StateDType,
)
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
HybridLinearAttnBackend,
@@ -188,7 +189,9 @@ class TinyKDAModelConfig:
self.attention_chunk_size = None
self.sliding_window_size = None
self.hf_config = SimpleNamespace(architectures=["TinyKDAForCausalLM"])
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
def get_num_kv_heads(self, tp_size: int) -> int:
assert self.num_key_value_heads % tp_size == 0
@@ -216,6 +219,7 @@ class MockKDAModelRunner(ModelRunner):
self.dtype = dtype
self.kv_cache_dtype = dtype
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.canary_manager = None
self.page_size = case.page_size
self.model_config = model_config
@@ -10,6 +10,7 @@ from sglang.srt.configs.mamba_utils import (
Mamba2StateShape,
)
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.linear.lightning_backend import (
LightningAttentionBackend,
@@ -198,7 +199,9 @@ class TinyLightningModelConfig:
num_hidden_layers=num_hidden_layers,
linear_backend=linear_backend,
)
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
def get_num_kv_heads(self, tp_size: int) -> int:
assert self.num_key_value_heads % tp_size == 0
@@ -224,6 +227,7 @@ class MockLightningModelRunner(ModelRunner):
self.dtype = dtype
self.kv_cache_dtype = dtype
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.canary_manager = None
self.page_size = case.page_size
self.model_config = model_config
@@ -18,12 +18,14 @@ _parallel_override.__enter__()
# Provide a stub group with world_size=1 so use_symmetric_memory short-circuits.
_linear_mod.get_tp_group = lambda: SimpleNamespace(world_size=1)
from sglang.srt.configs.falcon_h1 import FalconH1Config # noqa: E402
from sglang.srt.configs.mamba_utils import ( # noqa: E402
Mamba2CacheParams,
Mamba2StateDType,
Mamba2StateShape,
)
from sglang.srt.configs.model_config import AttentionArch # noqa: E402
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ( # noqa: E402
ATTENTION_BACKENDS,
)
@@ -278,14 +280,14 @@ class TinyMamba2ModelConfig:
self.is_local_attention_model = False
self.attention_chunk_size = None
self.sliding_window_size = None
# Mamba2AttnBackend reads mamba2_config.mamba_chunk_size; expose it
# through a SimpleNamespace-as-hf_config so runner.mamba2_config returns
# something non-None with the expected attribute.
self.hf_config = SimpleNamespace(
# Mamba2AttnBackend reads mamba2_config(model_config).mamba_chunk_size; expose it
self.hf_config = FalconH1Config(
architectures=["TinyMamba2ForCausalLM"],
mamba_chunk_size=case.mamba_chunk_size,
)
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
def get_num_kv_heads(self, tp_size: int) -> int:
assert self.num_key_value_heads % tp_size == 0
@@ -310,6 +312,7 @@ class MockMamba2ModelRunner(ModelRunner):
self.dtype = dtype
self.kv_cache_dtype = dtype
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.canary_manager = None
self.page_size = case.page_size
self.model_config = model_config
@@ -7,6 +7,7 @@ import torch.nn.functional as F
from torch import nn
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool, ReqToTokenPool
@@ -191,7 +192,9 @@ class TinyMLAModelConfig:
qk_rope_head_dim=qk_rope_head_dim,
v_head_dim=kv_lora_rank,
)
self.hf_config.get_text_config = lambda: self.hf_config
self.hf_text_config = self.hf_config
self.linear_attn_registry_result = None
def get_num_attention_heads(self, tp_size: int) -> int:
assert self.num_attention_heads % tp_size == 0
@@ -233,6 +236,7 @@ class MockMLAModelRunner(ModelRunner):
self.tp_size = 1
self.dp_size = 1
self.pp_size = 1
self.ps = ParallelState.trivial()
speculative_num_draft_tokens = (
max(case.input_lens)
if case.forward_mode.is_target_verify()