[sgl] fix incorrect behavior in cuda graph draft extend (#22832)
This commit is contained in:
@@ -538,20 +538,28 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
|
|||||||
|
|
||||||
# Update req_to_hidden_states_pool for KV Cache reversion
|
# Update req_to_hidden_states_pool for KV Cache reversion
|
||||||
if (
|
if (
|
||||||
self.cuda_graph_runner_for_draft_extend is not None
|
forward_batch.extend_seq_lens is not None
|
||||||
and forward_batch.extend_seq_lens is not None
|
and self.cuda_graph_runner_for_draft_extend is not None
|
||||||
):
|
):
|
||||||
last_cuda_graph_runner = (
|
if can_cuda_graph:
|
||||||
self.cuda_graph_runner_for_draft_extend.get_last_runner()
|
last_runner = self.cuda_graph_runner_for_draft_extend.get_last_runner()
|
||||||
)
|
hidden_states = last_runner.buffers.hidden_states
|
||||||
|
req_pool_indices = last_runner.buffers.req_pool_indices
|
||||||
|
extend_seq_lens = last_runner.buffers.extend_seq_lens
|
||||||
|
extend_start_loc = last_runner.buffers.extend_start_loc
|
||||||
|
else:
|
||||||
|
hidden_states = draft_logits_output.logits_output.hidden_states
|
||||||
|
req_pool_indices = forward_batch.req_pool_indices
|
||||||
|
extend_seq_lens = forward_batch.extend_seq_lens
|
||||||
|
extend_start_loc = forward_batch.extend_start_loc
|
||||||
assign_hidden_states_pool_triton(
|
assign_hidden_states_pool_triton(
|
||||||
last_cuda_graph_runner.buffers.hidden_states,
|
hidden_states,
|
||||||
last_cuda_graph_runner.buffers.req_pool_indices,
|
req_pool_indices,
|
||||||
self.req_to_hidden_states_pool,
|
self.req_to_hidden_states_pool,
|
||||||
self.speculative_num_steps - 1,
|
self.speculative_num_steps - 1,
|
||||||
forward_batch.batch_size,
|
forward_batch.batch_size,
|
||||||
last_cuda_graph_runner.buffers.extend_seq_lens,
|
extend_seq_lens,
|
||||||
last_cuda_graph_runner.buffers.extend_start_loc,
|
extend_start_loc,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Reorganize the spec info for the next batch
|
# Reorganize the spec info for the next batch
|
||||||
|
|||||||
Reference in New Issue
Block a user