[core] WAR barrier for overlap schedule buffer writes, without fwd occupancy cost (#26380)
This commit is contained in:
@@ -105,6 +105,16 @@ class FutureMap:
|
||||
self.new_seq_lens_buf = torch.empty(
|
||||
(self.req_pool_size,), dtype=torch.int64, device=self.device
|
||||
)
|
||||
# Pinned host copy of new_seq_lens_buf + private stream for fwd-prepare
|
||||
# D2H pulls (gated only on publish, off the schedule stream).
|
||||
if _is_cuda or _is_hip:
|
||||
self.new_seq_lens_cpu_pinned = torch.empty(
|
||||
(self.req_pool_size,), dtype=torch.int64, pin_memory=True
|
||||
)
|
||||
self.fwd_prepare_d2h_stream = torch.get_device_module(self.device).Stream()
|
||||
else:
|
||||
self.new_seq_lens_cpu_pinned = None
|
||||
self.fwd_prepare_d2h_stream = None
|
||||
if self.spec_algo.is_some():
|
||||
self._forward_buf_initialized = False
|
||||
|
||||
@@ -194,9 +204,20 @@ class FutureMap:
|
||||
return
|
||||
if self.publish_ready is not None:
|
||||
self.publish_ready.wait()
|
||||
new_seq_lens = self.new_seq_lens_buf[fi]
|
||||
batch.seq_lens = new_seq_lens
|
||||
batch.seq_lens_cpu = new_seq_lens.cpu()
|
||||
batch.seq_lens = self.new_seq_lens_buf[fi]
|
||||
|
||||
if self.fwd_prepare_d2h_stream is None or self.publish_ready is None:
|
||||
batch.seq_lens_cpu = batch.seq_lens.cpu() # bootstrap / non-CUDA
|
||||
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.
|
||||
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()
|
||||
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())
|
||||
|
||||
def publish(self, future_indices: torch.Tensor, new_seq_lens: torch.Tensor) -> None:
|
||||
|
||||
@@ -1408,6 +1408,9 @@ class Scheduler(
|
||||
if self._engine_paused:
|
||||
continue
|
||||
|
||||
# WAR barrier: this iter's schedule writes to shared GPU buffers wait for prev forward's reads.
|
||||
self.schedule_stream.wait_stream(self.forward_stream)
|
||||
|
||||
# Get the next batch to run
|
||||
batch = self.get_next_batch_to_run()
|
||||
self.cur_batch = batch
|
||||
|
||||
Reference in New Issue
Block a user