refactor(attn): init hisparse_coordinator before attn_backend; replace lazy property with init-time capture (#26012)
Co-authored-by: Claude Sonnet 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Sonnet 4.6
parent
c5251a98a9
commit
d765dfd043
@@ -354,14 +354,9 @@ class DeepseekV4AttnBackend(
|
|||||||
self.page_size = model_runner.page_size
|
self.page_size = model_runner.page_size
|
||||||
assert self.page_size == 256, "the system hardcodes page_size=256"
|
assert self.page_size == 256, "the system hardcodes page_size=256"
|
||||||
|
|
||||||
# Pool refs — captured at construction so they survive deletion of the
|
|
||||||
# corresponding ForwardBatch fields.
|
|
||||||
self.req_to_token_pool = model_runner.req_to_token_pool
|
self.req_to_token_pool = model_runner.req_to_token_pool
|
||||||
self.token_to_kv_pool: DeepSeekV4TokenToKVPool = model_runner.token_to_kv_pool
|
self.token_to_kv_pool: DeepSeekV4TokenToKVPool = model_runner.token_to_kv_pool
|
||||||
# Keep a runner ref to read live state set after backend construction
|
self.hisparse_coordinator = model_runner.hisparse_coordinator
|
||||||
# (e.g. hisparse_coordinator is built in model_runner *after*
|
|
||||||
# init_attention_backend()).
|
|
||||||
self.model_runner = model_runner
|
|
||||||
self.req_to_token = model_runner.req_to_token_pool.req_to_token
|
self.req_to_token = model_runner.req_to_token_pool.req_to_token
|
||||||
self.MAX_SEQ_LEN_FOR_CAPTURE = self.req_to_token.shape[1]
|
self.MAX_SEQ_LEN_FOR_CAPTURE = self.req_to_token.shape[1]
|
||||||
|
|
||||||
@@ -385,12 +380,6 @@ class DeepseekV4AttnBackend(
|
|||||||
] = None
|
] = None
|
||||||
self._replay_forward_batch: Optional[ForwardBatch] = None # FIXME: out-of-band
|
self._replay_forward_batch: Optional[ForwardBatch] = None # FIXME: out-of-band
|
||||||
|
|
||||||
@property
|
|
||||||
def hisparse_coordinator(self):
|
|
||||||
# Live read: model_runner builds the coordinator *after*
|
|
||||||
# init_attention_backend(), so we cannot capture at __init__ time.
|
|
||||||
return self.model_runner.hisparse_coordinator
|
|
||||||
|
|
||||||
def _move_to_device(self, x: List[int]) -> torch.Tensor:
|
def _move_to_device(self, x: List[int]) -> torch.Tensor:
|
||||||
pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True)
|
pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True)
|
||||||
return pin_tensor.to(self.device, non_blocking=True)
|
return pin_tensor.to(self.device, non_blocking=True)
|
||||||
@@ -1196,7 +1185,6 @@ class DeepseekV4MultiStepBackend(DeepseekV4AttnBackend):
|
|||||||
self, model_runner: ModelRunner, topk: int, speculative_num_steps: int
|
self, model_runner: ModelRunner, topk: int, speculative_num_steps: int
|
||||||
):
|
):
|
||||||
super().__init__(model_runner)
|
super().__init__(model_runner)
|
||||||
self.model_runner = model_runner
|
|
||||||
self.topk = topk
|
self.topk = topk
|
||||||
self.speculative_num_steps = speculative_num_steps
|
self.speculative_num_steps = speculative_num_steps
|
||||||
self.attn_backends: List[DeepseekV4AttnBackend] = []
|
self.attn_backends: List[DeepseekV4AttnBackend] = []
|
||||||
|
|||||||
@@ -348,14 +348,9 @@ class DeepseekV4HipRadixBackend(
|
|||||||
self.page_size = model_runner.page_size
|
self.page_size = model_runner.page_size
|
||||||
assert self.page_size == 256, "the system hardcodes page_size=256"
|
assert self.page_size == 256, "the system hardcodes page_size=256"
|
||||||
|
|
||||||
# Pool refs — captured at construction so they survive deletion of the
|
|
||||||
# corresponding ForwardBatch fields.
|
|
||||||
self.req_to_token_pool = model_runner.req_to_token_pool
|
self.req_to_token_pool = model_runner.req_to_token_pool
|
||||||
self.token_to_kv_pool: DeepSeekV4TokenToKVPool = model_runner.token_to_kv_pool
|
self.token_to_kv_pool: DeepSeekV4TokenToKVPool = model_runner.token_to_kv_pool
|
||||||
# Keep a runner ref to read live state set after backend construction
|
self.hisparse_coordinator = model_runner.hisparse_coordinator
|
||||||
# (e.g. hisparse_coordinator is built in model_runner *after*
|
|
||||||
# init_attention_backend()).
|
|
||||||
self.model_runner = model_runner
|
|
||||||
self.req_to_token = model_runner.req_to_token_pool.req_to_token
|
self.req_to_token = model_runner.req_to_token_pool.req_to_token
|
||||||
self.MAX_SEQ_LEN_FOR_CAPTURE = self.req_to_token.shape[1]
|
self.MAX_SEQ_LEN_FOR_CAPTURE = self.req_to_token.shape[1]
|
||||||
|
|
||||||
@@ -379,12 +374,6 @@ class DeepseekV4HipRadixBackend(
|
|||||||
] = None
|
] = None
|
||||||
self._replay_forward_batch: Optional[ForwardBatch] = None # FIXME: out-of-band
|
self._replay_forward_batch: Optional[ForwardBatch] = None # FIXME: out-of-band
|
||||||
|
|
||||||
@property
|
|
||||||
def hisparse_coordinator(self):
|
|
||||||
# Live read: model_runner builds the coordinator *after*
|
|
||||||
# init_attention_backend(), so we cannot capture at __init__ time.
|
|
||||||
return self.model_runner.hisparse_coordinator
|
|
||||||
|
|
||||||
def _move_to_device(self, x: List[int]) -> torch.Tensor:
|
def _move_to_device(self, x: List[int]) -> torch.Tensor:
|
||||||
pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True)
|
pin_tensor = torch.tensor(x, dtype=torch.int32, pin_memory=True)
|
||||||
return pin_tensor.to(self.device, non_blocking=True)
|
return pin_tensor.to(self.device, non_blocking=True)
|
||||||
@@ -1196,7 +1185,6 @@ class DeepseekV4MultiStepBackend(DeepseekV4HipRadixBackend):
|
|||||||
self, model_runner: ModelRunner, topk: int, speculative_num_steps: int
|
self, model_runner: ModelRunner, topk: int, speculative_num_steps: int
|
||||||
):
|
):
|
||||||
super().__init__(model_runner)
|
super().__init__(model_runner)
|
||||||
self.model_runner = model_runner
|
|
||||||
self.topk = topk
|
self.topk = topk
|
||||||
self.speculative_num_steps = speculative_num_steps
|
self.speculative_num_steps = speculative_num_steps
|
||||||
self.attn_backends: List[DeepseekV4HipRadixBackend] = []
|
self.attn_backends: List[DeepseekV4HipRadixBackend] = []
|
||||||
|
|||||||
@@ -330,14 +330,9 @@ class DeepseekSparseAttnBackend(
|
|||||||
self.qk_rope_head_dim = model_runner.model_config.qk_rope_head_dim
|
self.qk_rope_head_dim = model_runner.model_config.qk_rope_head_dim
|
||||||
|
|
||||||
assert model_runner.req_to_token_pool is not None
|
assert model_runner.req_to_token_pool is not None
|
||||||
# Pool refs — captured at construction so they survive deletion of the
|
|
||||||
# corresponding ForwardBatch fields.
|
|
||||||
self.req_to_token_pool = model_runner.req_to_token_pool
|
self.req_to_token_pool = model_runner.req_to_token_pool
|
||||||
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||||
# Keep a runner ref to read live state set after backend construction
|
self.hisparse_coordinator = model_runner.hisparse_coordinator
|
||||||
# (e.g. hisparse_coordinator is built in model_runner *after*
|
|
||||||
# init_attention_backend()).
|
|
||||||
self.model_runner = model_runner
|
|
||||||
self.req_to_token = model_runner.req_to_token_pool.req_to_token
|
self.req_to_token = model_runner.req_to_token_pool.req_to_token
|
||||||
|
|
||||||
self.use_mha: bool = False
|
self.use_mha: bool = False
|
||||||
@@ -400,12 +395,6 @@ class DeepseekSparseAttnBackend(
|
|||||||
else:
|
else:
|
||||||
self.workspace_buffer = None
|
self.workspace_buffer = None
|
||||||
|
|
||||||
@property
|
|
||||||
def hisparse_coordinator(self):
|
|
||||||
# Live read: model_runner builds the coordinator *after*
|
|
||||||
# init_attention_backend(), so we cannot capture at __init__ time.
|
|
||||||
return self.model_runner.hisparse_coordinator
|
|
||||||
|
|
||||||
def get_device_int32_arange(self, l: int) -> torch.Tensor:
|
def get_device_int32_arange(self, l: int) -> torch.Tensor:
|
||||||
if l > len(self._arange_buf):
|
if l > len(self._arange_buf):
|
||||||
next_pow_of_2 = 1 << (l - 1).bit_length()
|
next_pow_of_2 = 1 << (l - 1).bit_length()
|
||||||
@@ -2318,7 +2307,6 @@ class DeepseekSparseAttnMultiStepBackend:
|
|||||||
def __init__(
|
def __init__(
|
||||||
self, model_runner: ModelRunner, topk: int, speculative_num_steps: int
|
self, model_runner: ModelRunner, topk: int, speculative_num_steps: int
|
||||||
):
|
):
|
||||||
self.model_runner = model_runner
|
|
||||||
self.topk = topk
|
self.topk = topk
|
||||||
self.speculative_num_steps = speculative_num_steps
|
self.speculative_num_steps = speculative_num_steps
|
||||||
self.attn_backends = []
|
self.attn_backends = []
|
||||||
|
|||||||
@@ -761,9 +761,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
|
|
||||||
if self.device == "cuda" or self.device == "musa":
|
if self.device == "cuda" or self.device == "musa":
|
||||||
self.init_cublas()
|
self.init_cublas()
|
||||||
self.init_attention_backend()
|
|
||||||
self.kernel_warmup()
|
|
||||||
# Init hisparse coordinator (must happen before CUDA graph capture)
|
|
||||||
if self.enable_hisparse:
|
if self.enable_hisparse:
|
||||||
from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator
|
from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator
|
||||||
from sglang.srt.mem_cache.sparsity import parse_hisparse_config
|
from sglang.srt.mem_cache.sparsity import parse_hisparse_config
|
||||||
@@ -785,6 +782,8 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
),
|
),
|
||||||
host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
|
host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
|
||||||
)
|
)
|
||||||
|
self.init_attention_backend()
|
||||||
|
self.kernel_warmup()
|
||||||
self._pre_initialize_flashinfer_allreduce_workspace()
|
self._pre_initialize_flashinfer_allreduce_workspace()
|
||||||
self.init_device_graphs()
|
self.init_device_graphs()
|
||||||
elif self.device == "cpu":
|
elif self.device == "cpu":
|
||||||
|
|||||||
@@ -91,6 +91,7 @@ class MockModelRunner:
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
enable_memory_saver=False,
|
enable_memory_saver=False,
|
||||||
)
|
)
|
||||||
|
self.hisparse_coordinator = None
|
||||||
|
|
||||||
|
|
||||||
@unittest.skipIf(not torch.cuda.is_available(), "Test requires CUDA")
|
@unittest.skipIf(not torch.cuda.is_available(), "Test requires CUDA")
|
||||||
|
|||||||
@@ -76,6 +76,7 @@ class MockModelRunner:
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
enable_memory_saver=False,
|
enable_memory_saver=False,
|
||||||
)
|
)
|
||||||
|
self.hisparse_coordinator = None
|
||||||
|
|
||||||
|
|
||||||
class MockReqToTokenPool:
|
class MockReqToTokenPool:
|
||||||
|
|||||||
@@ -268,6 +268,7 @@ class MockModelRunner:
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
enable_memory_saver=False,
|
enable_memory_saver=False,
|
||||||
)
|
)
|
||||||
|
self.hisparse_coordinator = None
|
||||||
|
|
||||||
|
|
||||||
def compare_outputs(trtllm_out, reference_out, tolerance=1e-2):
|
def compare_outputs(trtllm_out, reference_out, tolerance=1e-2):
|
||||||
|
|||||||
@@ -252,6 +252,7 @@ class MockModelRunner:
|
|||||||
"dsa_decode_backend": "fa3",
|
"dsa_decode_backend": "fa3",
|
||||||
},
|
},
|
||||||
)()
|
)()
|
||||||
|
self.hisparse_coordinator = None
|
||||||
|
|
||||||
|
|
||||||
@unittest.skipIf(not torch.cuda.is_available(), "Test requires CUDA")
|
@unittest.skipIf(not torch.cuda.is_available(), "Test requires CUDA")
|
||||||
|
|||||||
Reference in New Issue
Block a user