Introduce ModelRunner.ps ParallelState (#31161)
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user