Refactor graph input buffers (#18991)
This commit is contained in:
@@ -1,7 +1,8 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import bisect
|
||||
from typing import TYPE_CHECKING, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Callable, Optional
|
||||
|
||||
import torch
|
||||
|
||||
@@ -22,6 +23,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardBatch,
|
||||
ForwardMode,
|
||||
)
|
||||
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
||||
from sglang.srt.speculative.eagle_info import EagleDraftInput
|
||||
from sglang.srt.utils import (
|
||||
require_attn_tp_gather,
|
||||
@@ -34,6 +36,23 @@ if TYPE_CHECKING:
|
||||
from sglang.srt.speculative.eagle_worker import EAGLEWorker
|
||||
|
||||
|
||||
@dataclass
|
||||
class EagleDraftInputBuffers(ForwardInputBuffers):
|
||||
input_ids: torch.Tensor
|
||||
req_pool_indices: torch.Tensor
|
||||
out_cache_loc: torch.Tensor
|
||||
positions: torch.Tensor
|
||||
mrope_positions: torch.Tensor
|
||||
seq_lens: torch.Tensor
|
||||
seq_lens_cpu: torch.Tensor
|
||||
extend_seq_lens: torch.Tensor
|
||||
topk_p: torch.Tensor
|
||||
topk_index: torch.Tensor
|
||||
hidden_states: torch.Tensor
|
||||
global_num_tokens_gpu: Optional[torch.Tensor]
|
||||
global_num_tokens_for_logprob_gpu: Optional[torch.Tensor]
|
||||
|
||||
|
||||
class EAGLEDraftCudaGraphRunner:
|
||||
def __init__(self, eagle_worker: EAGLEWorker):
|
||||
# Parse args
|
||||
@@ -75,7 +94,7 @@ class EAGLEDraftCudaGraphRunner:
|
||||
self.seq_len_fill_value = self.model_runner.draft_attn_backend.attn_backends[
|
||||
0
|
||||
].get_cuda_graph_seq_len_fill_value()
|
||||
self.seq_lens_cpu = torch.full(
|
||||
seq_lens_cpu = torch.full(
|
||||
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int32
|
||||
)
|
||||
self.extend_seq_lens_cpu = [self.seq_len_fill_value] * self.max_bs
|
||||
@@ -85,44 +104,59 @@ class EAGLEDraftCudaGraphRunner:
|
||||
|
||||
# Graph inputs
|
||||
with torch.device(model_runner.device):
|
||||
self.input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64)
|
||||
self.req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int32)
|
||||
self.out_cache_loc = torch.zeros(
|
||||
input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64)
|
||||
req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int32)
|
||||
out_cache_loc = torch.zeros(
|
||||
(self.max_num_token * self.speculative_num_steps,),
|
||||
dtype=self._cache_loc_dtype(),
|
||||
)
|
||||
self.positions = torch.zeros((self.max_num_token,), dtype=torch.int64)
|
||||
self.mrope_positions = torch.zeros(
|
||||
(3, self.max_num_token), dtype=torch.int64
|
||||
)
|
||||
self.seq_lens = torch.full(
|
||||
positions = torch.zeros((self.max_num_token,), dtype=torch.int64)
|
||||
mrope_positions = torch.zeros((3, self.max_num_token), dtype=torch.int64)
|
||||
seq_lens = torch.full(
|
||||
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int32
|
||||
)
|
||||
self.extend_seq_lens = torch.ones((self.max_bs,), dtype=torch.int32)
|
||||
self.topk_p = torch.zeros((self.max_bs, self.topk), dtype=torch.float32)
|
||||
self.topk_index = torch.zeros((self.max_bs, self.topk), dtype=torch.int64)
|
||||
self.hidden_states = torch.zeros(
|
||||
extend_seq_lens = torch.ones((self.max_bs,), dtype=torch.int32)
|
||||
topk_p = torch.zeros((self.max_bs, self.topk), dtype=torch.float32)
|
||||
topk_index = torch.zeros((self.max_bs, self.topk), dtype=torch.int64)
|
||||
hidden_states = torch.zeros(
|
||||
(self.max_bs, self.model_runner.model_config.hidden_size),
|
||||
dtype=self.model_runner.dtype,
|
||||
)
|
||||
|
||||
if self.require_gathered_buffer:
|
||||
if self.require_mlp_tp_gather:
|
||||
self.global_num_tokens_gpu = torch.zeros(
|
||||
global_num_tokens_gpu = torch.zeros(
|
||||
(self.dp_size,), dtype=torch.int32
|
||||
)
|
||||
self.global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||
global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||
(self.dp_size,), dtype=torch.int32
|
||||
)
|
||||
else:
|
||||
assert self.require_attn_tp_gather
|
||||
self.global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
|
||||
self.global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||
global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
|
||||
global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||
(1,), dtype=torch.int32
|
||||
)
|
||||
else:
|
||||
self.global_num_tokens_gpu = None
|
||||
self.global_num_tokens_for_logprob_gpu = None
|
||||
global_num_tokens_gpu = None
|
||||
global_num_tokens_for_logprob_gpu = None
|
||||
|
||||
self.buffers = EagleDraftInputBuffers(
|
||||
input_ids=input_ids,
|
||||
req_pool_indices=req_pool_indices,
|
||||
out_cache_loc=out_cache_loc,
|
||||
positions=positions,
|
||||
mrope_positions=mrope_positions,
|
||||
seq_lens=seq_lens,
|
||||
seq_lens_cpu=seq_lens_cpu,
|
||||
extend_seq_lens=extend_seq_lens,
|
||||
topk_p=topk_p,
|
||||
topk_index=topk_index,
|
||||
hidden_states=hidden_states,
|
||||
global_num_tokens_gpu=global_num_tokens_gpu,
|
||||
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
|
||||
)
|
||||
self.buffers.share_buffers()
|
||||
|
||||
# Capture
|
||||
try:
|
||||
@@ -181,59 +215,60 @@ class EAGLEDraftCudaGraphRunner:
|
||||
def capture_one_batch_size(
|
||||
self, num_seqs: int, forward: Callable, stream_idx: int = 0
|
||||
):
|
||||
buffers = self.buffers
|
||||
graph = self._create_graph()
|
||||
stream = self.stream
|
||||
num_tokens = num_seqs * self.num_tokens_per_bs
|
||||
|
||||
# Graph inputs
|
||||
req_pool_indices = self.req_pool_indices[:num_seqs]
|
||||
seq_lens = self.seq_lens[:num_seqs]
|
||||
seq_lens_cpu = self.seq_lens_cpu[:num_seqs]
|
||||
extend_seq_lens = self.extend_seq_lens[:num_seqs]
|
||||
req_pool_indices = buffers.req_pool_indices[:num_seqs]
|
||||
seq_lens = buffers.seq_lens[:num_seqs]
|
||||
seq_lens_cpu = buffers.seq_lens_cpu[:num_seqs]
|
||||
extend_seq_lens = buffers.extend_seq_lens[:num_seqs]
|
||||
extend_seq_lens_cpu = self.extend_seq_lens_cpu[:num_seqs]
|
||||
out_cache_loc = self.out_cache_loc[: num_tokens * self.speculative_num_steps]
|
||||
positions = self.positions[:num_tokens]
|
||||
mrope_positions = self.mrope_positions[:, :num_tokens]
|
||||
hidden_states = self.hidden_states[:num_seqs]
|
||||
topk_p = self.topk_p[:num_seqs]
|
||||
topk_index = self.topk_index[:num_seqs]
|
||||
out_cache_loc = buffers.out_cache_loc[: num_tokens * self.speculative_num_steps]
|
||||
positions = buffers.positions[:num_tokens]
|
||||
mrope_positions = buffers.mrope_positions[:, :num_tokens]
|
||||
hidden_states = buffers.hidden_states[:num_seqs]
|
||||
topk_p = buffers.topk_p[:num_seqs]
|
||||
topk_index = buffers.topk_index[:num_seqs]
|
||||
|
||||
if self.require_mlp_tp_gather:
|
||||
self.global_num_tokens_gpu.copy_(
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens] * self.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=self.input_ids.device,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
self.global_num_tokens_for_logprob_gpu.copy_(
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens] * self.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=self.input_ids.device,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
global_num_tokens = self.global_num_tokens_gpu
|
||||
global_num_tokens = buffers.global_num_tokens_gpu
|
||||
global_dp_buffer_len = num_tokens * self.dp_size
|
||||
global_num_tokens_for_logprob = self.global_num_tokens_for_logprob_gpu
|
||||
global_num_tokens_for_logprob = buffers.global_num_tokens_for_logprob_gpu
|
||||
elif self.require_attn_tp_gather:
|
||||
self.global_num_tokens_gpu.copy_(
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens],
|
||||
dtype=torch.int32,
|
||||
device=self.input_ids.device,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
self.global_num_tokens_for_logprob_gpu.copy_(
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens],
|
||||
dtype=torch.int32,
|
||||
device=self.input_ids.device,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
global_num_tokens = self.global_num_tokens_gpu
|
||||
global_num_tokens = buffers.global_num_tokens_gpu
|
||||
global_dp_buffer_len = num_tokens
|
||||
global_num_tokens_for_logprob = self.global_num_tokens_for_logprob_gpu
|
||||
global_num_tokens_for_logprob = buffers.global_num_tokens_for_logprob_gpu
|
||||
else:
|
||||
global_num_tokens = None
|
||||
global_dp_buffer_len = None
|
||||
@@ -319,6 +354,7 @@ class EAGLEDraftCudaGraphRunner:
|
||||
def replay(self, forward_batch: ForwardBatch):
|
||||
assert forward_batch.out_cache_loc is not None
|
||||
self.deepep_adapter.replay()
|
||||
buffers = self.buffers
|
||||
|
||||
raw_bs = forward_batch.batch_size
|
||||
raw_num_token = raw_bs * self.num_tokens_per_bs
|
||||
@@ -338,40 +374,40 @@ class EAGLEDraftCudaGraphRunner:
|
||||
|
||||
bs = self.capture_bs[index]
|
||||
if bs != raw_bs:
|
||||
self.seq_lens.fill_(self.seq_len_fill_value)
|
||||
self.out_cache_loc.zero_()
|
||||
self.positions.zero_()
|
||||
buffers.seq_lens.fill_(self.seq_len_fill_value)
|
||||
buffers.out_cache_loc.zero_()
|
||||
buffers.positions.zero_()
|
||||
|
||||
num_tokens = bs * self.num_tokens_per_bs
|
||||
|
||||
# Common inputs
|
||||
self.seq_lens[:raw_bs].copy_(forward_batch.seq_lens)
|
||||
self.out_cache_loc[: raw_num_token * self.speculative_num_steps].copy_(
|
||||
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
|
||||
)
|
||||
self.positions[:raw_num_token].copy_(forward_batch.positions)
|
||||
self.topk_p[:raw_bs].copy_(forward_batch.spec_info.topk_p)
|
||||
self.topk_index[:raw_bs].copy_(forward_batch.spec_info.topk_index)
|
||||
self.hidden_states[:raw_bs].copy_(forward_batch.spec_info.hidden_states)
|
||||
self.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices)
|
||||
buffers.positions[:raw_num_token].copy_(forward_batch.positions)
|
||||
buffers.topk_p[:raw_bs].copy_(forward_batch.spec_info.topk_p)
|
||||
buffers.topk_index[:raw_bs].copy_(forward_batch.spec_info.topk_index)
|
||||
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
|
||||
if self.require_gathered_buffer:
|
||||
self.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs)
|
||||
self.global_num_tokens_for_logprob_gpu.fill_(bs * self.num_tokens_per_bs)
|
||||
buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs)
|
||||
buffers.global_num_tokens_for_logprob_gpu.fill_(bs * self.num_tokens_per_bs)
|
||||
|
||||
# Attention backend
|
||||
if bs != raw_bs:
|
||||
forward_batch.batch_size = bs
|
||||
forward_batch.seq_lens = self.seq_lens[:bs]
|
||||
forward_batch.req_pool_indices = self.req_pool_indices[:bs]
|
||||
forward_batch.positions = self.positions[:num_tokens]
|
||||
forward_batch.seq_lens = buffers.seq_lens[:bs]
|
||||
forward_batch.req_pool_indices = buffers.req_pool_indices[:bs]
|
||||
forward_batch.positions = buffers.positions[:num_tokens]
|
||||
|
||||
if forward_batch.seq_lens_cpu is not None:
|
||||
if bs != raw_bs:
|
||||
self.seq_lens_cpu.fill_(self.seq_len_fill_value)
|
||||
self.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
|
||||
forward_batch.seq_lens_cpu = self.seq_lens_cpu[:bs]
|
||||
buffers.seq_lens_cpu.fill_(self.seq_len_fill_value)
|
||||
buffers.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
|
||||
forward_batch.seq_lens_cpu = buffers.seq_lens_cpu[:bs]
|
||||
|
||||
self.model_runner.draft_attn_backend.init_forward_metadata_replay_cuda_graph(
|
||||
forward_batch, bs
|
||||
@@ -387,10 +423,10 @@ class EAGLEDraftCudaGraphRunner:
|
||||
if bs != raw_bs:
|
||||
out = self._postprocess_output_to_raw_bs(out, raw_bs)
|
||||
forward_batch.batch_size = raw_bs
|
||||
forward_batch.positions = self.positions[:raw_num_token]
|
||||
forward_batch.seq_lens = self.seq_lens[:raw_bs]
|
||||
forward_batch.req_pool_indices = self.req_pool_indices[:raw_bs]
|
||||
forward_batch.positions = buffers.positions[:raw_num_token]
|
||||
forward_batch.seq_lens = buffers.seq_lens[:raw_bs]
|
||||
forward_batch.req_pool_indices = buffers.req_pool_indices[:raw_bs]
|
||||
if forward_batch.seq_lens_cpu is not None:
|
||||
forward_batch.seq_lens_cpu = self.seq_lens_cpu[:raw_bs]
|
||||
forward_batch.seq_lens_cpu = buffers.seq_lens_cpu[:raw_bs]
|
||||
|
||||
return out
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import bisect
|
||||
from typing import TYPE_CHECKING, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Callable, Optional
|
||||
|
||||
import torch
|
||||
|
||||
@@ -23,6 +24,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardBatch,
|
||||
ForwardMode,
|
||||
)
|
||||
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
||||
from sglang.srt.speculative.eagle_info import EagleDraftInput
|
||||
from sglang.srt.speculative.spec_utils import fast_topk
|
||||
from sglang.srt.utils import (
|
||||
@@ -36,6 +38,23 @@ if TYPE_CHECKING:
|
||||
from sglang.srt.speculative.eagle_worker import EAGLEWorker
|
||||
|
||||
|
||||
@dataclass
|
||||
class EagleDraftExtendInputBuffers(ForwardInputBuffers):
|
||||
input_ids: torch.Tensor
|
||||
req_pool_indices: torch.Tensor
|
||||
out_cache_loc: torch.Tensor
|
||||
positions: torch.Tensor
|
||||
mrope_positions: torch.Tensor
|
||||
hidden_states: torch.Tensor
|
||||
seq_lens: torch.Tensor
|
||||
seq_lens_cpu: torch.Tensor
|
||||
extend_seq_lens: torch.Tensor
|
||||
accept_length: torch.Tensor
|
||||
next_token_logits_buffer: torch.Tensor
|
||||
global_num_tokens_gpu: Optional[torch.Tensor]
|
||||
global_num_tokens_for_logprob_gpu: Optional[torch.Tensor]
|
||||
|
||||
|
||||
class EAGLEDraftExtendCudaGraphRunner:
|
||||
def __init__(self, eagle_worker: EAGLEWorker):
|
||||
# Parse args
|
||||
@@ -80,7 +99,7 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
self.seq_len_fill_value = (
|
||||
self.eagle_worker.draft_extend_attn_backend.get_cuda_graph_seq_len_fill_value()
|
||||
)
|
||||
self.seq_lens_cpu = torch.full(
|
||||
seq_lens_cpu = torch.full(
|
||||
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int32
|
||||
)
|
||||
self.extend_seq_lens_cpu = [self.num_tokens_per_bs] * self.max_bs
|
||||
@@ -90,21 +109,19 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
|
||||
# Graph inputs
|
||||
with torch.device(model_runner.device):
|
||||
self.input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64)
|
||||
self.req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int32)
|
||||
self.out_cache_loc = torch.ones(
|
||||
input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64)
|
||||
req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int32)
|
||||
out_cache_loc = torch.ones(
|
||||
(self.max_num_token,), dtype=self._cache_loc_dtype()
|
||||
)
|
||||
self.positions = torch.zeros((self.max_num_token,), dtype=torch.int64)
|
||||
self.mrope_positions = torch.zeros(
|
||||
(3, self.max_num_token), dtype=torch.int64
|
||||
)
|
||||
positions = torch.zeros((self.max_num_token,), dtype=torch.int64)
|
||||
mrope_positions = torch.zeros((3, self.max_num_token), dtype=torch.int64)
|
||||
|
||||
if (
|
||||
self.eagle_worker.speculative_algorithm.is_eagle3()
|
||||
and self.eagle_worker.eagle_use_aux_hidden_state
|
||||
):
|
||||
self.hidden_states = torch.zeros(
|
||||
hidden_states = torch.zeros(
|
||||
(
|
||||
self.max_num_token,
|
||||
(
|
||||
@@ -120,40 +137,40 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
dtype=self.model_runner.dtype,
|
||||
)
|
||||
else:
|
||||
self.hidden_states = torch.zeros(
|
||||
hidden_states = torch.zeros(
|
||||
(self.max_num_token, self.model_runner.model_config.hidden_size),
|
||||
dtype=self.model_runner.dtype,
|
||||
)
|
||||
self.seq_len_fill_value = (
|
||||
self.model_runner.attn_backend.get_cuda_graph_seq_len_fill_value()
|
||||
)
|
||||
self.seq_lens = torch.full(
|
||||
seq_lens = torch.full(
|
||||
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int32
|
||||
)
|
||||
self.extend_seq_lens = torch.full(
|
||||
extend_seq_lens = torch.full(
|
||||
(self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32
|
||||
)
|
||||
self.accept_length = torch.full(
|
||||
accept_length = torch.full(
|
||||
(self.max_bs,), self.num_tokens_per_bs, dtype=torch.int32
|
||||
)
|
||||
|
||||
if self.require_gathered_buffer:
|
||||
if self.require_mlp_tp_gather:
|
||||
self.global_num_tokens_gpu = torch.zeros(
|
||||
global_num_tokens_gpu = torch.zeros(
|
||||
(self.dp_size,), dtype=torch.int32
|
||||
)
|
||||
self.global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||
global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||
(self.dp_size,), dtype=torch.int32
|
||||
)
|
||||
else:
|
||||
assert self.require_attn_tp_gather
|
||||
self.global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
|
||||
self.global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||
global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
|
||||
global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||
(1,), dtype=torch.int32
|
||||
)
|
||||
else:
|
||||
self.global_num_tokens_gpu = None
|
||||
self.global_num_tokens_for_logprob_gpu = None
|
||||
global_num_tokens_gpu = None
|
||||
global_num_tokens_for_logprob_gpu = None
|
||||
|
||||
if hasattr(
|
||||
self.model_runner.model_config.hf_config, "draft_vocab_size"
|
||||
@@ -166,7 +183,7 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
else:
|
||||
vocab_size = self.model_runner.model_config.vocab_size
|
||||
|
||||
self.next_token_logits_buffer = torch.zeros(
|
||||
next_token_logits_buffer = torch.zeros(
|
||||
(
|
||||
(
|
||||
self.max_bs * self.num_tokens_per_bs
|
||||
@@ -178,6 +195,23 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
dtype=torch.float,
|
||||
)
|
||||
|
||||
self.buffers = EagleDraftExtendInputBuffers(
|
||||
input_ids=input_ids,
|
||||
req_pool_indices=req_pool_indices,
|
||||
out_cache_loc=out_cache_loc,
|
||||
positions=positions,
|
||||
mrope_positions=mrope_positions,
|
||||
hidden_states=hidden_states,
|
||||
seq_lens=seq_lens,
|
||||
seq_lens_cpu=seq_lens_cpu,
|
||||
extend_seq_lens=extend_seq_lens,
|
||||
accept_length=accept_length,
|
||||
next_token_logits_buffer=next_token_logits_buffer,
|
||||
global_num_tokens_gpu=global_num_tokens_gpu,
|
||||
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
|
||||
)
|
||||
self.buffers.share_buffers()
|
||||
|
||||
# Capture
|
||||
try:
|
||||
with model_capture_mode():
|
||||
@@ -233,23 +267,24 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
CudaGraphRunner.capture(self)
|
||||
|
||||
def capture_one_batch_size(self, bs: int, forward: Callable, stream_idx: int = 0):
|
||||
buffers = self.buffers
|
||||
graph = self._create_graph()
|
||||
stream = self.stream
|
||||
num_tokens = bs * self.num_tokens_per_bs
|
||||
|
||||
# Graph inputs
|
||||
input_ids = self.input_ids[:num_tokens]
|
||||
req_pool_indices = self.req_pool_indices[:bs]
|
||||
seq_lens = self.seq_lens[:bs]
|
||||
seq_lens_cpu = self.seq_lens_cpu[:bs]
|
||||
extend_seq_lens = self.extend_seq_lens[:bs]
|
||||
input_ids = buffers.input_ids[:num_tokens]
|
||||
req_pool_indices = buffers.req_pool_indices[:bs]
|
||||
seq_lens = buffers.seq_lens[:bs]
|
||||
seq_lens_cpu = buffers.seq_lens_cpu[:bs]
|
||||
extend_seq_lens = buffers.extend_seq_lens[:bs]
|
||||
extend_seq_lens_cpu = self.extend_seq_lens_cpu[:bs]
|
||||
out_cache_loc = self.out_cache_loc[:num_tokens]
|
||||
positions = self.positions[:num_tokens]
|
||||
mrope_positions = self.mrope_positions[:, :num_tokens]
|
||||
hidden_states = self.hidden_states[:num_tokens]
|
||||
accept_length = self.accept_length[:bs]
|
||||
next_token_logits_buffer = self.next_token_logits_buffer[
|
||||
out_cache_loc = buffers.out_cache_loc[:num_tokens]
|
||||
positions = buffers.positions[:num_tokens]
|
||||
mrope_positions = buffers.mrope_positions[:, :num_tokens]
|
||||
hidden_states = buffers.hidden_states[:num_tokens]
|
||||
accept_length = buffers.accept_length[:bs]
|
||||
next_token_logits_buffer = buffers.next_token_logits_buffer[
|
||||
: bs if self.forward_mode == ForwardMode.DRAFT_EXTEND else num_tokens
|
||||
]
|
||||
|
||||
@@ -260,34 +295,34 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
)
|
||||
|
||||
if self.require_mlp_tp_gather:
|
||||
self.global_num_tokens_gpu.copy_(
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens] * self.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=self.input_ids.device,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
self.global_num_tokens_for_logprob_gpu.copy_(
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens_for_logprob] * self.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=self.input_ids.device,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
global_dp_buffer_len = num_tokens * self.dp_size
|
||||
elif self.require_attn_tp_gather:
|
||||
self.global_num_tokens_gpu.copy_(
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens],
|
||||
dtype=torch.int32,
|
||||
device=self.input_ids.device,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
self.global_num_tokens_for_logprob_gpu.copy_(
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens_for_logprob],
|
||||
dtype=torch.int32,
|
||||
device=self.input_ids.device,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
global_dp_buffer_len = num_tokens
|
||||
@@ -320,8 +355,8 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
return_logprob=False,
|
||||
positions=positions,
|
||||
mrope_positions=mrope_positions,
|
||||
global_num_tokens_gpu=self.global_num_tokens_gpu,
|
||||
global_num_tokens_for_logprob_gpu=self.global_num_tokens_for_logprob_gpu,
|
||||
global_num_tokens_gpu=buffers.global_num_tokens_gpu,
|
||||
global_num_tokens_for_logprob_gpu=buffers.global_num_tokens_for_logprob_gpu,
|
||||
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
|
||||
global_dp_buffer_len=global_dp_buffer_len,
|
||||
spec_algorithm=self.model_runner.spec_algorithm,
|
||||
@@ -380,6 +415,7 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
def replay(self, forward_batch: ForwardBatch):
|
||||
assert forward_batch.out_cache_loc is not None
|
||||
self.deepep_adapter.replay()
|
||||
buffers = self.buffers
|
||||
|
||||
# batch_size and num_seqs can be different in case there are finished examples
|
||||
# in the batch, which will not be counted as num_seqs
|
||||
@@ -398,45 +434,47 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
|
||||
bs = self.capture_bs[index]
|
||||
if bs * self.num_tokens_per_bs != num_tokens:
|
||||
self.seq_lens.fill_(self.seq_len_fill_value)
|
||||
self.out_cache_loc.zero_()
|
||||
self.positions.zero_()
|
||||
self.accept_length.fill_(self.num_tokens_per_bs)
|
||||
self.extend_seq_lens.fill_(self.num_tokens_per_bs)
|
||||
buffers.seq_lens.fill_(self.seq_len_fill_value)
|
||||
buffers.out_cache_loc.zero_()
|
||||
buffers.positions.zero_()
|
||||
buffers.accept_length.fill_(self.num_tokens_per_bs)
|
||||
buffers.extend_seq_lens.fill_(self.num_tokens_per_bs)
|
||||
|
||||
# Common inputs
|
||||
self.input_ids[:num_tokens].copy_(forward_batch.input_ids)
|
||||
self.seq_lens[:raw_bs].copy_(forward_batch.seq_lens)
|
||||
buffers.input_ids[:num_tokens].copy_(forward_batch.input_ids)
|
||||
buffers.seq_lens[:raw_bs].copy_(forward_batch.seq_lens)
|
||||
if forward_batch.extend_seq_lens is not None:
|
||||
self.extend_seq_lens[:raw_bs].copy_(forward_batch.extend_seq_lens)
|
||||
buffers.extend_seq_lens[:raw_bs].copy_(forward_batch.extend_seq_lens)
|
||||
else:
|
||||
self.extend_seq_lens[:raw_bs].fill_(self.num_tokens_per_bs)
|
||||
self.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc)
|
||||
self.positions[:num_tokens].copy_(forward_batch.positions)
|
||||
buffers.extend_seq_lens[:raw_bs].fill_(self.num_tokens_per_bs)
|
||||
buffers.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc)
|
||||
buffers.positions[:num_tokens].copy_(forward_batch.positions)
|
||||
if (
|
||||
forward_batch.spec_info.hidden_states.shape[1]
|
||||
== self.hidden_states.shape[1]
|
||||
== buffers.hidden_states.shape[1]
|
||||
):
|
||||
self.hidden_states[:num_tokens].copy_(forward_batch.spec_info.hidden_states)
|
||||
buffers.hidden_states[:num_tokens].copy_(
|
||||
forward_batch.spec_info.hidden_states
|
||||
)
|
||||
if forward_batch.spec_info.accept_length is not None:
|
||||
self.accept_length[:raw_bs].copy_(forward_batch.spec_info.accept_length)
|
||||
self.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices)
|
||||
buffers.accept_length[:raw_bs].copy_(forward_batch.spec_info.accept_length)
|
||||
buffers.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices)
|
||||
|
||||
# TODO(ch-wan): support num_token_non_padded
|
||||
if self.require_gathered_buffer:
|
||||
self.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs)
|
||||
buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs)
|
||||
# V1: pruned_states = bs; V2: pruned_states = num_tokens
|
||||
if self.forward_mode.is_draft_extend_v2():
|
||||
self.global_num_tokens_for_logprob_gpu.fill_(
|
||||
buffers.global_num_tokens_for_logprob_gpu.fill_(
|
||||
bs * self.num_tokens_per_bs
|
||||
)
|
||||
else:
|
||||
self.global_num_tokens_for_logprob_gpu.fill_(bs)
|
||||
buffers.global_num_tokens_for_logprob_gpu.fill_(bs)
|
||||
|
||||
if forward_batch.seq_lens_cpu is not None:
|
||||
if bs != raw_bs:
|
||||
self.seq_lens_cpu.fill_(self.seq_len_fill_value)
|
||||
self.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
|
||||
buffers.seq_lens_cpu.fill_(self.seq_len_fill_value)
|
||||
buffers.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
|
||||
|
||||
if forward_batch.extend_seq_lens_cpu is not None:
|
||||
self.extend_seq_lens_cpu[:raw_bs] = forward_batch.extend_seq_lens_cpu
|
||||
@@ -449,22 +487,22 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
forward_batch.spec_info.extend_seq_lens_cpu = list(
|
||||
self.extend_seq_lens_cpu[:bs]
|
||||
)
|
||||
forward_batch.spec_info.extend_seq_lens_tensor = self.extend_seq_lens[:bs]
|
||||
forward_batch.spec_info.extend_seq_lens_tensor = buffers.extend_seq_lens[:bs]
|
||||
|
||||
if bs != raw_bs:
|
||||
forward_batch.spec_info.positions = self.positions[:num_tokens]
|
||||
forward_batch.spec_info.accept_length = self.accept_length[:bs]
|
||||
forward_batch.spec_info.positions = buffers.positions[:num_tokens]
|
||||
forward_batch.spec_info.accept_length = buffers.accept_length[:bs]
|
||||
|
||||
self.eagle_worker.draft_extend_attn_backend.init_forward_metadata_replay_cuda_graph(
|
||||
bs=bs,
|
||||
req_pool_indices=self.req_pool_indices,
|
||||
seq_lens=self.seq_lens,
|
||||
req_pool_indices=buffers.req_pool_indices,
|
||||
seq_lens=buffers.seq_lens,
|
||||
seq_lens_sum=forward_batch.seq_lens_sum
|
||||
+ (bs - raw_bs) * self.seq_len_fill_value,
|
||||
encoder_lens=None,
|
||||
forward_mode=self.forward_mode,
|
||||
spec_info=forward_batch.spec_info,
|
||||
seq_lens_cpu=self.seq_lens_cpu,
|
||||
seq_lens_cpu=buffers.seq_lens_cpu,
|
||||
)
|
||||
|
||||
# Replay
|
||||
@@ -477,7 +515,7 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
# DRAFT_EXTEND_V2: all tokens calculations whether accepted or not.
|
||||
unpadding_bs = num_tokens
|
||||
elif bs != raw_bs:
|
||||
forward_batch.spec_info.accept_length = self.accept_length[:raw_bs]
|
||||
forward_batch.spec_info.accept_length = buffers.accept_length[:raw_bs]
|
||||
unpadding_bs = raw_bs
|
||||
else:
|
||||
unpadding_bs = None
|
||||
|
||||
@@ -17,7 +17,8 @@ from __future__ import annotations
|
||||
import bisect
|
||||
import logging
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Callable, List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
@@ -39,6 +40,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardBatch,
|
||||
ForwardMode,
|
||||
)
|
||||
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
||||
from sglang.srt.speculative.eagle_info import EagleDraftInput
|
||||
from sglang.srt.speculative.multi_layer_eagle_utils import assign_new_state_triton
|
||||
from sglang.srt.speculative.spec_utils import fast_topk
|
||||
@@ -59,6 +61,28 @@ if TYPE_CHECKING:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class MultiLayerEagleDraftExtendInputBuffers(ForwardInputBuffers):
|
||||
# Sliced from shared parent buffers
|
||||
input_ids: torch.Tensor
|
||||
out_cache_loc: torch.Tensor
|
||||
swa_out_cache_loc: torch.Tensor
|
||||
positions: torch.Tensor
|
||||
# Shared from parent
|
||||
seq_lens: torch.Tensor
|
||||
seq_lens_cpu: torch.Tensor
|
||||
req_pool_indices: torch.Tensor
|
||||
accept_length: torch.Tensor
|
||||
# Per-step buffers
|
||||
extend_seq_lens: torch.Tensor
|
||||
extend_start_loc: torch.Tensor
|
||||
mrope_positions: torch.Tensor
|
||||
hidden_states: torch.Tensor
|
||||
next_token_logits_buffer: torch.Tensor
|
||||
global_num_tokens_gpu: Optional[torch.Tensor]
|
||||
global_num_tokens_for_logprob_gpu: Optional[torch.Tensor]
|
||||
|
||||
|
||||
class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
def __init__(self, eagle_worker: MultiLayerEagleDraftWorker, step: int):
|
||||
# Parse args
|
||||
@@ -109,7 +133,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
next_cuda_graph_runner,
|
||||
):
|
||||
self.next_cuda_graph_runner = next_cuda_graph_runner
|
||||
self.seq_lens_cpu = cuda_graph_buffers["seq_lens_cpu"]
|
||||
seq_lens_cpu = cuda_graph_buffers["seq_lens_cpu"]
|
||||
self.extend_seq_lens_cpu = [self.num_tokens_per_bs] * self.max_bs
|
||||
|
||||
if self.enable_torch_compile:
|
||||
@@ -119,62 +143,60 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
with torch.device(self.model_runner.device):
|
||||
# sliced buffers
|
||||
# slice according to max_num_token
|
||||
self.input_ids = cuda_graph_buffers["input_ids"][
|
||||
input_ids = cuda_graph_buffers["input_ids"][
|
||||
offset : offset + self.max_num_token
|
||||
]
|
||||
self.out_cache_loc = cuda_graph_buffers["out_cache_loc"][
|
||||
out_cache_loc = cuda_graph_buffers["out_cache_loc"][
|
||||
offset : offset + self.max_num_token
|
||||
]
|
||||
self.swa_out_cache_loc = cuda_graph_buffers["swa_out_cache_loc"][
|
||||
swa_out_cache_loc = cuda_graph_buffers["swa_out_cache_loc"][
|
||||
offset : offset + self.max_num_token
|
||||
]
|
||||
self.positions = cuda_graph_buffers["positions"][
|
||||
positions = cuda_graph_buffers["positions"][
|
||||
offset : offset + self.max_num_token
|
||||
]
|
||||
|
||||
# shared states
|
||||
self.seq_lens = cuda_graph_buffers["seq_lens"]
|
||||
self.req_pool_indices = cuda_graph_buffers["req_pool_indices"]
|
||||
self.accept_length = cuda_graph_buffers["accept_length"]
|
||||
seq_lens = cuda_graph_buffers["seq_lens"]
|
||||
req_pool_indices = cuda_graph_buffers["req_pool_indices"]
|
||||
accept_length = cuda_graph_buffers["accept_length"]
|
||||
|
||||
self.extend_seq_lens = torch.full(
|
||||
extend_seq_lens = torch.full(
|
||||
(self.max_bs,),
|
||||
self.num_tokens_per_bs,
|
||||
dtype=torch.int32,
|
||||
)
|
||||
self.extend_start_loc = torch.arange(
|
||||
extend_start_loc = torch.arange(
|
||||
0,
|
||||
self.max_bs * self.num_tokens_per_bs,
|
||||
step=self.num_tokens_per_bs,
|
||||
dtype=torch.int32,
|
||||
)
|
||||
|
||||
self.mrope_positions = torch.zeros(
|
||||
(3, self.max_num_token), dtype=torch.int64
|
||||
)
|
||||
mrope_positions = torch.zeros((3, self.max_num_token), dtype=torch.int64)
|
||||
|
||||
self.hidden_states = torch.zeros(
|
||||
hidden_states = torch.zeros(
|
||||
(self.max_num_token, self.model_runner.model_config.hidden_size),
|
||||
dtype=self.model_runner.dtype,
|
||||
)
|
||||
|
||||
if self.require_gathered_buffer:
|
||||
if self.require_mlp_tp_gather:
|
||||
self.global_num_tokens_gpu = torch.zeros(
|
||||
global_num_tokens_gpu = torch.zeros(
|
||||
(self.dp_size,), dtype=torch.int32
|
||||
)
|
||||
self.global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||
global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||
(self.dp_size,), dtype=torch.int32
|
||||
)
|
||||
else:
|
||||
assert self.require_attn_tp_gather
|
||||
self.global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
|
||||
self.global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||
global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
|
||||
global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||
(1,), dtype=torch.int32
|
||||
)
|
||||
else:
|
||||
self.global_num_tokens_gpu = None
|
||||
self.global_num_tokens_for_logprob_gpu = None
|
||||
global_num_tokens_gpu = None
|
||||
global_num_tokens_for_logprob_gpu = None
|
||||
|
||||
if hasattr(
|
||||
self.model_runner.model_config.hf_config, "draft_vocab_size"
|
||||
@@ -187,7 +209,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
else:
|
||||
vocab_size = self.model_runner.model_config.vocab_size
|
||||
|
||||
self.next_token_logits_buffer = torch.zeros(
|
||||
next_token_logits_buffer = torch.zeros(
|
||||
(
|
||||
(
|
||||
self.max_bs * self.num_tokens_per_bs
|
||||
@@ -199,6 +221,25 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
dtype=torch.float,
|
||||
)
|
||||
|
||||
self.buffers = MultiLayerEagleDraftExtendInputBuffers(
|
||||
input_ids=input_ids,
|
||||
out_cache_loc=out_cache_loc,
|
||||
swa_out_cache_loc=swa_out_cache_loc,
|
||||
positions=positions,
|
||||
seq_lens=seq_lens,
|
||||
seq_lens_cpu=seq_lens_cpu,
|
||||
req_pool_indices=req_pool_indices,
|
||||
accept_length=accept_length,
|
||||
extend_seq_lens=extend_seq_lens,
|
||||
extend_start_loc=extend_start_loc,
|
||||
mrope_positions=mrope_positions,
|
||||
hidden_states=hidden_states,
|
||||
next_token_logits_buffer=next_token_logits_buffer,
|
||||
global_num_tokens_gpu=global_num_tokens_gpu,
|
||||
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
|
||||
)
|
||||
self.buffers.share_buffers()
|
||||
|
||||
# Capture
|
||||
try:
|
||||
with model_capture_mode():
|
||||
@@ -250,54 +291,55 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
CudaGraphRunner.capture(self)
|
||||
|
||||
def get_forward_batch(self, bs: int) -> ForwardBatch:
|
||||
buffers = self.buffers
|
||||
num_tokens = bs * self.num_tokens_per_bs
|
||||
|
||||
# Graph inputs
|
||||
input_ids = self.input_ids[:num_tokens]
|
||||
req_pool_indices = self.req_pool_indices[:bs]
|
||||
seq_lens = self.seq_lens[:bs]
|
||||
seq_lens_cpu = self.seq_lens_cpu[:bs]
|
||||
extend_seq_lens = self.extend_seq_lens[:bs]
|
||||
input_ids = buffers.input_ids[:num_tokens]
|
||||
req_pool_indices = buffers.req_pool_indices[:bs]
|
||||
seq_lens = buffers.seq_lens[:bs]
|
||||
seq_lens_cpu = buffers.seq_lens_cpu[:bs]
|
||||
extend_seq_lens = buffers.extend_seq_lens[:bs]
|
||||
extend_seq_lens_cpu = self.extend_seq_lens_cpu[:bs]
|
||||
extend_start_loc = self.extend_start_loc[:bs]
|
||||
accept_length = self.accept_length[:bs]
|
||||
out_cache_loc = self.out_cache_loc[:num_tokens]
|
||||
positions = self.positions[:num_tokens]
|
||||
mrope_positions = self.mrope_positions[:, :num_tokens]
|
||||
hidden_states = self.hidden_states[:num_tokens]
|
||||
next_token_logits_buffer = self.next_token_logits_buffer[
|
||||
extend_start_loc = buffers.extend_start_loc[:bs]
|
||||
accept_length = buffers.accept_length[:bs]
|
||||
out_cache_loc = buffers.out_cache_loc[:num_tokens]
|
||||
positions = buffers.positions[:num_tokens]
|
||||
mrope_positions = buffers.mrope_positions[:, :num_tokens]
|
||||
hidden_states = buffers.hidden_states[:num_tokens]
|
||||
next_token_logits_buffer = buffers.next_token_logits_buffer[
|
||||
: bs if self.forward_mode == ForwardMode.DRAFT_EXTEND else num_tokens
|
||||
]
|
||||
|
||||
if self.require_mlp_tp_gather:
|
||||
self.global_num_tokens_gpu.copy_(
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens] * self.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=self.input_ids.device,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
self.global_num_tokens_for_logprob_gpu.copy_(
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens] * self.dp_size,
|
||||
dtype=torch.int32,
|
||||
device=self.input_ids.device,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
global_dp_buffer_len = num_tokens * self.dp_size
|
||||
elif self.require_attn_tp_gather:
|
||||
self.global_num_tokens_gpu.copy_(
|
||||
buffers.global_num_tokens_gpu.copy_(
|
||||
torch.tensor(
|
||||
[num_tokens],
|
||||
dtype=torch.int32,
|
||||
device=self.input_ids.device,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
self.global_num_tokens_for_logprob_gpu.copy_(
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(
|
||||
torch.tensor(
|
||||
[bs],
|
||||
dtype=torch.int32,
|
||||
device=self.input_ids.device,
|
||||
device=buffers.input_ids.device,
|
||||
)
|
||||
)
|
||||
global_dp_buffer_len = num_tokens
|
||||
@@ -326,8 +368,8 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
return_logprob=False,
|
||||
positions=positions,
|
||||
mrope_positions=mrope_positions,
|
||||
global_num_tokens_gpu=self.global_num_tokens_gpu,
|
||||
global_num_tokens_for_logprob_gpu=self.global_num_tokens_for_logprob_gpu,
|
||||
global_num_tokens_gpu=buffers.global_num_tokens_gpu,
|
||||
global_num_tokens_for_logprob_gpu=buffers.global_num_tokens_for_logprob_gpu,
|
||||
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
|
||||
global_dp_buffer_len=global_dp_buffer_len,
|
||||
spec_algorithm=self.model_runner.spec_algorithm,
|
||||
@@ -346,6 +388,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
return forward_batch
|
||||
|
||||
def capture_one_batch_size(self, bs: int, forward: Callable, stream_idx: int = 0):
|
||||
buffers = self.buffers
|
||||
graph = self._create_graph()
|
||||
stream = self.stream
|
||||
|
||||
@@ -390,7 +433,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
select_index = (
|
||||
torch.arange(bs, device=self.model_runner.device)
|
||||
* (self.speculative_num_draft_tokens + self.step)
|
||||
+ self.accept_length[:bs]
|
||||
+ buffers.accept_length[:bs]
|
||||
- 1
|
||||
+ self.step
|
||||
)
|
||||
@@ -399,24 +442,25 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
ret.topk_p, ret.topk_index = fast_topk(probs, self.topk, dim=-1)
|
||||
|
||||
if self.next_cuda_graph_runner is not None:
|
||||
next_buffers = self.next_cuda_graph_runner.buffers
|
||||
padding_lens = (
|
||||
self.speculative_num_draft_tokens - self.accept_length[:bs]
|
||||
self.speculative_num_draft_tokens - buffers.accept_length[:bs]
|
||||
)
|
||||
assign_new_state_triton(
|
||||
ret.topk_index,
|
||||
self.input_ids,
|
||||
self.positions,
|
||||
self.hidden_states,
|
||||
self.out_cache_loc,
|
||||
self.extend_seq_lens,
|
||||
self.extend_start_loc,
|
||||
self.next_cuda_graph_runner.input_ids,
|
||||
self.next_cuda_graph_runner.positions,
|
||||
self.next_cuda_graph_runner.hidden_states,
|
||||
self.next_cuda_graph_runner.out_cache_loc,
|
||||
self.next_cuda_graph_runner.extend_seq_lens,
|
||||
self.next_cuda_graph_runner.extend_start_loc,
|
||||
self.next_cuda_graph_runner.seq_lens,
|
||||
buffers.input_ids,
|
||||
buffers.positions,
|
||||
buffers.hidden_states,
|
||||
buffers.out_cache_loc,
|
||||
buffers.extend_seq_lens,
|
||||
buffers.extend_start_loc,
|
||||
next_buffers.input_ids,
|
||||
next_buffers.positions,
|
||||
next_buffers.hidden_states,
|
||||
next_buffers.out_cache_loc,
|
||||
next_buffers.extend_seq_lens,
|
||||
next_buffers.extend_start_loc,
|
||||
next_buffers.seq_lens,
|
||||
padding_lens,
|
||||
forward_batch.batch_size,
|
||||
self.step,
|
||||
@@ -424,9 +468,9 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
forward_batch.req_to_token_pool.req_to_token,
|
||||
self.eagle_worker.req_to_hidden_states_pool,
|
||||
)
|
||||
self.next_cuda_graph_runner.swa_out_cache_loc.copy_(
|
||||
next_buffers.swa_out_cache_loc.copy_(
|
||||
self.model_runner.token_to_kv_pool.translate_loc_from_full_to_swa(
|
||||
self.next_cuda_graph_runner.out_cache_loc
|
||||
next_buffers.out_cache_loc
|
||||
)
|
||||
)
|
||||
|
||||
@@ -446,27 +490,30 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
def init_replay_state(
|
||||
self, forward_batch: ForwardBatch, bs: int, raw_bs: int, num_tokens: int
|
||||
):
|
||||
buffers = self.buffers
|
||||
# Common inputs
|
||||
self.input_ids[:num_tokens].copy_(forward_batch.input_ids)
|
||||
self.seq_lens[:raw_bs].copy_(forward_batch.seq_lens)
|
||||
buffers.input_ids[:num_tokens].copy_(forward_batch.input_ids)
|
||||
buffers.seq_lens[:raw_bs].copy_(forward_batch.seq_lens)
|
||||
if forward_batch.extend_seq_lens is not None:
|
||||
self.extend_seq_lens[:raw_bs].copy_(forward_batch.extend_seq_lens)
|
||||
self.extend_start_loc[:raw_bs].copy_(forward_batch.extend_start_loc)
|
||||
self.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc)
|
||||
self.positions[:num_tokens].copy_(forward_batch.positions)
|
||||
buffers.extend_seq_lens[:raw_bs].copy_(forward_batch.extend_seq_lens)
|
||||
buffers.extend_start_loc[:raw_bs].copy_(forward_batch.extend_start_loc)
|
||||
buffers.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc)
|
||||
buffers.positions[:num_tokens].copy_(forward_batch.positions)
|
||||
if (
|
||||
forward_batch.spec_info.hidden_states.shape[1]
|
||||
== self.hidden_states.shape[1]
|
||||
== buffers.hidden_states.shape[1]
|
||||
):
|
||||
self.hidden_states[:num_tokens].copy_(forward_batch.spec_info.hidden_states)
|
||||
buffers.hidden_states[:num_tokens].copy_(
|
||||
forward_batch.spec_info.hidden_states
|
||||
)
|
||||
if forward_batch.spec_info.accept_length is not None:
|
||||
self.accept_length[:raw_bs].copy_(forward_batch.spec_info.accept_length)
|
||||
self.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices)
|
||||
buffers.accept_length[:raw_bs].copy_(forward_batch.spec_info.accept_length)
|
||||
buffers.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices)
|
||||
|
||||
if forward_batch.seq_lens_cpu is not None:
|
||||
if bs != raw_bs:
|
||||
self.seq_lens_cpu.fill_(self.seq_len_fill_value)
|
||||
self.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
|
||||
buffers.seq_lens_cpu.fill_(self.seq_len_fill_value)
|
||||
buffers.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
|
||||
|
||||
if forward_batch.extend_seq_lens_cpu is not None:
|
||||
self.extend_seq_lens_cpu[:raw_bs] = forward_batch.extend_seq_lens_cpu
|
||||
@@ -474,6 +521,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
def replay(self, forward_batch: ForwardBatch, init_state: bool = True):
|
||||
assert forward_batch.out_cache_loc is not None
|
||||
self.deepep_adapter.replay()
|
||||
buffers = self.buffers
|
||||
|
||||
# batch_size and num_seqs can be different in case there are finished examples
|
||||
# in the batch, which will not be counted as num_seqs
|
||||
@@ -492,28 +540,28 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
self.init_replay_state(forward_batch, bs, raw_bs, num_tokens)
|
||||
|
||||
if self.require_gathered_buffer:
|
||||
self.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs)
|
||||
self.global_num_tokens_for_logprob_gpu.fill_(bs * self.num_tokens_per_bs)
|
||||
buffers.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs)
|
||||
buffers.global_num_tokens_for_logprob_gpu.fill_(bs * self.num_tokens_per_bs)
|
||||
|
||||
forward_batch.spec_info.hidden_states = self.hidden_states[:num_tokens]
|
||||
forward_batch.spec_info.accept_length = self.accept_length[:bs]
|
||||
forward_batch.spec_info.hidden_states = buffers.hidden_states[:num_tokens]
|
||||
forward_batch.spec_info.accept_length = buffers.accept_length[:bs]
|
||||
forward_batch.spec_info.num_tokens_per_req = self.num_tokens_per_bs
|
||||
forward_batch.spec_info.num_tokens_for_logprob_per_req = 1
|
||||
forward_batch.spec_info.positions = self.positions[:num_tokens]
|
||||
forward_batch.spec_info.extend_seq_lens_tensor = self.extend_seq_lens[:bs]
|
||||
forward_batch.spec_info.positions = buffers.positions[:num_tokens]
|
||||
forward_batch.spec_info.extend_seq_lens_tensor = buffers.extend_seq_lens[:bs]
|
||||
|
||||
self.eagle_worker.draft_extend_attn_backend_list[
|
||||
self.step
|
||||
].init_forward_metadata_replay_cuda_graph(
|
||||
bs=bs,
|
||||
req_pool_indices=self.req_pool_indices,
|
||||
seq_lens=self.seq_lens,
|
||||
req_pool_indices=buffers.req_pool_indices,
|
||||
seq_lens=buffers.seq_lens,
|
||||
seq_lens_sum=forward_batch.seq_lens_sum
|
||||
+ (bs - raw_bs) * self.seq_len_fill_value,
|
||||
encoder_lens=None,
|
||||
forward_mode=self.forward_mode,
|
||||
spec_info=forward_batch.spec_info,
|
||||
seq_lens_cpu=self.seq_lens_cpu,
|
||||
seq_lens_cpu=buffers.seq_lens_cpu,
|
||||
)
|
||||
|
||||
# Replay
|
||||
@@ -526,7 +574,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
# DRAFT_EXTEND_V2: all tokens calculations whether accepted or not.
|
||||
unpadding_bs = num_tokens
|
||||
elif bs != raw_bs:
|
||||
forward_batch.spec_info.accept_length = self.accept_length[:raw_bs]
|
||||
forward_batch.spec_info.accept_length = buffers.accept_length[:raw_bs]
|
||||
unpadding_bs = raw_bs
|
||||
else:
|
||||
unpadding_bs = None
|
||||
@@ -565,8 +613,8 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
|
||||
self.runners = [None] * self.speculative_num_steps
|
||||
return
|
||||
|
||||
self.runners = []
|
||||
buffer_len_list = []
|
||||
self.runners: List[Optional[MultiLayerEagleDraftExtendCudaGraphRunner]] = []
|
||||
buffer_len_list: List[int] = []
|
||||
|
||||
# 1. Capture loop
|
||||
for step in range(self.speculative_num_steps):
|
||||
|
||||
@@ -498,13 +498,13 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
|
||||
self.cuda_graph_runner_for_draft_extend.get_last_runner()
|
||||
)
|
||||
assign_hidden_states_pool_triton(
|
||||
last_cuda_graph_runner.hidden_states,
|
||||
last_cuda_graph_runner.req_pool_indices,
|
||||
last_cuda_graph_runner.buffers.hidden_states,
|
||||
last_cuda_graph_runner.buffers.req_pool_indices,
|
||||
self.req_to_hidden_states_pool,
|
||||
self.speculative_num_steps - 1,
|
||||
forward_batch.batch_size,
|
||||
last_cuda_graph_runner.extend_seq_lens,
|
||||
last_cuda_graph_runner.extend_start_loc,
|
||||
last_cuda_graph_runner.buffers.extend_seq_lens,
|
||||
last_cuda_graph_runner.buffers.extend_start_loc,
|
||||
)
|
||||
|
||||
# Reorganize the spec info for the next batch
|
||||
|
||||
Reference in New Issue
Block a user