Fix Qwen Next GDN w/ Radix Cache (#16053)
This commit is contained in:
@@ -127,6 +127,13 @@ def set_torch_compile_config():
|
||||
class PiecewiseCudaGraphRunner:
|
||||
"""A PiecewiseCudaGraphRunner runs the forward pass of a model with cuda graph and torch.compile."""
|
||||
|
||||
def is_mamba_track_enabled(self):
|
||||
return (
|
||||
self.model_runner.server_args.enable_mamba_extra_buffer()
|
||||
and not self.model_runner.server_args.disable_radix_cache
|
||||
and self.model_runner.spec_algorithm.is_none()
|
||||
)
|
||||
|
||||
def __init__(self, model_runner: ModelRunner):
|
||||
# Parse args
|
||||
self.model_runner = model_runner
|
||||
@@ -175,8 +182,10 @@ class PiecewiseCudaGraphRunner:
|
||||
self.capture_hidden_mode = CaptureHiddenMode.FULL
|
||||
|
||||
self.max_num_tokens = max(self.capture_num_tokens)
|
||||
self.max_bs = model_runner.req_to_token_pool.size
|
||||
|
||||
self.is_multimodal = model_runner.is_multimodal
|
||||
self.mamba_track_enabled = self.is_mamba_track_enabled()
|
||||
|
||||
# Graph inputs
|
||||
with torch.device(self.device):
|
||||
@@ -189,6 +198,21 @@ class PiecewiseCudaGraphRunner:
|
||||
if model_runner.is_hybrid_swa
|
||||
else None
|
||||
)
|
||||
self.mamba_track_indices = (
|
||||
torch.zeros((self.max_bs,), dtype=torch.int64)
|
||||
if self.mamba_track_enabled
|
||||
else None
|
||||
)
|
||||
self.mamba_track_mask = (
|
||||
torch.zeros((self.max_bs,), dtype=torch.bool)
|
||||
if self.mamba_track_enabled
|
||||
else None
|
||||
)
|
||||
self.mamba_track_seqlens = (
|
||||
torch.zeros((self.max_bs,), dtype=torch.int32)
|
||||
if self.mamba_track_enabled
|
||||
else None
|
||||
)
|
||||
self.positions = torch.zeros((self.max_num_tokens,), dtype=torch.int64)
|
||||
|
||||
self.tbo_plugin = TboCudaGraphRunnerPlugin()
|
||||
@@ -269,6 +293,19 @@ class PiecewiseCudaGraphRunner:
|
||||
if self.out_cache_loc_swa is not None
|
||||
else None
|
||||
)
|
||||
mamba_track_indices = (
|
||||
self.mamba_track_indices[:1]
|
||||
if self.mamba_track_indices is not None
|
||||
else None
|
||||
)
|
||||
mamba_track_mask = (
|
||||
self.mamba_track_mask[:1] if self.mamba_track_mask is not None else None
|
||||
)
|
||||
mamba_track_seqlens = (
|
||||
self.mamba_track_seqlens[:1]
|
||||
if self.mamba_track_seqlens is not None
|
||||
else None
|
||||
)
|
||||
with torch.device(self.device):
|
||||
forward_batch = ForwardBatch(
|
||||
forward_mode=ForwardMode.EXTEND,
|
||||
@@ -286,6 +323,9 @@ class PiecewiseCudaGraphRunner:
|
||||
out_cache_loc=out_cache_loc,
|
||||
out_cache_loc_swa=out_cache_loc_swa,
|
||||
seq_lens_sum=num_tokens,
|
||||
mamba_track_indices=mamba_track_indices,
|
||||
mamba_track_mask=mamba_track_mask,
|
||||
mamba_track_seqlens=mamba_track_seqlens,
|
||||
encoder_lens=None,
|
||||
return_logprob=False,
|
||||
extend_num_tokens=num_tokens,
|
||||
@@ -386,6 +426,19 @@ class PiecewiseCudaGraphRunner:
|
||||
if self.out_cache_loc_swa is not None
|
||||
else None
|
||||
)
|
||||
mamba_track_indices = (
|
||||
self.mamba_track_indices[:bs]
|
||||
if self.mamba_track_indices is not None
|
||||
else None
|
||||
)
|
||||
mamba_track_mask = (
|
||||
self.mamba_track_mask[:bs] if self.mamba_track_mask is not None else None
|
||||
)
|
||||
mamba_track_seqlens = (
|
||||
self.mamba_track_seqlens[:bs]
|
||||
if self.mamba_track_seqlens is not None
|
||||
else None
|
||||
)
|
||||
positions = self.positions[:num_tokens]
|
||||
mrope_positions = (
|
||||
self.mrope_positions[:, :num_tokens] if self.is_multimodal else None
|
||||
@@ -417,6 +470,9 @@ class PiecewiseCudaGraphRunner:
|
||||
out_cache_loc=out_cache_loc,
|
||||
out_cache_loc_swa=out_cache_loc_swa,
|
||||
seq_lens_sum=num_tokens,
|
||||
mamba_track_indices=mamba_track_indices,
|
||||
mamba_track_mask=mamba_track_mask,
|
||||
mamba_track_seqlens=mamba_track_seqlens,
|
||||
encoder_lens=None,
|
||||
return_logprob=False,
|
||||
extend_num_tokens=num_tokens,
|
||||
@@ -504,6 +560,23 @@ class PiecewiseCudaGraphRunner:
|
||||
forward_batch.out_cache_loc
|
||||
)
|
||||
)
|
||||
|
||||
if (
|
||||
self.mamba_track_indices is not None
|
||||
and forward_batch.mamba_track_indices is not None
|
||||
):
|
||||
self.mamba_track_indices[:bs].copy_(forward_batch.mamba_track_indices)
|
||||
if (
|
||||
self.mamba_track_mask is not None
|
||||
and forward_batch.mamba_track_mask is not None
|
||||
):
|
||||
self.mamba_track_mask[:bs].copy_(forward_batch.mamba_track_mask)
|
||||
if (
|
||||
self.mamba_track_seqlens is not None
|
||||
and forward_batch.mamba_track_seqlens is not None
|
||||
):
|
||||
self.mamba_track_seqlens[:bs].copy_(forward_batch.mamba_track_seqlens)
|
||||
|
||||
input_ids = self.input_ids[:static_num_tokens]
|
||||
positions = self.positions[:static_num_tokens]
|
||||
out_cache_loc = self.out_cache_loc[:static_num_tokens]
|
||||
@@ -514,6 +587,19 @@ class PiecewiseCudaGraphRunner:
|
||||
else None
|
||||
)
|
||||
|
||||
mamba_track_indices = (
|
||||
self.mamba_track_indices[:bs]
|
||||
if self.mamba_track_indices is not None
|
||||
else None
|
||||
)
|
||||
mamba_track_mask = (
|
||||
self.mamba_track_mask[:bs] if self.mamba_track_mask is not None else None
|
||||
)
|
||||
mamba_track_seqlens = (
|
||||
self.mamba_track_seqlens[:bs]
|
||||
if self.mamba_track_seqlens is not None
|
||||
else None
|
||||
)
|
||||
if forward_batch.mrope_positions is not None:
|
||||
self.mrope_positions[:, :num_tokens].copy_(forward_batch.mrope_positions)
|
||||
|
||||
@@ -546,6 +632,9 @@ class PiecewiseCudaGraphRunner:
|
||||
out_cache_loc=out_cache_loc,
|
||||
out_cache_loc_swa=out_cache_loc_swa,
|
||||
seq_lens_sum=forward_batch.seq_lens_sum,
|
||||
mamba_track_indices=mamba_track_indices,
|
||||
mamba_track_mask=mamba_track_mask,
|
||||
mamba_track_seqlens=mamba_track_seqlens,
|
||||
encoder_lens=forward_batch.encoder_lens,
|
||||
return_logprob=False,
|
||||
extend_seq_lens=forward_batch.extend_seq_lens,
|
||||
|
||||
Reference in New Issue
Block a user