Stabilize CP shared-KV batch padding semantics
CP shared-KV bs>1 exposed three distinct padding domains: valid cache rows, CP page-tail compute rows, and MLP-sync flattened static padding. The previous implementation mixed these domains in direct-write and index top-k paths, so real requests failed when q/out_cache_loc lengths matched valid rows while metadata aliases described compute rows.\n\nThis change makes compute split strip only proven flattened static padding, keeps valid cache writes strict except for extend_num_tokens-proven static tails, marks CP-local EAGLE draft hidden state explicitly, and selects NSA index top-k query metadata by the actual q/weight row count.\n\nConstraint: CP shared-KV cache writes must never persist dummy page-tail or MLP static padding rows.\nConstraint: EAGLE draft hidden state can be CP-local before full CP metadata is visible in prepare_mlp_sync_batch.\nRejected: Use compute_padding_enabled as direct-write truncation proof | it silently accepts unknown out_cache_loc tails.\nRejected: Always consume compute q metadata in index top-k | actual q/weights can be valid-only after CP split.\nConfidence: medium\nScope-risk: moderate\nDirective: Do not collapse valid rows, CP compute padding, and MLP static padding into one length condition; use explicit provenance.\nTested: remote py_compile for touched NSA files\nTested: remote targeted CP shared-KV padding/top-k regressions\nTested: remote pytest test_nsa_cp_utils.py test_cp_shared_kv_layout.py test_cp_shared_kv_runtime.py -k 'not test_tai_current_slot_fill_sparse_page_self_test_passes_on_installed_kernel' => 228 passed, 1 deselected, 5 warnings, 2 subtests passed\nNot-tested: full ETE replay after the final index top-k fix\nNot-tested: TAI current-index fast path dtype fallback
This commit is contained in:
@@ -3,6 +3,7 @@ from __future__ import annotations
|
||||
import contextlib
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
@@ -120,6 +121,120 @@ def _compute_contiguous_valid_cp_query_count(
|
||||
return max(0, min(actual_seq_q, valid_count))
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BatchTopKQueryLengths:
|
||||
request_seq_q_prev: List[int]
|
||||
request_seq_q_next: List[int]
|
||||
request_valid_seq_q_prev: List[int]
|
||||
request_valid_seq_q_next: List[int]
|
||||
uses_compute_query_rows: bool
|
||||
valid_token_count: int
|
||||
compute_token_count: int
|
||||
|
||||
|
||||
def _select_batch_topk_query_lengths(
|
||||
*,
|
||||
cp_metadata,
|
||||
batch_plan,
|
||||
batch_size: int,
|
||||
q_tokens: int,
|
||||
weights_tokens: int,
|
||||
layer_id: Optional[int] = None,
|
||||
) -> BatchTopKQueryLengths:
|
||||
"""Select CP top-k segment lengths that match actual q/weight rows."""
|
||||
|
||||
def metadata_list(name: str, fallback_name: Optional[str] = None) -> List[int]:
|
||||
values = getattr(cp_metadata, name, None)
|
||||
if values is None and batch_plan is not None:
|
||||
values = getattr(batch_plan, name, None)
|
||||
if values is None and fallback_name is not None:
|
||||
values = getattr(cp_metadata, fallback_name, None)
|
||||
if values is None and batch_plan is not None:
|
||||
values = getattr(batch_plan, fallback_name, None)
|
||||
return [int(x) for x in (values or [])]
|
||||
|
||||
request_compute_seq_q_prev = metadata_list(
|
||||
"request_compute_seq_q_prev",
|
||||
fallback_name="request_actual_seq_q_prev",
|
||||
)
|
||||
request_compute_seq_q_next = metadata_list(
|
||||
"request_compute_seq_q_next",
|
||||
fallback_name="request_actual_seq_q_next",
|
||||
)
|
||||
request_valid_seq_q_prev = metadata_list(
|
||||
"request_valid_seq_q_prev",
|
||||
fallback_name="request_valid_actual_seq_q_prev",
|
||||
)
|
||||
request_valid_seq_q_next = metadata_list(
|
||||
"request_valid_seq_q_next",
|
||||
fallback_name="request_valid_actual_seq_q_next",
|
||||
)
|
||||
if not request_valid_seq_q_prev:
|
||||
request_valid_seq_q_prev = metadata_list("request_actual_seq_q_prev")
|
||||
if not request_valid_seq_q_next:
|
||||
request_valid_seq_q_next = metadata_list("request_actual_seq_q_next")
|
||||
|
||||
if not (
|
||||
len(request_compute_seq_q_prev) == batch_size
|
||||
and len(request_compute_seq_q_next) == batch_size
|
||||
and len(request_valid_seq_q_prev) == batch_size
|
||||
and len(request_valid_seq_q_next) == batch_size
|
||||
):
|
||||
raise RuntimeError(
|
||||
"[CP_SHARED_KV_FAIL_FAST][index_topk] "
|
||||
"reason=batch_gt1_index_query_metadata_incomplete "
|
||||
f"batch_size={batch_size} layer_id={layer_id} "
|
||||
f"compute_q_prev={request_compute_seq_q_prev} "
|
||||
f"compute_q_next={request_compute_seq_q_next} "
|
||||
f"valid_q_prev={request_valid_seq_q_prev} "
|
||||
f"valid_q_next={request_valid_seq_q_next}"
|
||||
)
|
||||
|
||||
valid_token_count = sum(request_valid_seq_q_prev) + sum(request_valid_seq_q_next)
|
||||
compute_token_count = sum(request_compute_seq_q_prev) + sum(
|
||||
request_compute_seq_q_next
|
||||
)
|
||||
q_tokens = int(q_tokens)
|
||||
weights_tokens = int(weights_tokens)
|
||||
if q_tokens != weights_tokens:
|
||||
raise RuntimeError(
|
||||
"[CP_SHARED_KV_FAIL_FAST][index_topk] "
|
||||
"reason=batch_gt1_index_q_weight_length_mismatch "
|
||||
f"batch_size={batch_size} layer_id={layer_id} q_tokens={q_tokens} "
|
||||
f"weights_tokens={weights_tokens} valid_tokens={valid_token_count} "
|
||||
f"compute_tokens={compute_token_count}"
|
||||
)
|
||||
|
||||
if q_tokens == valid_token_count:
|
||||
return BatchTopKQueryLengths(
|
||||
request_seq_q_prev=request_valid_seq_q_prev,
|
||||
request_seq_q_next=request_valid_seq_q_next,
|
||||
request_valid_seq_q_prev=request_valid_seq_q_prev,
|
||||
request_valid_seq_q_next=request_valid_seq_q_next,
|
||||
uses_compute_query_rows=False,
|
||||
valid_token_count=valid_token_count,
|
||||
compute_token_count=compute_token_count,
|
||||
)
|
||||
if q_tokens == compute_token_count:
|
||||
return BatchTopKQueryLengths(
|
||||
request_seq_q_prev=request_compute_seq_q_prev,
|
||||
request_seq_q_next=request_compute_seq_q_next,
|
||||
request_valid_seq_q_prev=request_valid_seq_q_prev,
|
||||
request_valid_seq_q_next=request_valid_seq_q_next,
|
||||
uses_compute_query_rows=True,
|
||||
valid_token_count=valid_token_count,
|
||||
compute_token_count=compute_token_count,
|
||||
)
|
||||
|
||||
raise RuntimeError(
|
||||
"[CP_SHARED_KV_FAIL_FAST][index_topk] "
|
||||
"reason=batch_gt1_index_q_length_mismatch "
|
||||
f"batch_size={batch_size} layer_id={layer_id} q_tokens={q_tokens} "
|
||||
f"weights_tokens={weights_tokens} valid_tokens={valid_token_count} "
|
||||
f"compute_tokens={compute_token_count}"
|
||||
)
|
||||
|
||||
|
||||
def _log_cp_shared_kv_index_prefetch_fallback(
|
||||
reason: str,
|
||||
message: str,
|
||||
@@ -1747,20 +1862,6 @@ class Indexer(MultiPlatformOp):
|
||||
assert cp_metadata is not None
|
||||
batch_size = int(getattr(cp_metadata, "batch_size", 1) or 1)
|
||||
batch_plan = get_cp_shared_kv_batch_plan(forward_batch)
|
||||
compute_padding_enabled = bool(
|
||||
getattr(cp_metadata, "compute_padding_enabled", False)
|
||||
or bool(getattr(batch_plan, "compute_padding_enabled", False))
|
||||
)
|
||||
|
||||
def metadata_list(name: str, fallback_name: Optional[str] = None) -> List[int]:
|
||||
values = getattr(cp_metadata, name, None)
|
||||
if values is None and batch_plan is not None:
|
||||
values = getattr(batch_plan, name, None)
|
||||
if values is None and fallback_name is not None:
|
||||
values = getattr(cp_metadata, fallback_name, None)
|
||||
if values is None and batch_plan is not None:
|
||||
values = getattr(batch_plan, fallback_name, None)
|
||||
return list(values or [])
|
||||
|
||||
request_kv_len_prev = list(getattr(cp_metadata, "request_kv_len_prev", []) or [])
|
||||
request_kv_len_next = list(getattr(cp_metadata, "request_kv_len_next", []) or [])
|
||||
@@ -1772,35 +1873,20 @@ class Indexer(MultiPlatformOp):
|
||||
request_kv_len_next = list(
|
||||
getattr(batch_plan, "request_kv_len_next", []) or []
|
||||
)
|
||||
request_actual_seq_q_prev = metadata_list(
|
||||
"request_compute_seq_q_prev"
|
||||
if compute_padding_enabled
|
||||
else "request_actual_seq_q_prev",
|
||||
fallback_name="request_actual_seq_q_prev",
|
||||
|
||||
query_lengths = _select_batch_topk_query_lengths(
|
||||
cp_metadata=cp_metadata,
|
||||
batch_plan=batch_plan,
|
||||
batch_size=batch_size,
|
||||
q_tokens=int(q_fp8.shape[0]),
|
||||
weights_tokens=int(weights.shape[0]),
|
||||
layer_id=layer_id,
|
||||
)
|
||||
request_actual_seq_q_next = metadata_list(
|
||||
"request_compute_seq_q_next"
|
||||
if compute_padding_enabled
|
||||
else "request_actual_seq_q_next",
|
||||
fallback_name="request_actual_seq_q_next",
|
||||
)
|
||||
request_valid_seq_q_prev = metadata_list(
|
||||
"request_valid_seq_q_prev",
|
||||
fallback_name="request_valid_actual_seq_q_prev",
|
||||
)
|
||||
request_valid_seq_q_next = metadata_list(
|
||||
"request_valid_seq_q_next",
|
||||
fallback_name="request_valid_actual_seq_q_next",
|
||||
)
|
||||
if not compute_padding_enabled:
|
||||
# Older bs>1 metadata did not have explicit valid-q aliases because
|
||||
# actual q length was also the valid q length. Keep that path
|
||||
# compatible while compute-padding remains fail-fast if valid
|
||||
# lengths are missing.
|
||||
if not request_valid_seq_q_prev:
|
||||
request_valid_seq_q_prev = request_actual_seq_q_prev
|
||||
if not request_valid_seq_q_next:
|
||||
request_valid_seq_q_next = request_actual_seq_q_next
|
||||
request_actual_seq_q_prev = query_lengths.request_seq_q_prev
|
||||
request_actual_seq_q_next = query_lengths.request_seq_q_next
|
||||
request_valid_seq_q_prev = query_lengths.request_valid_seq_q_prev
|
||||
request_valid_seq_q_next = query_lengths.request_valid_seq_q_next
|
||||
|
||||
if not (
|
||||
len(request_kv_len_prev) == batch_size
|
||||
and len(request_kv_len_next) == batch_size
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from itertools import accumulate
|
||||
from typing import TYPE_CHECKING, List, Tuple, Union
|
||||
from typing import TYPE_CHECKING, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
@@ -106,6 +106,23 @@ def _is_cp_shared_kv_forward_batch(forward_batch: "ForwardBatch") -> bool:
|
||||
return bool(getattr(forward_batch, "uses_cp_shared_kv", False))
|
||||
|
||||
|
||||
def _is_cp_shared_kv_draft_extend(forward_batch: "ForwardBatch") -> bool:
|
||||
"""Return whether EAGLE/NextN draft extend should keep CP-local semantics."""
|
||||
|
||||
if not _is_cp_shared_kv_forward_batch(forward_batch):
|
||||
return False
|
||||
if not envs.SGLANG_CP_DRAFT_SHARED_KV.get():
|
||||
return False
|
||||
forward_mode = getattr(forward_batch, "forward_mode", None)
|
||||
is_draft_extend = getattr(forward_mode, "is_draft_extend", None)
|
||||
if not callable(is_draft_extend):
|
||||
return False
|
||||
try:
|
||||
return bool(is_draft_extend(include_v2=True))
|
||||
except TypeError:
|
||||
return bool(is_draft_extend())
|
||||
|
||||
|
||||
def _fail_if_cp_shared_kv_round_robin(
|
||||
forward_batch: "ForwardBatch",
|
||||
*,
|
||||
@@ -751,12 +768,25 @@ def get_cp_shared_kv_batch_plan(forward_batch: "ForwardBatch"):
|
||||
return None
|
||||
|
||||
|
||||
def _get_forward_batch_static_padded_tokens(
|
||||
forward_batch: "ForwardBatch",
|
||||
) -> Optional[int]:
|
||||
static_padded_tokens = getattr(forward_batch, "extend_num_tokens", None)
|
||||
if static_padded_tokens is None:
|
||||
return None
|
||||
try:
|
||||
return int(static_padded_tokens)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def split_tensor_by_cp_batch_plan(
|
||||
tensor: torch.Tensor,
|
||||
plan,
|
||||
*,
|
||||
mode: str = "data",
|
||||
split_kind: str = "compute",
|
||||
static_padded_tokens: Optional[int] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Split a flattened batch tensor by per-request in-seq CP plan.
|
||||
|
||||
@@ -802,11 +832,42 @@ def split_tensor_by_cp_batch_plan(
|
||||
)
|
||||
|
||||
expected_tokens = sum(int(x) for x in request_extend_lens)
|
||||
if int(tensor.shape[0]) != expected_tokens:
|
||||
raise RuntimeError(
|
||||
"[CP_SHARED_KV_FAIL_FAST][batch_gt1_split_input_len_mismatch] "
|
||||
f"input tokens={int(tensor.shape[0])} expected={expected_tokens}"
|
||||
input_tokens = int(tensor.shape[0])
|
||||
if input_tokens != expected_tokens:
|
||||
max_compute_tokens = (
|
||||
sum(int(x) for x in request_target_lens)
|
||||
if split_kind == "compute" and compute_padding_enabled
|
||||
else expected_tokens
|
||||
)
|
||||
max_static_tokens = (
|
||||
int(static_padded_tokens)
|
||||
if static_padded_tokens is not None
|
||||
else expected_tokens
|
||||
)
|
||||
max_allowed_tokens = max(max_compute_tokens, max_static_tokens)
|
||||
if (
|
||||
split_kind == "compute"
|
||||
and input_tokens > expected_tokens
|
||||
and input_tokens <= max_allowed_tokens
|
||||
and (
|
||||
compute_padding_enabled
|
||||
or (
|
||||
static_padded_tokens is not None
|
||||
and int(static_padded_tokens) > expected_tokens
|
||||
)
|
||||
)
|
||||
):
|
||||
tensor = tensor[:expected_tokens]
|
||||
else:
|
||||
expected_detail = f"expected={expected_tokens}"
|
||||
if split_kind == "compute" and compute_padding_enabled:
|
||||
expected_detail += f" max_compute={max_compute_tokens}"
|
||||
if split_kind == "compute" and static_padded_tokens is not None:
|
||||
expected_detail += f" static_padded={int(static_padded_tokens)}"
|
||||
raise RuntimeError(
|
||||
"[CP_SHARED_KV_FAIL_FAST][batch_gt1_split_input_len_mismatch] "
|
||||
f"input tokens={input_tokens} {expected_detail}"
|
||||
)
|
||||
|
||||
local_chunks = []
|
||||
request_tensors = torch.split(tensor, [int(x) for x in request_extend_lens], dim=0)
|
||||
@@ -1462,11 +1523,15 @@ def can_cp_split(seq_len: int, cp_size: int, use_nsa: bool, forward_batch):
|
||||
min_extend_token_count = 1
|
||||
else:
|
||||
min_extend_token_count = cp_size
|
||||
is_context_parallel_extend = (
|
||||
forward_batch.forward_mode.is_context_parallel_extend()
|
||||
or _is_cp_shared_kv_draft_extend(forward_batch)
|
||||
)
|
||||
if (
|
||||
cur_cp_seq_len != 0
|
||||
and cp_size > 1
|
||||
and use_nsa
|
||||
and forward_batch.forward_mode.is_context_parallel_extend()
|
||||
and is_context_parallel_extend
|
||||
and is_nsa_enable_prefill_cp()
|
||||
and extend_token_count >= min_extend_token_count
|
||||
):
|
||||
@@ -1523,6 +1588,7 @@ def _cp_split_and_rebuild_batch_in_seq(forward_batch, input_: torch.Tensor):
|
||||
input_,
|
||||
get_cp_shared_kv_batch_plan(forward_batch),
|
||||
mode="1d" if input_.dim() == 1 else "data",
|
||||
static_padded_tokens=_get_forward_batch_static_padded_tokens(forward_batch),
|
||||
)
|
||||
|
||||
|
||||
@@ -1623,12 +1689,31 @@ def get_cp_shared_kv_local_out_cache_loc(forward_batch: "ForwardBatch"):
|
||||
mismatch_reason = "split_out_cache_len_mismatch"
|
||||
out_cache_tokens = int(out_cache_loc.numel())
|
||||
if split_tokens != out_cache_tokens:
|
||||
raise_cp_shared_kv_direct_write_error(
|
||||
mismatch_reason,
|
||||
"split_list tokens=%s out_cache_loc tokens=%s",
|
||||
split_tokens,
|
||||
out_cache_tokens,
|
||||
)
|
||||
static_padded_tokens = _get_forward_batch_static_padded_tokens(forward_batch)
|
||||
if (
|
||||
static_padded_tokens is not None
|
||||
and out_cache_tokens > split_tokens
|
||||
and out_cache_tokens <= static_padded_tokens
|
||||
):
|
||||
# Model-runner/speculative warmup can append one global block of
|
||||
# static padding rows after the valid flattened batch. These rows
|
||||
# may carry dummy cache locs, but CP shared-KV direct-write is a
|
||||
# valid-token operation: never split/write dummy compute rows.
|
||||
#
|
||||
# Do not use CP compute-padding metadata as an implicit allowance
|
||||
# here. Page-tail/owner-lane compute padding is produced by
|
||||
# split_tensor_by_cp_batch_plan() after valid input rows are split;
|
||||
# only forward_batch.extend_num_tokens proves that out_cache_loc
|
||||
# already contains global trailing static padding rows.
|
||||
out_cache_loc = out_cache_loc[:split_tokens]
|
||||
else:
|
||||
raise_cp_shared_kv_direct_write_error(
|
||||
mismatch_reason,
|
||||
"split_list tokens=%s out_cache_loc tokens=%s static_padded=%s",
|
||||
split_tokens,
|
||||
out_cache_tokens,
|
||||
static_padded_tokens,
|
||||
)
|
||||
|
||||
if batch_plan is not None:
|
||||
local_out_cache_loc = split_tensor_by_cp_batch_plan(
|
||||
@@ -1724,6 +1809,7 @@ def cp_split_and_rebuild_position(forward_batch, positions: torch.Tensor):
|
||||
positions,
|
||||
get_cp_shared_kv_batch_plan(forward_batch),
|
||||
mode="position",
|
||||
static_padded_tokens=_get_forward_batch_static_padded_tokens(forward_batch),
|
||||
)
|
||||
|
||||
position_id_list = list(
|
||||
@@ -1833,10 +1919,15 @@ def nsa_cp_round_robin_split_q_seqs(
|
||||
def nsa_use_prefill_cp(forward_batch, nsa_enable_prefill_cp=None):
|
||||
if nsa_enable_prefill_cp is None:
|
||||
nsa_enable_prefill_cp = is_nsa_enable_prefill_cp()
|
||||
forward_mode = getattr(forward_batch, "forward_mode", None)
|
||||
is_context_parallel_extend = (
|
||||
forward_mode is not None
|
||||
and forward_mode.is_context_parallel_extend()
|
||||
) or _is_cp_shared_kv_draft_extend(forward_batch)
|
||||
if (
|
||||
forward_batch.nsa_cp_metadata is not None
|
||||
and nsa_enable_prefill_cp
|
||||
and forward_batch.forward_mode.is_context_parallel_extend()
|
||||
and is_context_parallel_extend
|
||||
):
|
||||
return True
|
||||
else:
|
||||
|
||||
@@ -1002,9 +1002,28 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
spec_info.accept_length = self._pad_tensor_to_size(
|
||||
spec_info.accept_length, bs
|
||||
)
|
||||
spec_info.hidden_states = self._pad_tensor_to_size(
|
||||
spec_info.hidden_states, num_tokens
|
||||
)
|
||||
# EAGLE/NextN draft extend can receive a CP-local hidden side
|
||||
# channel captured by the target model before CP output collect.
|
||||
# This block is already guarded by spec_info.is_draft_input().
|
||||
# prepare_mlp_sync_batch may temporarily rewrite draft forward
|
||||
# modes to EXTEND while static DP padding is prepared, so the
|
||||
# contract must be carried by EagleDraftInput rather than
|
||||
# inferred from forward_mode or tensor length.
|
||||
keep_cp_local_hidden = getattr(spec_info, "cp_local_hidden_states", False)
|
||||
if not keep_cp_local_hidden:
|
||||
if spec_info.hidden_states.shape[0] > num_tokens:
|
||||
raise RuntimeError(
|
||||
"[CP_SHARED_KV_FAIL_FAST]"
|
||||
"[draft_hidden_static_padding_mismatch] "
|
||||
"EAGLE draft hidden_states is larger than the static "
|
||||
"MLP-sync token count but is not marked as CP-local. "
|
||||
f"hidden_tokens={spec_info.hidden_states.shape[0]} "
|
||||
f"num_tokens={num_tokens} "
|
||||
f"forward_mode={self.forward_mode}"
|
||||
)
|
||||
spec_info.hidden_states = self._pad_tensor_to_size(
|
||||
spec_info.hidden_states, num_tokens
|
||||
)
|
||||
|
||||
def prepare_attn_tp_scatter_input(self, model_runner: ModelRunner):
|
||||
from sglang.srt.layers.communicator import get_attn_tp_context
|
||||
|
||||
@@ -622,6 +622,10 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
|
||||
topk_index: torch.Tensor = None
|
||||
# shape: (b, hidden_size)
|
||||
hidden_states: torch.Tensor = None
|
||||
# True when hidden_states is a CP-local side channel captured by the target
|
||||
# model before CP output collect. This is a semantic marker; consumers must
|
||||
# not infer it from tensor length alone.
|
||||
cp_local_hidden_states: bool = False
|
||||
capture_hidden_mode: CaptureHiddenMode = CaptureHiddenMode.FULL
|
||||
|
||||
# Inputs for extend
|
||||
|
||||
@@ -334,6 +334,9 @@ class EAGLEWorker(TpModelWorker):
|
||||
if logits_output.draft_hidden_states is not None
|
||||
else logits_output.hidden_states
|
||||
)
|
||||
cp_local_draft_hidden_states = (
|
||||
logits_output.draft_hidden_states is not None
|
||||
)
|
||||
if (
|
||||
envs.SGLANG_CP_DRAFT_SHARED_KV.get()
|
||||
and draft_hidden_states is None
|
||||
@@ -347,6 +350,7 @@ class EAGLEWorker(TpModelWorker):
|
||||
next_token_ids,
|
||||
seq_lens_cpu,
|
||||
logits_output.mm_input_embeds,
|
||||
cp_local_hidden_states=cp_local_draft_hidden_states,
|
||||
)
|
||||
return GenerationBatchResult(
|
||||
logits_output=logits_output,
|
||||
@@ -935,6 +939,8 @@ class EAGLEWorker(TpModelWorker):
|
||||
next_token_ids: torch.Tensor,
|
||||
seq_lens_cpu: Optional[torch.Tensor],
|
||||
mm_input_embeds: Optional[torch.Tensor] = None,
|
||||
*,
|
||||
cp_local_hidden_states: bool = False,
|
||||
):
|
||||
"""Run draft model extend. This API modifies the states of the batch.
|
||||
|
||||
@@ -945,6 +951,7 @@ class EAGLEWorker(TpModelWorker):
|
||||
"""
|
||||
batch.spec_info = EagleDraftInput(
|
||||
hidden_states=hidden_states,
|
||||
cp_local_hidden_states=cp_local_hidden_states,
|
||||
verified_id=next_token_ids,
|
||||
num_tokens_per_req=1,
|
||||
num_tokens_for_logprob_per_req=1,
|
||||
|
||||
Reference in New Issue
Block a user