spec_v2: consolidate seq_lens_cpu/sum maintenance into helper (#25818)

This commit is contained in:
Liangsheng Yin
2026-05-20 04:42:26 -07:00
committed by GitHub
parent 9b005d3608
commit 34d3e23232
4 changed files with 29 additions and 9 deletions
+18 -3
View File
@@ -2420,7 +2420,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.seq_lens.add_(1)
self.seq_lens_cpu.add_(1)
self.orig_seq_lens.add_(1)
self.seq_lens_sum += bs
# Defer compute to refresh_seq_lens_cpu (either pre-forward in scheduler.py
# or lazily in ForwardBatch.init_new).
self.seq_lens_sum = None
if self.hisparse_coordinator is not None:
self.hisparse_coordinator.map_last_loc_to_buffer(
@@ -2455,6 +2457,15 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
if draft_input.verify_done is not None:
draft_input.verify_done.wait()
def refresh_seq_lens_cpu(self, sync: bool = True):
# sync=True: D2H from seq_lens (needed when seq_lens_cpu is stale
# relative to seq_lens, i.e. spec v2's mid-forward GPU rebind).
# sync=False: caller asserts seq_lens_cpu already fresh — skip D2H,
# only recompute the cached sum.
if sync and self.is_spec_v2:
self.seq_lens_cpu = self.seq_lens.cpu()
self.seq_lens_sum = int(self.seq_lens_cpu.sum())
def filter_batch(
self,
chunked_req_to_exclude: Optional[Union[Req, List[Req]]] = None,
@@ -2505,7 +2516,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.seq_lens_cpu = self.seq_lens_cpu[keep_indices]
self.orig_seq_lens = self.orig_seq_lens[keep_indices_device]
self.out_cache_loc = None
self.seq_lens_sum = self.seq_lens.sum().item()
# Defer compute to refresh_seq_lens_cpu (either pre-forward in scheduler.py
# or lazily in ForwardBatch.init_new).
self.seq_lens_sum = None
if self.input_ids is not None:
self.input_ids = self.input_ids[keep_indices_device]
@@ -2565,7 +2578,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.seq_lens_cpu = torch.cat([self.seq_lens_cpu, other.seq_lens_cpu])
self.orig_seq_lens = torch.cat([self.orig_seq_lens, other.orig_seq_lens])
self.out_cache_loc = None
self.seq_lens_sum += other.seq_lens_sum
# Defer compute to refresh_seq_lens_cpu (either pre-forward in scheduler.py
# or lazily in ForwardBatch.init_new).
self.seq_lens_sum = None
if self.input_ids is not None:
self.input_ids = torch.cat([self.input_ids, other.input_ids])
self.mamba_track_indices = None
+4
View File
@@ -2835,6 +2835,10 @@ class Scheduler(
# Run forward
if self.is_generation:
if self.enable_overlap:
# Refresh BEFORE _overlap_forward_isolation so snapshot
# captures fresh values and restore preserves them.
batch.refresh_seq_lens_cpu()
with self._overlap_forward_isolation(batch):
bs = len(batch.seq_lens)
future_indices = self.future_map.alloc_future_indices(bs)
@@ -503,6 +503,9 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
else:
seq_lens_cpu = batch.seq_lens_cpu
if batch.seq_lens_sum is None:
batch.refresh_seq_lens_cpu(sync=False)
ret = cls(
forward_mode=batch.forward_mode,
batch_size=len(batch.seq_lens),
@@ -173,10 +173,6 @@ class EagleDraftInputV2Mixin:
bs,
)
# FIXME(lsyin): make this sync optional
batch.seq_lens_cpu = batch.seq_lens.cpu()
batch.seq_lens_sum = batch.seq_lens_cpu.sum().item()
def prepare_for_v2_draft(
self: EagleDraftInput,
req_to_token_pool: ReqToTokenPool,
@@ -235,7 +231,9 @@ class EagleDraftInputV2Mixin:
batch.input_ids = predict
batch.seq_lens = batch.seq_lens + num_draft_tokens
batch.seq_lens_cpu = batch.seq_lens_cpu + num_draft_tokens
batch.seq_lens_sum += extend_num_tokens
# seq_lens_cpu was just CPU-updated in tandem — sync=False avoids
# a redundant D2H on the draft hot path.
batch.refresh_seq_lens_cpu(sync=False)
batch.extend_lens = [num_draft_tokens for _ in range(len(batch.seq_lens))]
batch.prefix_lens = seq_lens_cpu_.tolist()
batch.extend_num_tokens = extend_num_tokens
@@ -288,7 +286,7 @@ class EagleVerifyInputV2Mixin:
# Populate seq_lens_cpu/seq_lens_sum on the verify input so that
# TBO's split_spec_info can slice the custom_mask correctly.
self.seq_lens_cpu = batch.seq_lens_cpu
self.seq_lens_sum = batch.seq_lens_sum
self.seq_lens_sum = int(batch.seq_lens_cpu.sum())
# Get a forward batch
batch.forward_mode = (