Zero req_pool_indices padding in cuda-graph populate (#26292)
This commit is contained in:
@@ -786,6 +786,9 @@ class CPUGraphRunner:
|
|||||||
assert captured_forward_batch is not None
|
assert captured_forward_batch is not None
|
||||||
captured_forward_batch.seq_lens.fill_(self.seq_len_fill_value)
|
captured_forward_batch.seq_lens.fill_(self.seq_len_fill_value)
|
||||||
captured_forward_batch.out_cache_loc.zero_()
|
captured_forward_batch.out_cache_loc.zero_()
|
||||||
|
# Pair with seq_lens fill: padded rows must point at reserved
|
||||||
|
# req_pool slot 0 (req_to_token[0, :] is all zeros from init).
|
||||||
|
captured_forward_batch.req_pool_indices.zero_()
|
||||||
captured_forward_batch.input_ids[:raw_num_token].copy_(forward_batch.input_ids)
|
captured_forward_batch.input_ids[:raw_num_token].copy_(forward_batch.input_ids)
|
||||||
captured_forward_batch.req_pool_indices[:raw_bs].copy_(
|
captured_forward_batch.req_pool_indices[:raw_bs].copy_(
|
||||||
forward_batch.req_pool_indices
|
forward_batch.req_pool_indices
|
||||||
|
|||||||
@@ -286,6 +286,11 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
if bs != raw_bs:
|
if bs != raw_bs:
|
||||||
self.seq_lens.fill_(seq_len_fill_value)
|
self.seq_lens.fill_(seq_len_fill_value)
|
||||||
self.out_cache_loc.zero_()
|
self.out_cache_loc.zero_()
|
||||||
|
# Pair with seq_lens fill: padded rows must point at reserved
|
||||||
|
# req_pool slot 0 (req_to_token[0, :] is all zeros from init),
|
||||||
|
# so dummy attention reads land on slot 0 instead of a stale
|
||||||
|
# req_to_token row left by an earlier replay.
|
||||||
|
self.req_pool_indices.zero_()
|
||||||
if self.mamba_track_indices is not None:
|
if self.mamba_track_indices is not None:
|
||||||
self.mamba_track_indices.zero_()
|
self.mamba_track_indices.zero_()
|
||||||
if self.mamba_track_mask is not None:
|
if self.mamba_track_mask is not None:
|
||||||
|
|||||||
@@ -460,6 +460,9 @@ class EAGLEDraftExtendCudaGraphRunner:
|
|||||||
buffers.seq_lens.fill_(self.seq_len_fill_value)
|
buffers.seq_lens.fill_(self.seq_len_fill_value)
|
||||||
buffers.out_cache_loc.zero_()
|
buffers.out_cache_loc.zero_()
|
||||||
buffers.positions.zero_()
|
buffers.positions.zero_()
|
||||||
|
# Pair with seq_lens fill: padded rows must point at reserved
|
||||||
|
# req_pool slot 0 (req_to_token[0, :] is all zeros from init).
|
||||||
|
buffers.req_pool_indices.zero_()
|
||||||
buffers.num_correct_drafts.fill_(self.num_tokens_per_bs)
|
buffers.num_correct_drafts.fill_(self.num_tokens_per_bs)
|
||||||
buffers.num_accept_tokens.fill_(self.num_tokens_per_bs)
|
buffers.num_accept_tokens.fill_(self.num_tokens_per_bs)
|
||||||
buffers.extend_seq_lens.fill_(self.num_tokens_per_bs)
|
buffers.extend_seq_lens.fill_(self.num_tokens_per_bs)
|
||||||
|
|||||||
@@ -352,6 +352,9 @@ class FrozenKVMTPCudaGraphRunner:
|
|||||||
if bs != raw_bs:
|
if bs != raw_bs:
|
||||||
buffers.seq_lens.fill_(self.seq_len_fill_value)
|
buffers.seq_lens.fill_(self.seq_len_fill_value)
|
||||||
buffers.positions.zero_()
|
buffers.positions.zero_()
|
||||||
|
# Pair with seq_lens fill: padded rows must point at reserved
|
||||||
|
# req_pool slot 0 (req_to_token[0, :] is all zeros from init).
|
||||||
|
buffers.req_pool_indices.zero_()
|
||||||
|
|
||||||
num_tokens = expanded_bs
|
num_tokens = expanded_bs
|
||||||
buffers.seq_lens[:raw_expanded_bs].copy_(forward_batch.seq_lens)
|
buffers.seq_lens[:raw_expanded_bs].copy_(forward_batch.seq_lens)
|
||||||
|
|||||||
Reference in New Issue
Block a user