Move deferred mamba cow and clear (#29945)
This commit is contained in:
@@ -53,46 +53,6 @@ class MambaAttnBackendBase(AttentionBackend):
|
|||||||
self.cached_cuda_graph_verify_query_start_loc: torch.Tensor = None
|
self.cached_cuda_graph_verify_query_start_loc: torch.Tensor = None
|
||||||
self.conv_states_shape: tuple[int, int] = 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:
|
def _translate_mamba_indices(self, mamba_indices: torch.Tensor) -> torch.Tensor:
|
||||||
"""Virtual->physical mamba slot-id translate (identity for the non-unified
|
"""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
|
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):
|
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)
|
self.forward_metadata = self._forward_metadata(forward_batch)
|
||||||
|
|
||||||
def _init_track_conv_indices(
|
def _init_track_conv_indices(
|
||||||
@@ -767,7 +726,6 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
||||||
self._execute_deferred_mamba_cow_and_clear(forward_batch)
|
|
||||||
metadata = self._forward_metadata(forward_batch)
|
metadata = self._forward_metadata(forward_batch)
|
||||||
self.forward_metadata = Mamba2Metadata.prepare_mixed(
|
self.forward_metadata = Mamba2Metadata.prepare_mixed(
|
||||||
metadata,
|
metadata,
|
||||||
|
|||||||
@@ -142,7 +142,7 @@ from sglang.srt.lora.lora_manager import LoRAManager
|
|||||||
from sglang.srt.lora.lora_registry import LoRARef
|
from sglang.srt.lora.lora_registry import LoRARef
|
||||||
from sglang.srt.managers.schedule_batch import sanity_check_mm_pad_shift_value
|
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.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.cpu_graph_runner import CPUGraphRunner
|
||||||
from sglang.srt.model_executor.cuda_graph_config import (
|
from sglang.srt.model_executor.cuda_graph_config import (
|
||||||
Backend,
|
Backend,
|
||||||
@@ -3072,6 +3072,54 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
|
|
||||||
return output
|
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(
|
def _forward_raw(
|
||||||
self,
|
self,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
@@ -3119,6 +3167,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
# and the collectives depend on.
|
# and the collectives depend on.
|
||||||
self._prepare_eager_forward_batch(forward_batch)
|
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():
|
if forward_batch.forward_mode.is_split_prefill():
|
||||||
# Layer-split mode; stays on ModelRunner, not the eager runner.
|
# Layer-split mode; stays on ModelRunner, not the eager runner.
|
||||||
ret = self.forward_split_prefill(
|
ret = self.forward_split_prefill(
|
||||||
|
|||||||
Reference in New Issue
Block a user