Batch EAGLE draft/draft-extend replay memcpys via grouped foreach copy (#28465)
Co-authored-by: Cursor <cursoragent@cursor.com> Co-authored-by: Qiaolin Yu <liin1211@outlook.com>
This commit is contained in:
co-authored by
Cursor
Qiaolin Yu
parent
3bc618485a
commit
b54f8432ad
@@ -24,6 +24,7 @@ from sglang.srt.model_executor.runner import (
|
|||||||
DecodeCudaGraphRunner,
|
DecodeCudaGraphRunner,
|
||||||
DeepEPCudaGraphRunnerAdapter,
|
DeepEPCudaGraphRunnerAdapter,
|
||||||
ShapeKey,
|
ShapeKey,
|
||||||
|
_grouped_foreach_copy_,
|
||||||
get_batch_sizes_to_capture,
|
get_batch_sizes_to_capture,
|
||||||
model_capture_mode,
|
model_capture_mode,
|
||||||
)
|
)
|
||||||
@@ -459,21 +460,6 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
|
|
||||||
num_tokens = bs * self.num_tokens_per_bs
|
num_tokens = bs * self.num_tokens_per_bs
|
||||||
|
|
||||||
# Common inputs
|
|
||||||
buffers.seq_lens[:raw_bs].copy_(forward_batch.seq_lens)
|
|
||||||
buffers.out_cache_loc[: raw_num_token * self.speculative_num_steps].copy_(
|
|
||||||
forward_batch.out_cache_loc
|
|
||||||
)
|
|
||||||
buffers.positions[:raw_num_token].copy_(forward_batch.positions)
|
|
||||||
if buffers.rids_int is not None and forward_batch.rids_int is not None:
|
|
||||||
buffers.rids_int[:raw_bs].copy_(forward_batch.rids_int)
|
|
||||||
if (
|
|
||||||
buffers.bootstrap_room_ids_int is not None
|
|
||||||
and forward_batch.bootstrap_room_ids_int is not None
|
|
||||||
):
|
|
||||||
buffers.bootstrap_room_ids_int[:raw_bs].copy_(
|
|
||||||
forward_batch.bootstrap_room_ids_int
|
|
||||||
)
|
|
||||||
maybe_detect_nan(
|
maybe_detect_nan(
|
||||||
forward_batch.spec_info.topk_p,
|
forward_batch.spec_info.topk_p,
|
||||||
"EagleDraftCudaGraphRunner.replay: topk_p",
|
"EagleDraftCudaGraphRunner.replay: topk_p",
|
||||||
@@ -485,14 +471,45 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
"EagleDraftCudaGraphRunner.replay: topk_index vs vocab_size="
|
"EagleDraftCudaGraphRunner.replay: topk_index vs vocab_size="
|
||||||
f"{self.model_runner.model_config.vocab_size}",
|
f"{self.model_runner.model_config.vocab_size}",
|
||||||
)
|
)
|
||||||
buffers.topk_p[:raw_bs].copy_(forward_batch.spec_info.topk_p)
|
|
||||||
buffers.topk_index[:raw_bs].copy_(forward_batch.spec_info.topk_index)
|
# Common inputs — batch the small per-field device copies into a grouped
|
||||||
|
# foreach copy (one foreach call per dtype pair) to cut launch overhead.
|
||||||
|
# hidden_states is handled separately below (see note), and seq_lens_cpu
|
||||||
|
# is handled further down since it lives on host.
|
||||||
|
copy_dsts = [
|
||||||
|
buffers.seq_lens[:raw_bs],
|
||||||
|
buffers.out_cache_loc[: raw_num_token * self.speculative_num_steps],
|
||||||
|
buffers.positions[:raw_num_token],
|
||||||
|
buffers.topk_p[:raw_bs],
|
||||||
|
buffers.topk_index[:raw_bs],
|
||||||
|
buffers.req_pool_indices[:raw_bs],
|
||||||
|
]
|
||||||
|
copy_srcs = [
|
||||||
|
forward_batch.seq_lens,
|
||||||
|
forward_batch.out_cache_loc,
|
||||||
|
forward_batch.positions,
|
||||||
|
forward_batch.spec_info.topk_p,
|
||||||
|
forward_batch.spec_info.topk_index,
|
||||||
|
forward_batch.req_pool_indices,
|
||||||
|
]
|
||||||
|
if buffers.rids_int is not None and forward_batch.rids_int is not None:
|
||||||
|
copy_dsts.append(buffers.rids_int[:raw_bs])
|
||||||
|
copy_srcs.append(forward_batch.rids_int)
|
||||||
|
if (
|
||||||
|
buffers.bootstrap_room_ids_int is not None
|
||||||
|
and forward_batch.bootstrap_room_ids_int is not None
|
||||||
|
):
|
||||||
|
copy_dsts.append(buffers.bootstrap_room_ids_int[:raw_bs])
|
||||||
|
copy_srcs.append(forward_batch.bootstrap_room_ids_int)
|
||||||
|
_grouped_foreach_copy_(copy_dsts, copy_srcs)
|
||||||
|
|
||||||
|
# hidden_states is large + contiguous: copy_() uses the cudaMemcpyAsync
|
||||||
|
# DMA engine; foreach would force the ~3x slower compute-kernel copy.
|
||||||
if (
|
if (
|
||||||
buffers.hidden_states is not None
|
buffers.hidden_states is not None
|
||||||
and forward_batch.spec_info.hidden_states is not None
|
and forward_batch.spec_info.hidden_states is not None
|
||||||
):
|
):
|
||||||
buffers.hidden_states[:raw_bs].copy_(forward_batch.spec_info.hidden_states)
|
buffers.hidden_states[:raw_bs].copy_(forward_batch.spec_info.hidden_states)
|
||||||
buffers.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices)
|
|
||||||
|
|
||||||
# TODO(ch-wan): support num_token_non_padded
|
# TODO(ch-wan): support num_token_non_padded
|
||||||
if self.require_gathered_buffer:
|
if self.require_gathered_buffer:
|
||||||
|
|||||||
@@ -24,6 +24,7 @@ from sglang.srt.model_executor.runner import (
|
|||||||
DecodeCudaGraphRunner,
|
DecodeCudaGraphRunner,
|
||||||
DeepEPCudaGraphRunnerAdapter,
|
DeepEPCudaGraphRunnerAdapter,
|
||||||
ShapeKey,
|
ShapeKey,
|
||||||
|
_grouped_foreach_copy_,
|
||||||
get_batch_sizes_to_capture,
|
get_batch_sizes_to_capture,
|
||||||
model_capture_mode,
|
model_capture_mode,
|
||||||
)
|
)
|
||||||
@@ -455,14 +456,38 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
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)
|
||||||
|
|
||||||
buffers.input_ids[:num_tokens].copy_(forward_batch.input_ids)
|
# Batch the small per-field device copies into a grouped foreach copy
|
||||||
buffers.seq_lens[:raw_bs].copy_(forward_batch.seq_lens)
|
# (one foreach call per dtype pair) to cut launch overhead. hidden_states
|
||||||
|
# is handled separately below (see note), and seq_lens_cpu is handled
|
||||||
|
# further down since it lives on host.
|
||||||
|
copy_dsts = [
|
||||||
|
buffers.input_ids[:num_tokens],
|
||||||
|
buffers.seq_lens[:raw_bs],
|
||||||
|
buffers.out_cache_loc[:num_tokens],
|
||||||
|
buffers.positions[:num_tokens],
|
||||||
|
buffers.req_pool_indices[:raw_bs],
|
||||||
|
]
|
||||||
|
copy_srcs = [
|
||||||
|
forward_batch.input_ids,
|
||||||
|
forward_batch.seq_lens,
|
||||||
|
forward_batch.out_cache_loc,
|
||||||
|
forward_batch.positions,
|
||||||
|
forward_batch.req_pool_indices,
|
||||||
|
]
|
||||||
if forward_batch.extend_seq_lens is not None:
|
if forward_batch.extend_seq_lens is not None:
|
||||||
buffers.extend_seq_lens[:raw_bs].copy_(forward_batch.extend_seq_lens)
|
copy_dsts.append(buffers.extend_seq_lens[:raw_bs])
|
||||||
|
copy_srcs.append(forward_batch.extend_seq_lens)
|
||||||
else:
|
else:
|
||||||
buffers.extend_seq_lens[:raw_bs].fill_(self.num_tokens_per_bs)
|
buffers.extend_seq_lens[:raw_bs].fill_(self.num_tokens_per_bs)
|
||||||
buffers.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc)
|
if forward_batch.spec_info.num_correct_drafts is not None:
|
||||||
buffers.positions[:num_tokens].copy_(forward_batch.positions)
|
copy_dsts.append(buffers.num_correct_drafts[:raw_bs])
|
||||||
|
copy_srcs.append(forward_batch.spec_info.num_correct_drafts)
|
||||||
|
copy_dsts.append(buffers.num_accept_tokens[:raw_bs])
|
||||||
|
copy_srcs.append(forward_batch.spec_info.num_accept_tokens)
|
||||||
|
_grouped_foreach_copy_(copy_dsts, copy_srcs)
|
||||||
|
|
||||||
|
# hidden_states is large + contiguous: copy_() uses the cudaMemcpyAsync
|
||||||
|
# DMA engine; foreach would force the ~3x slower compute-kernel copy.
|
||||||
if (
|
if (
|
||||||
buffers.hidden_states is not None
|
buffers.hidden_states is not None
|
||||||
and forward_batch.spec_info.hidden_states is not None
|
and forward_batch.spec_info.hidden_states is not None
|
||||||
@@ -472,14 +497,6 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
buffers.hidden_states[:num_tokens].copy_(
|
buffers.hidden_states[:num_tokens].copy_(
|
||||||
forward_batch.spec_info.hidden_states
|
forward_batch.spec_info.hidden_states
|
||||||
)
|
)
|
||||||
if forward_batch.spec_info.num_correct_drafts is not None:
|
|
||||||
buffers.num_correct_drafts[:raw_bs].copy_(
|
|
||||||
forward_batch.spec_info.num_correct_drafts
|
|
||||||
)
|
|
||||||
buffers.num_accept_tokens[:raw_bs].copy_(
|
|
||||||
forward_batch.spec_info.num_accept_tokens
|
|
||||||
)
|
|
||||||
buffers.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices)
|
|
||||||
|
|
||||||
# TODO(ch-wan): support num_token_non_padded
|
# TODO(ch-wan): support num_token_non_padded
|
||||||
if self.require_gathered_buffer:
|
if self.require_gathered_buffer:
|
||||||
|
|||||||
Reference in New Issue
Block a user