[Diffusion] Bound reusable Ulysses A2A staging buffers across shapes (#36327)
This commit is contained in:
@@ -1,6 +1,7 @@
|
|||||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import math
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -37,7 +38,7 @@ def _maybe_wait(tensor: torch.Tensor) -> torch.Tensor:
|
|||||||
return tensor
|
return tensor
|
||||||
|
|
||||||
|
|
||||||
_A2A_STAGING_BUFFERS: dict[tuple, torch.Tensor] = {}
|
_A2A_STAGING_BUFFERS: dict[tuple[str, torch.dtype, int], torch.Tensor] = {}
|
||||||
|
|
||||||
|
|
||||||
def _a2a_staging_buffer(
|
def _a2a_staging_buffer(
|
||||||
@@ -46,8 +47,10 @@ def _a2a_staging_buffer(
|
|||||||
"""Reusable staging buffer for a Ulysses collective.
|
"""Reusable staging buffer for a Ulysses collective.
|
||||||
|
|
||||||
A buffer of a given role is fully consumed (in stream order) before the
|
A buffer of a given role is fully consumed (in stream order) before the
|
||||||
next collective with the same role overwrites it, so caching by
|
next collective with the same role overwrites it. Keep one grow-only
|
||||||
(role, shape, dtype) is exact and removes per-block allocator churn.
|
backing allocation per (role, dtype, device), then return an exact-shape
|
||||||
|
view into that allocation. This removes per-block allocator churn without
|
||||||
|
retaining one CUDA tensor for every request shape seen by the worker.
|
||||||
Bypassed under autograd and CUDA graph capture: a buffer first allocated
|
Bypassed under autograd and CUDA graph capture: a buffer first allocated
|
||||||
while capturing would live in the graph's private memory pool and must
|
while capturing would live in the graph's private memory pool and must
|
||||||
not be shared with eager replays.
|
not be shared with eager replays.
|
||||||
@@ -59,12 +62,22 @@ def _a2a_staging_buffer(
|
|||||||
or torch.cuda.is_current_stream_capturing()
|
or torch.cuda.is_current_stream_capturing()
|
||||||
):
|
):
|
||||||
return torch.empty(shape, dtype=dtype, device=device)
|
return torch.empty(shape, dtype=dtype, device=device)
|
||||||
key = (role, tuple(shape), dtype, device.index)
|
|
||||||
|
device_index = device.index
|
||||||
|
if device_index is None:
|
||||||
|
device_index = torch.cuda.current_device()
|
||||||
|
key = (role, dtype, device_index)
|
||||||
|
required_numel = math.prod(shape)
|
||||||
buffer = _A2A_STAGING_BUFFERS.get(key)
|
buffer = _A2A_STAGING_BUFFERS.get(key)
|
||||||
if buffer is None:
|
if buffer is None or buffer.numel() < required_numel:
|
||||||
buffer = torch.empty(shape, dtype=dtype, device=device)
|
# The previous same-role collective is fully consumed by contract, so
|
||||||
|
# drop its cache reference before allocating a larger backing buffer.
|
||||||
|
# Any outstanding tensor view still keeps the old storage alive.
|
||||||
|
_A2A_STAGING_BUFFERS.pop(key, None)
|
||||||
|
del buffer
|
||||||
|
buffer = torch.empty(required_numel, dtype=dtype, device=device)
|
||||||
_A2A_STAGING_BUFFERS[key] = buffer
|
_A2A_STAGING_BUFFERS[key] = buffer
|
||||||
return buffer
|
return buffer[:required_numel].view(shape)
|
||||||
|
|
||||||
|
|
||||||
def _usp_all_to_all_single(x: torch.Tensor, role: str | None = None) -> torch.Tensor:
|
def _usp_all_to_all_single(x: torch.Tensor, role: str | None = None) -> torch.Tensor:
|
||||||
|
|||||||
@@ -3,18 +3,130 @@ unpacked path. The collective is emulated in-process with exact
|
|||||||
``all_to_all_single`` chunk semantics (rank r's j-th chunk goes to rank j's
|
``all_to_all_single`` chunk semantics (rank r's j-th chunk goes to rank j's
|
||||||
r-th chunk); the pack kernel and unpack views run unmodified on CUDA."""
|
r-th chunk); the pack kernel and unpack views run unmodified on CUDA."""
|
||||||
|
|
||||||
|
import math
|
||||||
import unittest
|
import unittest
|
||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.layers import usp as usp_mod
|
from sglang.multimodal_gen.runtime.layers import usp as usp_mod
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
_USP = "sglang.multimodal_gen.runtime.layers.usp"
|
_USP = "sglang.multimodal_gen.runtime.layers.usp"
|
||||||
|
|
||||||
|
|
||||||
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
|
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
|
||||||
class TestPackedQKVInputA2A(unittest.TestCase):
|
class TestA2AStagingBuffer(CustomTestCase):
|
||||||
|
def setUp(self):
|
||||||
|
super().setUp()
|
||||||
|
usp_mod._A2A_STAGING_BUFFERS.clear()
|
||||||
|
|
||||||
|
def tearDown(self):
|
||||||
|
usp_mod._A2A_STAGING_BUFFERS.clear()
|
||||||
|
super().tearDown()
|
||||||
|
|
||||||
|
def test_cache_capacity_is_bounded_across_shapes(self):
|
||||||
|
device = torch.device("cuda", torch.cuda.current_device())
|
||||||
|
role = "test_role"
|
||||||
|
shapes = ((2, 3), (4, 5), (5, 4), (1, 7), (3, 11), (2, 4))
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
for shape in shapes:
|
||||||
|
actual = usp_mod._a2a_staging_buffer(
|
||||||
|
role, shape, torch.bfloat16, device
|
||||||
|
)
|
||||||
|
self.assertEqual(actual.shape, shape)
|
||||||
|
self.assertTrue(actual.is_contiguous())
|
||||||
|
|
||||||
|
key = (role, torch.bfloat16, device.index)
|
||||||
|
self.assertEqual(list(usp_mod._A2A_STAGING_BUFFERS), [key])
|
||||||
|
self.assertEqual(
|
||||||
|
usp_mod._A2A_STAGING_BUFFERS[key].numel(),
|
||||||
|
max(math.prod(shape) for shape in shapes),
|
||||||
|
)
|
||||||
|
retained_bytes = sum(
|
||||||
|
tensor.numel() * tensor.element_size()
|
||||||
|
for tensor in usp_mod._A2A_STAGING_BUFFERS.values()
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
retained_bytes,
|
||||||
|
max(math.prod(shape) for shape in shapes)
|
||||||
|
* torch.empty((), dtype=torch.bfloat16).element_size(),
|
||||||
|
)
|
||||||
|
self.assertLess(
|
||||||
|
retained_bytes,
|
||||||
|
sum(math.prod(shape) for shape in shapes)
|
||||||
|
* torch.empty((), dtype=torch.bfloat16).element_size(),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_smaller_shape_reuses_larger_backing_buffer(self):
|
||||||
|
device = torch.device("cuda", torch.cuda.current_device())
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
large = usp_mod._a2a_staging_buffer(
|
||||||
|
"test_role", (8, 16), torch.float16, device
|
||||||
|
)
|
||||||
|
small = usp_mod._a2a_staging_buffer(
|
||||||
|
"test_role", (2, 7), torch.float16, device
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(large.untyped_storage().data_ptr(), small.data_ptr())
|
||||||
|
self.assertEqual(small.shape, (2, 7))
|
||||||
|
|
||||||
|
def test_role_and_dtype_are_separate_cache_keys(self):
|
||||||
|
device = torch.device("cuda", torch.cuda.current_device())
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
usp_mod._a2a_staging_buffer("input", (8,), torch.float16, device)
|
||||||
|
usp_mod._a2a_staging_buffer("output", (8,), torch.float16, device)
|
||||||
|
usp_mod._a2a_staging_buffer("input", (8,), torch.bfloat16, device)
|
||||||
|
|
||||||
|
self.assertEqual(len(usp_mod._A2A_STAGING_BUFFERS), 3)
|
||||||
|
|
||||||
|
def test_bypass_paths_do_not_replace_cached_storage(self):
|
||||||
|
cuda_device = torch.device("cuda", torch.cuda.current_device())
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
cached = usp_mod._a2a_staging_buffer(
|
||||||
|
"test_role", (8,), torch.float16, cuda_device
|
||||||
|
)
|
||||||
|
key = ("test_role", torch.float16, cuda_device.index)
|
||||||
|
cached_storage = usp_mod._A2A_STAGING_BUFFERS[key]
|
||||||
|
|
||||||
|
with torch.enable_grad():
|
||||||
|
grad_buffer = usp_mod._a2a_staging_buffer(
|
||||||
|
"test_role", (16,), torch.float16, cuda_device
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
torch.no_grad(),
|
||||||
|
patch(f"{_USP}.torch.compiler.is_compiling", return_value=True),
|
||||||
|
):
|
||||||
|
compile_buffer = usp_mod._a2a_staging_buffer(
|
||||||
|
"test_role", (16,), torch.float16, cuda_device
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
torch.no_grad(),
|
||||||
|
patch(f"{_USP}.torch.cuda.is_current_stream_capturing", return_value=True),
|
||||||
|
):
|
||||||
|
capture_buffer = usp_mod._a2a_staging_buffer(
|
||||||
|
"test_role", (16,), torch.float16, cuda_device
|
||||||
|
)
|
||||||
|
with torch.no_grad():
|
||||||
|
cpu_buffer = usp_mod._a2a_staging_buffer(
|
||||||
|
"cpu", (8,), torch.float16, torch.device("cpu")
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(list(usp_mod._A2A_STAGING_BUFFERS), [key])
|
||||||
|
self.assertIs(usp_mod._A2A_STAGING_BUFFERS[key], cached_storage)
|
||||||
|
self.assertEqual(cached.numel(), 8)
|
||||||
|
self.assertEqual(grad_buffer.device.type, "cuda")
|
||||||
|
self.assertEqual(compile_buffer.device.type, "cuda")
|
||||||
|
self.assertEqual(capture_buffer.device.type, "cuda")
|
||||||
|
self.assertEqual(cpu_buffer.device.type, "cpu")
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
|
||||||
|
class TestPackedQKVInputA2A(CustomTestCase):
|
||||||
def _run_all_ranks(self, fn, world):
|
def _run_all_ranks(self, fn, world):
|
||||||
sends, recvs = [], None
|
sends, recvs = [], None
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user