[cuda-graph] Size breakable-graph shared buffer from warmup output; slice by produced row count (#30834)

Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
Ming Yang
2026-07-11 10:41:44 -07:00
committed by GitHub
co-authored by Claude
parent 32cb89d412
commit 9b4bb415dd
@@ -110,10 +110,11 @@ class BreakableCudaGraphBackend(DedupedCudaGraphMixin, BaseCudaGraphBackend):
dummies: Optional[Any] = None,
post_warmup_hook: Optional[Callable[[], None]] = None,
) -> None:
warmup_out = None
for _ in range(2):
self._device_module.synchronize()
self._tp_group.barrier()
forward_fn()
warmup_out = forward_fn()
if post_warmup_hook is not None:
post_warmup_hook()
@@ -122,24 +123,55 @@ class BreakableCudaGraphBackend(DedupedCudaGraphMixin, BaseCudaGraphBackend):
eager_on_graph(True)(forward_fn) if self._debug_eager else forward_fn
)
size = shape_key.size
if self._shared_output_buffer is None:
self._shared_output_buffer = self._alloc_full_buffer(warmup_out, size)
with BreakableCUDAGraphCapture(
cuda_graph=graph,
pool=self._pool,
stream=self._capture_stream,
):
out = captured_fn()
if self._shared_output_buffer is not None:
self._copy_output_to_buffer(out, self._shared_output_buffer, size)
if self._shared_output_buffer is None:
self._shared_output_buffer = out
stored = self._slice_output(out, size)
else:
stored = self._slice_output(self._shared_output_buffer, size)
out_rows = self._output_rows(out, size)
self._copy_output_to_buffer(out, self._shared_output_buffer, out_rows)
stored = self._slice_output(self._shared_output_buffer, out_rows)
self._graphs[shape_key] = graph
self._outputs[shape_key] = stored
def _output_rows(self, output: Any, cap: int) -> int:
"""Leading-dim row count actually produced by the body, clamped to ``cap``.
A body that shards or prunes its output along dim 0 returns fewer than
``cap`` rows; everything else returns exactly ``cap``.
"""
if torch.is_tensor(output):
return min(cap, output.shape[0])
if isinstance(output, PPProxyTensors):
rows = [t.shape[0] for t in output.tensors.values()]
return min([cap, *rows])
if isinstance(output, (list, tuple)) and output:
return min(self._output_rows(o, cap) for o in output if o is not None)
return cap
def _alloc_full_buffer(self, output: Any, size: int) -> Any:
"""A same-structure buffer as ``output`` but with ``size`` leading rows."""
if output is None:
return None
if torch.is_tensor(output):
return output.new_empty((size, *output.shape[1:]))
if isinstance(output, PPProxyTensors):
return PPProxyTensors(
{
key: t.new_empty((size, *t.shape[1:]))
for key, t in output.tensors.items()
}
)
if isinstance(output, tuple):
return tuple(self._alloc_full_buffer(o, size) for o in output)
if isinstance(output, list):
return [self._alloc_full_buffer(o, size) for o in output]
raise TypeError(f"Unsupported BCG output type: {type(output)}")
def _slice_output(self, output: Any, num_tokens: int) -> Any:
if output is None:
return None