Fix DSpark and DP/EP (#33098)
Signed-off-by: Vladislav Nosivskoy <vladnosiv@gmail.com> Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com>
This commit is contained in:
co-authored by
Xinyuan Tong
parent
bfa4e4a57b
commit
154f0ac662
@@ -16,6 +16,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
CaptureHiddenMode,
|
||||
ForwardBatch,
|
||||
ForwardMode,
|
||||
enable_num_token_non_padded,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
|
||||
@@ -384,6 +385,13 @@ class DraftBlockProposer:
|
||||
batch.global_num_tokens_for_logprob,
|
||||
)
|
||||
device = self.draft_model_runner.device
|
||||
forward_batch.original_global_num_tokens_cpu = batch.global_num_tokens
|
||||
num_tokens = forward_batch.input_ids.numel()
|
||||
if enable_num_token_non_padded():
|
||||
forward_batch.num_token_non_padded = torch.tensor(
|
||||
num_tokens, dtype=torch.int32, device=device
|
||||
)
|
||||
forward_batch.num_token_non_padded_cpu = num_tokens
|
||||
forward_batch.global_num_tokens_cpu = gnt
|
||||
forward_batch.global_num_tokens_for_logprob_cpu = gnt_logprob
|
||||
forward_batch.global_num_tokens_gpu = torch.tensor(gnt, dtype=torch.int64).to(
|
||||
|
||||
Reference in New Issue
Block a user