Move deferred mamba cow and clear (#29945)

This commit is contained in:
Ke Bao
2026-07-03 11:29:05 +08:00
committed by GitHub
parent a2d7eb303e
commit 372a893744
2 changed files with 53 additions and 43 deletions
@@ -53,46 +53,6 @@ class MambaAttnBackendBase(AttentionBackend):
self.cached_cuda_graph_verify_query_start_loc: torch.Tensor = None
self.conv_states_shape: tuple[int, int] = None
def _execute_deferred_mamba_cow_and_clear(self, forward_batch: ForwardBatch):
"""Run deferred clear/COW ops on the forward stream to avoid races."""
if (
not forward_batch.forward_mode.is_extend()
or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend_v2()
or self.is_draft_worker
):
return
if (
forward_batch.mamba_clear_indices is not None
and len(forward_batch.mamba_clear_indices) > 0
):
# mamba_pool is a pure PHYSICAL store; translate before zeroing or
# clear_slots zeroes the wrong physical slots.
self.req_to_token_pool.mamba_pool.clear_slots(
self._translate_mamba_indices(forward_batch.mamba_clear_indices)
)
if (
forward_batch.mamba_cow_src_indices is not None
and len(forward_batch.mamba_cow_src_indices) > 0
):
ckpt_pool = getattr(self.req_to_token_pool, "mamba_ckpt_pool", None)
if ckpt_pool is not None:
# int8 checkpoints: dequantize src int8 ckpt slot into the active bf16 dst.
ckpt_pool.load_to_active(
self.req_to_token_pool.mamba_pool,
forward_batch.mamba_cow_src_indices,
forward_batch.mamba_cow_dst_indices,
)
else:
# mamba_pool is a pure PHYSICAL store; translate both COW slot ids.
self.req_to_token_pool.mamba_pool.copy_from(
self._translate_mamba_indices(forward_batch.mamba_cow_src_indices),
self._translate_mamba_indices(forward_batch.mamba_cow_dst_indices),
)
forward_batch.mamba_clear_indices = None
forward_batch.mamba_cow_src_indices = None
forward_batch.mamba_cow_dst_indices = None
def _translate_mamba_indices(self, mamba_indices: torch.Tensor) -> torch.Tensor:
"""Virtual->physical mamba slot-id translate (identity for the non-unified
pool). Must run everywhere mamba ids feed the SSM/conv kernels or mamba-pool
@@ -274,7 +234,6 @@ class MambaAttnBackendBase(AttentionBackend):
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
self._execute_deferred_mamba_cow_and_clear(forward_batch)
self.forward_metadata = self._forward_metadata(forward_batch)
def _init_track_conv_indices(
@@ -767,7 +726,6 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
self._execute_deferred_mamba_cow_and_clear(forward_batch)
metadata = self._forward_metadata(forward_batch)
self.forward_metadata = Mamba2Metadata.prepare_mixed(
metadata,
@@ -142,7 +142,7 @@ from sglang.srt.lora.lora_manager import LoRAManager
from sglang.srt.lora.lora_registry import LoRARef
from sglang.srt.managers.schedule_batch import sanity_check_mm_pad_shift_value
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
from sglang.srt.model_executor.cuda_graph_config import (
Backend,
@@ -3072,6 +3072,54 @@ class ModelRunner(ModelRunnerKVCacheMixin):
return output
def _maybe_execute_deferred_mamba_cow_and_clear(
self, forward_batch: ForwardBatch
) -> None:
"""Run deferred clear/COW on the forward stream, before the mamba layers
read the pool, so the copies don't race the scheduler copy stream.
No-op unless this is an extend forward on a mamba model's target worker;
COW/clear only happen at prefix match on extend.
"""
pool = self.req_to_token_pool
if (
not isinstance(pool, HybridReqToTokenPool)
or self.is_draft_worker
or not forward_batch.forward_mode.is_extend()
or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend_v2()
):
return
if (
forward_batch.mamba_clear_indices is not None
and len(forward_batch.mamba_clear_indices) > 0
):
# mamba_pool is a pure PHYSICAL store; translate before zeroing or
# clear_slots zeroes the wrong physical slots.
pool.mamba_pool.clear_slots(
pool.translate_mamba_indices(forward_batch.mamba_clear_indices)
)
if (
forward_batch.mamba_cow_src_indices is not None
and len(forward_batch.mamba_cow_src_indices) > 0
):
if pool.mamba_ckpt_pool is not None:
# int8 checkpoints: dequantize src int8 ckpt slot into the active bf16 dst.
pool.mamba_ckpt_pool.load_to_active(
pool.mamba_pool,
forward_batch.mamba_cow_src_indices,
forward_batch.mamba_cow_dst_indices,
)
else:
# mamba_pool is a pure PHYSICAL store; translate both COW slot ids.
pool.mamba_pool.copy_from(
pool.translate_mamba_indices(forward_batch.mamba_cow_src_indices),
pool.translate_mamba_indices(forward_batch.mamba_cow_dst_indices),
)
forward_batch.mamba_clear_indices = None
forward_batch.mamba_cow_src_indices = None
forward_batch.mamba_cow_dst_indices = None
def _forward_raw(
self,
forward_batch: ForwardBatch,
@@ -3119,6 +3167,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
# and the collectives depend on.
self._prepare_eager_forward_batch(forward_batch)
# Deferred mamba COW/clear on the forward stream, before the extend
# dispatch below reads the pool.
self._maybe_execute_deferred_mamba_cow_and_clear(forward_batch)
if forward_batch.forward_mode.is_split_prefill():
# Layer-split mode; stays on ModelRunner, not the eager runner.
ret = self.forward_split_prefill(