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)