Pass per-forward overrides to ForwardBatch.init_new as explicit arguments (#30670)

This commit is contained in:
fzyzcjy
2026-07-15 14:25:59 +08:00
committed by GitHub
parent 861d97d24d
commit e77d95c3d5
16 changed files with 147 additions and 55 deletions
@@ -0,0 +1,17 @@
---
paths:
- "**/*.py"
---
# `ForwardBatch.init_new` must not mutate the ScheduleBatch
`init_new` (and any ForwardBatch factory) treats the input `ScheduleBatch` as
read-only. Per-forward overrides go through the kw-only params of `init_new` /
`TpModelWorker.forward_batch_generation`, never a ScheduleBatch field. Batch-prep
writes (`out_cache_loc`, `seq_lens_*`, ...) before `init_new` are out of scope.
Tolerated exceptions (don't add new ones):
- `seq_lens_sum` backfill — the `seq_lens` family is slated for removal (→ kv-committed lengths).
- `sampling_info` sub-object writes (grammars, canary ids) — shared object, until the sampling forward-copy op.
- `_expand_mrope_from_input` memoizing `mrope_position_delta_repeated_cache` — pre-existing.
+10 -2
View File
@@ -507,7 +507,11 @@ def extend(reqs, model_runner):
) )
batch.prefill_input_ids_cpu = None batch.prefill_input_ids_cpu = None
forward_batch = ForwardBatch.init_new(batch, model_runner) forward_batch = ForwardBatch.init_new(
batch,
model_runner,
return_hidden_states_before_norm=False,
)
logits_output = model_runner.forward(forward_batch).logits_output logits_output = model_runner.forward(forward_batch).logits_output
next_token_ids = model_runner.sample(logits_output, forward_batch) next_token_ids = model_runner.sample(logits_output, forward_batch)
return next_token_ids, logits_output.next_token_logits, batch return next_token_ids, logits_output.next_token_logits, batch
@@ -518,7 +522,11 @@ def decode(input_token_ids, batch, model_runner):
batch.input_ids = input_token_ids.to(torch.int64) batch.input_ids = input_token_ids.to(torch.int64)
batch.prepare_for_decode() batch.prepare_for_decode()
_maybe_prepare_mlp_sync_batch(batch, model_runner) _maybe_prepare_mlp_sync_batch(batch, model_runner)
forward_batch = ForwardBatch.init_new(batch, model_runner) forward_batch = ForwardBatch.init_new(
batch,
model_runner,
return_hidden_states_before_norm=False,
)
logits_output = model_runner.forward(forward_batch).logits_output logits_output = model_runner.forward(forward_batch).logits_output
next_token_ids = model_runner.sample(logits_output, forward_batch) next_token_ids = model_runner.sample(logits_output, forward_batch)
return next_token_ids, logits_output.next_token_logits return next_token_ids, logits_output.next_token_logits
@@ -26,7 +26,11 @@ from sglang.srt.hardware_backend.mlx.model_runner import (
from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.managers.utils import GenerationBatchResult from sglang.srt.managers.utils import GenerationBatchResult
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode,
ForwardBatch,
PPProxyTensors,
)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -93,6 +97,8 @@ class MlxTpModelWorker(TpModelWorker):
pp_proxy_tensors: Optional[PPProxyTensors] = None, pp_proxy_tensors: Optional[PPProxyTensors] = None,
is_verify: bool = False, is_verify: bool = False,
skip_attn_backend_init: Optional[bool] = None, # deprecated skip_attn_backend_init: Optional[bool] = None, # deprecated
*,
capture_hidden_mode: Optional[CaptureHiddenMode] = None,
) -> GenerationBatchResult: ) -> GenerationBatchResult:
"""Override to route through MLX model runner.""" """Override to route through MLX model runner."""
if batch is not None: if batch is not None:
@@ -106,6 +112,7 @@ class MlxTpModelWorker(TpModelWorker):
pp_proxy_tensors, pp_proxy_tensors,
is_verify, is_verify,
skip_attn_backend_init, skip_attn_backend_init,
capture_hidden_mode=capture_hidden_mode,
) )
def _cleanup_stale_rids(self, forward_mode, current_rids: set[str]) -> None: def _cleanup_stale_rids(self, forward_mode, current_rids: set[str]) -> None:
+1 -9
View File
@@ -94,11 +94,7 @@ from sglang.srt.mem_cache.common import (
) )
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.mem_cache.radix_cache import RadixKey
from sglang.srt.model_executor.forward_batch_info import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
CaptureHiddenMode,
ForwardBatch,
ForwardMode,
)
from sglang.srt.observability.metrics_collector import ( from sglang.srt.observability.metrics_collector import (
DPCooperationInfo, DPCooperationInfo,
SchedulerMetricsCollector, SchedulerMetricsCollector,
@@ -1954,10 +1950,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# spec_info: Optional[SpecInput] = None # spec_info: Optional[SpecInput] = None
spec_info: Optional[SpecInput] = None spec_info: Optional[SpecInput] = None
# === One-shot per-forward overrides; init_new consumes and resets ===
capture_hidden_mode: Optional[CaptureHiddenMode] = None
return_hidden_states_before_norm: bool = False
@classmethod @classmethod
def init_new( def init_new(
cls, cls,
@@ -695,7 +695,11 @@ class SchedulerPPMixin:
) )
batch.prefill_input_ids_cpu = None batch.prefill_input_ids_cpu = None
forward_batch = ForwardBatch.init_new(batch, model_runner) forward_batch = ForwardBatch.init_new(
batch,
model_runner,
return_hidden_states_before_norm=False,
)
set_is_extend_in_batch(batch.forward_mode.is_extend()) set_is_extend_in_batch(batch.forward_mode.is_extend())
_ = model_runner.forward( _ = model_runner.forward(
+26 -4
View File
@@ -41,7 +41,11 @@ from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.managers.scheduler import GenerationBatchResult
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode,
ForwardBatch,
PPProxyTensors,
)
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import MultiprocessingSerializer, broadcast_pyobj, set_random_seed from sglang.srt.utils import MultiprocessingSerializer, broadcast_pyobj, set_random_seed
@@ -226,7 +230,11 @@ class BaseTpWorker(ABC):
return result return result
def forward_batch_embedding(self, batch: ScheduleBatch): def forward_batch_embedding(self, batch: ScheduleBatch):
forward_batch = ForwardBatch.init_new(batch, self.model_runner) forward_batch = ForwardBatch.init_new(
batch,
self.model_runner,
return_hidden_states_before_norm=False,
)
output = self.model_runner.forward(forward_batch).logits_output output = self.model_runner.forward(forward_batch).logits_output
return output # Returns EmbeddingPoolerOutput return output # Returns EmbeddingPoolerOutput
@@ -490,16 +498,26 @@ class TpModelWorker(BaseTpWorker):
pp_proxy_tensors: Optional[PPProxyTensors] = None, pp_proxy_tensors: Optional[PPProxyTensors] = None,
is_verify: bool = False, is_verify: bool = False,
skip_attn_backend_init: Optional[bool] = None, # deprecated skip_attn_backend_init: Optional[bool] = None, # deprecated
*,
capture_hidden_mode: Optional[CaptureHiddenMode] = None,
) -> GenerationBatchResult: ) -> GenerationBatchResult:
# Get forward batch from schedule batch # Get forward batch from schedule batch
if batch is not None: if batch is not None:
# update the consumer index of hicache to the running batch # update the consumer index of hicache to the running batch
self.set_hicache_consumer(batch.hicache_consumer_index) self.set_hicache_consumer(batch.hicache_consumer_index)
forward_batch = ForwardBatch.init_new(batch, self.model_runner) forward_batch = ForwardBatch.init_new(
batch,
self.model_runner,
capture_hidden_mode=capture_hidden_mode,
return_hidden_states_before_norm=False,
)
else: else:
# FIXME(lsyin): unify the interface of forward_batch # FIXME(lsyin): unify the interface of forward_batch
assert forward_batch is not None assert forward_batch is not None
assert (
capture_hidden_mode is None
), "capture_hidden_mode override requires a ScheduleBatch input"
# Deprecated kwarg: pre-planners mark the batch themselves now. # Deprecated kwarg: pre-planners mark the batch themselves now.
forward_batch.apply_deprecated_skip_attn_backend_init(skip_attn_backend_init) forward_batch.apply_deprecated_skip_attn_backend_init(skip_attn_backend_init)
@@ -577,7 +595,11 @@ class TpModelWorker(BaseTpWorker):
def forward_batch_split_prefill(self, batch: ScheduleBatch): def forward_batch_split_prefill(self, batch: ScheduleBatch):
if batch.split_index == 0: if batch.split_index == 0:
forward_batch = ForwardBatch.init_new(batch, self.model_runner) forward_batch = ForwardBatch.init_new(
batch,
self.model_runner,
return_hidden_states_before_norm=False,
)
batch.split_forward_batch = forward_batch batch.split_forward_batch = forward_batch
out = self.model_runner.forward( out = self.model_runner.forward(
@@ -433,7 +433,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# For dumper: request IDs for cross-step sequence tracking # For dumper: request IDs for cross-step sequence tracking
rids: Optional[List[str]] = None rids: Optional[List[str]] = None
# === Resolved from SB one-shot overrides (consumed + reset by init_new) === # === Per-forward overrides passed explicitly to init_new ===
capture_hidden_mode: CaptureHiddenMode = None capture_hidden_mode: CaptureHiddenMode = None
# For hidden states before normal # For hidden states before normal
return_hidden_states_before_norm: bool = False return_hidden_states_before_norm: bool = False
@@ -624,17 +624,15 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
cls, cls,
batch: ScheduleBatch, batch: ScheduleBatch,
model_runner: ModelRunner, model_runner: ModelRunner,
*,
capture_hidden_mode: Optional[CaptureHiddenMode] = None,
return_hidden_states_before_norm: bool,
): ):
# Consume one-shot per-forward overrides from SB; reset to defaults so # init_new must not mutate the input ScheduleBatch; per-forward
# the next forward on the same SB starts clean. See SB field comment # overrides go through explicit keyword arguments.
# for the contract.
capture_hidden_mode = batch.capture_hidden_mode
batch.capture_hidden_mode = None
return_hidden_states_before_norm = batch.return_hidden_states_before_norm
batch.return_hidden_states_before_norm = False
# capture_hidden_mode default: derive from SB.return_hidden_states / # capture_hidden_mode=None means no override: derive from
# spec_info.capture_hidden_mode when caller did not override. # SB.return_hidden_states / spec_info.capture_hidden_mode.
if capture_hidden_mode is None: if capture_hidden_mode is None:
if batch.return_hidden_states: if batch.return_hidden_states:
capture_hidden_mode = CaptureHiddenMode.FULL capture_hidden_mode = CaptureHiddenMode.FULL
@@ -666,6 +664,10 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# block there). Use it directly. # block there). Use it directly.
seq_lens_cpu = batch.seq_lens_cpu seq_lens_cpu = batch.seq_lens_cpu
# TODO(seq-lens-removal): the whole ScheduleBatch seq_lens family
# (incl. seq_lens_sum) is slated for removal in favor of kv-committed
# lengths, so this init_new-time backfill onto the ScheduleBatch is
# tolerated for now despite the init_new-must-not-mutate-SB rule.
if batch.seq_lens_sum is None and seq_lens_cpu is not None: if batch.seq_lens_sum is None and seq_lens_cpu is not None:
batch.seq_lens_sum = int(seq_lens_cpu.sum()) batch.seq_lens_sum = int(seq_lens_cpu.sum())
@@ -114,6 +114,8 @@ class EagleDraftWorkerBase(ABC):
num_draft_tokens: int, num_draft_tokens: int,
draft_model_runner: Any, draft_model_runner: Any,
cuda_graph_runner: Any, cuda_graph_runner: Any,
*,
return_hidden_states_before_norm: bool,
): ):
from sglang.srt.model_executor.forward_batch_info import ( from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode, CaptureHiddenMode,
@@ -161,8 +163,12 @@ class EagleDraftWorkerBase(ABC):
if batch.forward_mode.is_idle() if batch.forward_mode.is_idle()
else ForwardMode.DRAFT_EXTEND_V2 else ForwardMode.DRAFT_EXTEND_V2
) )
batch.capture_hidden_mode = capture_mode forward_batch = ForwardBatch.init_new(
forward_batch = ForwardBatch.init_new(batch, draft_model_runner) batch,
draft_model_runner,
capture_hidden_mode=capture_mode,
return_hidden_states_before_norm=return_hidden_states_before_norm,
)
# Forward sees post-write length (draft extend writes num_draft_tokens # Forward sees post-write length (draft extend writes num_draft_tokens
# slots); mutation stays on forward_batch to preserve SB.seq_lens. # slots); mutation stays on forward_batch to preserve SB.seq_lens.
forward_batch.seq_lens = forward_batch.seq_lens + num_draft_tokens forward_batch.seq_lens = forward_batch.seq_lens + num_draft_tokens
@@ -292,8 +298,12 @@ class EagleDraftWorkerBase(ABC):
else CaptureHiddenMode.LAST else CaptureHiddenMode.LAST
) )
draft_input.positions = batch.seq_lens.repeat_interleave(topk, dim=0) draft_input.positions = batch.seq_lens.repeat_interleave(topk, dim=0)
batch.capture_hidden_mode = capture_mode forward_batch = ForwardBatch.init_new(
forward_batch = ForwardBatch.init_new(batch, draft_model_runner) batch,
draft_model_runner,
capture_hidden_mode=capture_mode,
return_hidden_states_before_norm=False,
)
can_cuda_graph = cuda_graph_runner and cuda_graph_runner.can_run_graph( can_cuda_graph = cuda_graph_runner and cuda_graph_runner.can_run_graph(
forward_batch forward_batch
) )
+6 -2
View File
@@ -69,8 +69,12 @@ class DFlashVerifyInput(SpecInput):
if batch.forward_mode.is_idle() if batch.forward_mode.is_idle()
else ForwardMode.TARGET_VERIFY else ForwardMode.TARGET_VERIFY
) )
batch.capture_hidden_mode = self.capture_hidden_mode verify_forward_batch = ForwardBatch.init_new(
verify_forward_batch = ForwardBatch.init_new(batch, target_worker.model_runner) batch,
target_worker.model_runner,
capture_hidden_mode=self.capture_hidden_mode,
return_hidden_states_before_norm=False,
)
can_run_cuda_graph = bool( can_run_cuda_graph = bool(
target_worker.model_runner.decode_cuda_graph_runner target_worker.model_runner.decode_cuda_graph_runner
@@ -1226,8 +1226,9 @@ class DFlashWorkerV2(BaseSpecWorker):
if batch.forward_mode.is_extend() or batch.is_extend_in_batch: if batch.forward_mode.is_extend() or batch.is_extend_in_batch:
# Target prefill: capture DFlash aux hidden states for prompt tokens. # Target prefill: capture DFlash aux hidden states for prompt tokens.
batch.capture_hidden_mode = CaptureHiddenMode.FULL batch_output = self.target_worker.forward_batch_generation(
batch_output = self.target_worker.forward_batch_generation(batch) batch, capture_hidden_mode=CaptureHiddenMode.FULL
)
logits_output, next_token_ids = ( logits_output, next_token_ids = (
batch_output.logits_output, batch_output.logits_output,
@@ -376,12 +376,14 @@ class DSparkWorkerV2(BaseSpecWorker):
) -> GenerationBatchResult: ) -> GenerationBatchResult:
if batch.forward_mode.is_idle(): if batch.forward_mode.is_idle():
if self.server_args.enable_dp_attention: if self.server_args.enable_dp_attention:
batch.capture_hidden_mode = CaptureHiddenMode.FULL self.target_worker.forward_batch_generation(
self.target_worker.forward_batch_generation(batch) batch, capture_hidden_mode=CaptureHiddenMode.FULL
)
return self._decode_idle_result(on_publish=on_publish) return self._decode_idle_result(on_publish=on_publish)
batch.capture_hidden_mode = CaptureHiddenMode.FULL batch_output = self.target_worker.forward_batch_generation(
batch_output = self.target_worker.forward_batch_generation(batch) batch, capture_hidden_mode=CaptureHiddenMode.FULL
)
logits_output = batch_output.logits_output logits_output = batch_output.logits_output
next_token_ids = batch_output.next_token_ids next_token_ids = batch_output.next_token_ids
batch_output.new_seq_lens = batch.seq_lens batch_output.new_seq_lens = batch.seq_lens
+6 -2
View File
@@ -556,8 +556,12 @@ def eagle_prepare_for_verify(
if target_worker.model_runner.spec_algorithm.is_standalone() if target_worker.model_runner.spec_algorithm.is_standalone()
else CaptureHiddenMode.FULL else CaptureHiddenMode.FULL
) )
batch.capture_hidden_mode = capture_mode verify_forward_batch = ForwardBatch.init_new(
verify_forward_batch = ForwardBatch.init_new(batch, target_worker.model_runner) batch,
target_worker.model_runner,
capture_hidden_mode=capture_mode,
return_hidden_states_before_norm=False,
)
# Run attention backend plan and cuda graph preparation # Run attention backend plan and cuda graph preparation
can_run_cuda_graph = bool( can_run_cuda_graph = bool(
@@ -802,8 +802,12 @@ class EagleDraftWorker(EagleDraftWorkerBase):
if self.speculative_algorithm.is_standalone() if self.speculative_algorithm.is_standalone()
else CaptureHiddenMode.LAST else CaptureHiddenMode.LAST
) )
batch.capture_hidden_mode = capture_hidden_mode forward_batch = ForwardBatch.init_new(
forward_batch = ForwardBatch.init_new(batch, self.draft_runner) batch,
self.draft_runner,
capture_hidden_mode=capture_hidden_mode,
return_hidden_states_before_norm=False,
)
forward_batch.return_logprob = False forward_batch.return_logprob = False
if mm_input_embeds is not None: if mm_input_embeds is not None:
forward_batch.mm_input_embeds = mm_input_embeds forward_batch.mm_input_embeds = mm_input_embeds
@@ -916,6 +920,7 @@ class EagleDraftWorker(EagleDraftWorkerBase):
self.speculative_num_draft_tokens, self.speculative_num_draft_tokens,
self.draft_runner, self.draft_runner,
self.cuda_graph_runner_for_draft_extend, self.cuda_graph_runner_for_draft_extend,
return_hidden_states_before_norm=False,
) )
if self.plan_stream: if self.plan_stream:
@@ -1138,8 +1143,9 @@ class EAGLEWorkerV2(BaseSpecWorker):
if self.speculative_algorithm.is_standalone() if self.speculative_algorithm.is_standalone()
else CaptureHiddenMode.FULL else CaptureHiddenMode.FULL
) )
batch.capture_hidden_mode = target_capture_mode batch_output = self.target_worker.forward_batch_generation(
batch_output = self.target_worker.forward_batch_generation(batch) batch, capture_hidden_mode=target_capture_mode
)
# Spec_v2 convention: batch.seq_lens = length BEFORE this iter's tokens. # Spec_v2 convention: batch.seq_lens = length BEFORE this iter's tokens.
# Extend processed L prompt tokens; next verify iter expects same L. # Extend processed L prompt tokens; next verify iter expects same L.
@@ -427,7 +427,11 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
batch.seq_lens_sum = torch.sum(batch.seq_lens).item() batch.seq_lens_sum = torch.sum(batch.seq_lens).item()
batch.return_hidden_states = False batch.return_hidden_states = False
forward_batch = ForwardBatch.init_new(batch, self.draft_model_runner) forward_batch = ForwardBatch.init_new(
batch,
self.draft_model_runner,
return_hidden_states_before_norm=False,
)
assert forward_batch.capture_hidden_mode == CaptureHiddenMode.LAST assert forward_batch.capture_hidden_mode == CaptureHiddenMode.LAST
self._set_positions(forward_batch) self._set_positions(forward_batch)
self._expand_for_topk_draft(forward_batch) self._expand_for_topk_draft(forward_batch)
@@ -710,8 +714,9 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
# size). The draft / seed-based draft-extend hooks are FrozenKVMTPDraftWorker's. # size). The draft / seed-based draft-extend hooks are FrozenKVMTPDraftWorker's.
if batch.forward_mode.is_extend() or batch.is_extend_in_batch: if batch.forward_mode.is_extend() or batch.is_extend_in_batch:
# Target prefill (frozen is never standalone -> capture FULL hidden). # Target prefill (frozen is never standalone -> capture FULL hidden).
batch.capture_hidden_mode = CaptureHiddenMode.FULL batch_output = self.target_worker.forward_batch_generation(
batch_output = self.target_worker.forward_batch_generation(batch) batch, capture_hidden_mode=CaptureHiddenMode.FULL
)
# Spec_v2 convention: batch.seq_lens = length BEFORE this iter's tokens. # Spec_v2 convention: batch.seq_lens = length BEFORE this iter's tokens.
batch_output.new_seq_lens = batch.seq_lens batch_output.new_seq_lens = batch.seq_lens
@@ -432,9 +432,12 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
draft_capture_hidden_mode = CaptureHiddenMode.LAST draft_capture_hidden_mode = CaptureHiddenMode.LAST
# Run forward # Run forward
batch.capture_hidden_mode = draft_capture_hidden_mode forward_batch = ForwardBatch.init_new(
batch.return_hidden_states_before_norm = True batch,
forward_batch = ForwardBatch.init_new(batch, self.draft_runner_list[0]) self.draft_runner_list[0],
capture_hidden_mode=draft_capture_hidden_mode,
return_hidden_states_before_norm=True,
)
# Construct input_ids # Construct input_ids
# TODO: same chunked-prefill chain divergence as PR #26329. # TODO: same chunked-prefill chain divergence as PR #26329.
@@ -530,8 +533,8 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
self.speculative_num_draft_tokens, self.speculative_num_draft_tokens,
self.draft_runner_list[0], self.draft_runner_list[0],
self.cuda_graph_runner_for_draft_extend, self.cuda_graph_runner_for_draft_extend,
return_hidden_states_before_norm=True,
) )
forward_batch.return_hidden_states_before_norm = True
if self.plan_stream: if self.plan_stream:
torch.get_device_module(self.device).current_stream().wait_stream( torch.get_device_module(self.device).current_stream().wait_stream(
@@ -721,8 +724,9 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
if self.speculative_algorithm.is_standalone() if self.speculative_algorithm.is_standalone()
else CaptureHiddenMode.FULL else CaptureHiddenMode.FULL
) )
batch.capture_hidden_mode = target_capture_mode batch_output = self.target_worker.forward_batch_generation(
batch_output = self.target_worker.forward_batch_generation(batch) batch, capture_hidden_mode=target_capture_mode
)
# Spec_v2 convention: batch.seq_lens = length BEFORE this iter's tokens. # Spec_v2 convention: batch.seq_lens = length BEFORE this iter's tokens.
# Extend processed L prompt tokens; next verify iter expects same L. # Extend processed L prompt tokens; next verify iter expects same L.
+5 -1
View File
@@ -120,7 +120,11 @@ class TestForwardSplitPrefill(CustomTestCase):
batch.forward_mode = ForwardMode.SPLIT_PREFILL batch.forward_mode = ForwardMode.SPLIT_PREFILL
# Create forward batch # Create forward batch
forward_batch = ForwardBatch.init_new(batch, self.model_runner) forward_batch = ForwardBatch.init_new(
batch,
self.model_runner,
return_hidden_states_before_norm=False,
)
return forward_batch return forward_batch