Sync the changes on cuda graph runners (#6932)
This commit is contained in:
@@ -4,7 +4,7 @@ from typing import List
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.utils import is_cuda, is_hip
|
||||
from sglang.srt.utils import is_cuda, is_hip, rank0_print
|
||||
|
||||
if is_cuda() or is_hip():
|
||||
from sgl_kernel import (
|
||||
@@ -344,13 +344,13 @@ def test_build_tree_kernel_efficient():
|
||||
num_verify_tokens=num_draft_token,
|
||||
)
|
||||
|
||||
first_rank_print("=========== build tree kernel efficient ==========")
|
||||
# first_rank_print(f"{tree_mask=}", flush=True)
|
||||
first_rank_print(f"{position=}", flush=True)
|
||||
first_rank_print(f"{retrive_index=}", flush=True)
|
||||
first_rank_print(f"{retrive_next_token=}", flush=True)
|
||||
first_rank_print(f"{retrive_next_sibling=}", flush=True)
|
||||
first_rank_print(f"{draft_tokens=}", flush=True)
|
||||
rank0_print("=========== build tree kernel efficient ==========")
|
||||
# rank0_print(f"{tree_mask=}", flush=True)
|
||||
rank0_print(f"{position=}", flush=True)
|
||||
rank0_print(f"{retrive_index=}", flush=True)
|
||||
rank0_print(f"{retrive_next_token=}", flush=True)
|
||||
rank0_print(f"{retrive_next_sibling=}", flush=True)
|
||||
rank0_print(f"{draft_tokens=}", flush=True)
|
||||
assert position.tolist() == [5, 6, 6, 7, 7, 8, 8, 9, 10, 11, 12, 12, 12, 12, 13, 14]
|
||||
assert retrive_index.tolist() == [
|
||||
[0, 1, 2, 3, 4, 5, 6, 7],
|
||||
|
||||
@@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Callable
|
||||
import torch
|
||||
|
||||
from sglang.srt.model_executor.cuda_graph_runner import (
|
||||
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
||||
CudaGraphRunner,
|
||||
get_batch_sizes_to_capture,
|
||||
get_global_graph_memory_pool,
|
||||
@@ -73,7 +74,7 @@ class EAGLEDraftCudaGraphRunner:
|
||||
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(
|
||||
(self.max_bs, self.model_runner.model_config.hidden_size),
|
||||
(self.max_num_token, self.model_runner.model_config.hidden_size),
|
||||
dtype=self.model_runner.dtype,
|
||||
)
|
||||
|
||||
@@ -82,13 +83,7 @@ class EAGLEDraftCudaGraphRunner:
|
||||
self.capture()
|
||||
except RuntimeError as e:
|
||||
raise Exception(
|
||||
f"Capture CUDA graph failed: {e}\n"
|
||||
"Possible solutions:\n"
|
||||
"1. set --mem-fraction-static to a smaller value (e.g., 0.8 or 0.7)\n"
|
||||
"2. set --cuda-graph-max-bs to a smaller value (e.g., 16)\n"
|
||||
"3. disable torch compile by not using --enable-torch-compile\n"
|
||||
"4. disable CUDA graph by --disable-cuda-graph. (Not recommended. Huge performance loss)\n"
|
||||
"Open an issue on GitHub https://github.com/sgl-project/sglang/issues/new/choose \n"
|
||||
f"Capture cuda graph failed: {e}\n{CUDA_GRAPH_CAPTURE_FAILED_MSG}"
|
||||
)
|
||||
|
||||
def can_run(self, forward_batch: ForwardBatch):
|
||||
|
||||
@@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Callable
|
||||
import torch
|
||||
|
||||
from sglang.srt.model_executor.cuda_graph_runner import (
|
||||
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
||||
CudaGraphRunner,
|
||||
LogitsProcessorOutput,
|
||||
get_batch_sizes_to_capture,
|
||||
@@ -89,13 +90,7 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
self.capture()
|
||||
except RuntimeError as e:
|
||||
raise Exception(
|
||||
f"Capture CUDA graph failed: {e}\n"
|
||||
"Possible solutions:\n"
|
||||
"1. set --mem-fraction-static to a smaller value (e.g., 0.8 or 0.7)\n"
|
||||
"2. set --cuda-graph-max-bs to a smaller value (e.g., 16)\n"
|
||||
"3. disable torch compile by not using --enable-torch-compile\n"
|
||||
"4. disable CUDA graph by --disable-cuda-graph. (Not recommended. Huge performance loss)\n"
|
||||
"Open an issue on GitHub https://github.com/sgl-project/sglang/issues/new/choose \n"
|
||||
f"Capture cuda graph failed: {e}\n{CUDA_GRAPH_CAPTURE_FAILED_MSG}"
|
||||
)
|
||||
|
||||
def can_run(self, forward_batch: ForwardBatch):
|
||||
@@ -200,7 +195,6 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
# in the batch, which will not be counted as num_seqs
|
||||
raw_bs = forward_batch.batch_size
|
||||
num_tokens = forward_batch.input_ids.shape[0]
|
||||
assert raw_bs * self.num_tokens_per_bs == num_tokens
|
||||
|
||||
index = bisect.bisect_left(self.capture_bs, raw_bs)
|
||||
bs = self.capture_bs[index]
|
||||
@@ -224,9 +218,9 @@ class EAGLEDraftExtendCudaGraphRunner:
|
||||
self.seq_lens_cpu.fill_(1)
|
||||
self.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
|
||||
|
||||
forward_batch.spec_info.positions = None
|
||||
if bs != raw_bs:
|
||||
forward_batch.spec_info.accept_length = self.accept_length[:bs]
|
||||
forward_batch.spec_info.positions = None
|
||||
|
||||
self.eagle_worker.draft_extend_attn_backend.init_forward_metadata_replay_cuda_graph(
|
||||
bs=bs,
|
||||
|
||||
@@ -232,8 +232,9 @@ class EagleVerifyInput:
|
||||
retrive_next_token: torch.Tensor
|
||||
retrive_next_sibling: torch.Tensor
|
||||
retrive_cum_len: torch.Tensor
|
||||
draft_token_num: int
|
||||
spec_steps: int
|
||||
topk: int
|
||||
draft_token_num: int
|
||||
capture_hidden_mode: CaptureHiddenMode
|
||||
grammar: BaseGrammarObject = None
|
||||
|
||||
@@ -270,16 +271,17 @@ class EagleVerifyInput:
|
||||
)
|
||||
|
||||
return cls(
|
||||
draft_tokens,
|
||||
tree_mask,
|
||||
position,
|
||||
retrive_index,
|
||||
retrive_next_token,
|
||||
retrive_next_sibling,
|
||||
None,
|
||||
num_verify_tokens,
|
||||
spec_steps,
|
||||
CaptureHiddenMode.FULL,
|
||||
draft_token=draft_tokens,
|
||||
custom_mask=tree_mask,
|
||||
positions=position,
|
||||
retrive_index=retrive_index,
|
||||
retrive_next_token=retrive_next_token,
|
||||
retrive_next_sibling=retrive_next_sibling,
|
||||
retrive_cum_len=None,
|
||||
spec_steps=spec_steps,
|
||||
topk=topk,
|
||||
draft_token_num=num_verify_tokens,
|
||||
capture_hidden_mode=CaptureHiddenMode.FULL,
|
||||
)
|
||||
|
||||
def prepare_for_verify(self, batch: ScheduleBatch, page_size: int):
|
||||
|
||||
Reference in New Issue
Block a user