refresh resolve_seq_lens_cpu comments (#26463)
This commit is contained in:
@@ -196,9 +196,10 @@ class FutureMap:
|
||||
batch.input_ids = -future_indices
|
||||
|
||||
def resolve_seq_lens_cpu(self, batch: ScheduleBatch) -> None:
|
||||
# Lazy pull from new_seq_lens_buf for spec_v2 (accept_lens not known to
|
||||
# schedule). Write into both CPU and GPU so SB.seq_lens stays a faithful
|
||||
# seq_lens_cpu mirror.
|
||||
# seq_lens_cpu may be needed on the host for kernel-launch prep (some backends).
|
||||
# Run this D2H on a standalone stream to avoid chain-blocking forward_n ->
|
||||
# prepare_{n+1}: a sync on the schedule stream would inherit its WAR barrier and
|
||||
# stall the host until forward_n ends.
|
||||
fi = batch.spec_info.future_indices if batch.spec_info is not None else None
|
||||
if fi is None:
|
||||
return
|
||||
@@ -211,12 +212,14 @@ class FutureMap:
|
||||
batch.seq_lens_sum = int(batch.seq_lens_cpu.sum())
|
||||
return
|
||||
|
||||
# seq_lens_cpu off the schedule stream: D2H the relay buf on a private
|
||||
# stream (gated on publish), host-select via req_pool_indices_cpu.
|
||||
# Mechanism: don't sync the schedule stream; gate a private stream on the
|
||||
# publish event and copy into the static pinned buffer.
|
||||
self.fwd_prepare_d2h_stream.wait_event(self.publish_ready)
|
||||
with torch.get_device_module(self.device).stream(self.fwd_prepare_d2h_stream):
|
||||
self.new_seq_lens_cpu_pinned.copy_(self.new_seq_lens_buf, non_blocking=True)
|
||||
self.fwd_prepare_d2h_stream.synchronize()
|
||||
|
||||
# FIXME: fi == batch.req_pool_indices; unify future_indices and req_pool_indices.
|
||||
batch.seq_lens_cpu = self.new_seq_lens_cpu_pinned[batch.req_pool_indices_cpu]
|
||||
batch.seq_lens_sum = int(batch.seq_lens_cpu.sum())
|
||||
|
||||
@@ -225,7 +228,7 @@ class FutureMap:
|
||||
if indices.shape[0] == 0:
|
||||
return # DP idle
|
||||
self.new_seq_lens_buf[indices] = new_seq_lens.to(self.new_seq_lens_buf.dtype)
|
||||
# Fast path: only spec_v2 needs the event (schedule-stream D2H sync).
|
||||
# Only spec_v2 needs the event; it gates the seq_lens D2H on the private stream.
|
||||
if self.spec_algo.is_some():
|
||||
if self.publish_ready is None:
|
||||
self.publish_ready = torch.get_device_module(self.device).Event()
|
||||
|
||||
Reference in New Issue
Block a user