From a5888c956f9051f29ea05e1ed6896c004f8d4950 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Wed, 5 Aug 2026 19:15:46 +0800 Subject: [PATCH] [diffusion] Pack Ulysses Q/K/V input all-to-all into one collective + reusable a2a staging buffers (#33667) Co-authored-by: Claude Fable 5 --- .../ops/diffusion/triton/ulysses_qkv.py | 20 +-- .../runtime/layers/attention/layer.py | 9 +- .../multimodal_gen/runtime/layers/usp.py | 126 +++++++++++++++++- .../test/unit/test_minimax_h3_dit_contract.py | 2 +- .../test/unit/test_usp_packed_qkv_a2a.py | 66 +++++++++ 5 files changed, 202 insertions(+), 21 deletions(-) create mode 100644 python/sglang/multimodal_gen/test/unit/test_usp_packed_qkv_a2a.py diff --git a/python/sglang/kernels/ops/diffusion/triton/ulysses_qkv.py b/python/sglang/kernels/ops/diffusion/triton/ulysses_qkv.py index a265cf883..9364373f9 100644 --- a/python/sglang/kernels/ops/diffusion/triton/ulysses_qkv.py +++ b/python/sglang/kernels/ops/diffusion/triton/ulysses_qkv.py @@ -57,17 +57,21 @@ def pack_qkv_destination_major( k: torch.Tensor, v: torch.Tensor, world_size: int, + out: torch.Tensor | None = None, ) -> torch.Tensor: rows, global_heads, head_size = q.shape local_heads = global_heads // world_size - output = torch.empty( - world_size, - rows, - local_heads, - 3 * head_size, - dtype=q.dtype, - device=q.device, - ) + expected_shape = (world_size, rows, local_heads, 3 * head_size) + if out is not None: + assert out.shape == expected_shape and out.is_contiguous() + assert out.dtype == q.dtype and out.device == q.device + output = out + else: + output = torch.empty( + expected_shape, + dtype=q.dtype, + device=q.device, + ) total_elements = rows * global_heads * head_size if total_elements == 0: return output diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py index dee99e38e..63966182b 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py @@ -46,6 +46,7 @@ from sglang.multimodal_gen.runtime.layers.attention.turbo_layer import ( from sglang.multimodal_gen.runtime.layers.usp import ( _ipc_input_a2a_qkv, _usp_input_all_to_all, + _usp_input_all_to_all_qkv, _usp_input_all_to_all_varlen, _usp_output_all_to_all, _usp_output_all_to_all_varlen, @@ -784,9 +785,7 @@ class USPAttention(nn.Module): if qkv_fast is not None: q, k, v = qkv_fast else: - q = _usp_input_all_to_all(q, head_dim=2) - k = _usp_input_all_to_all(k, head_dim=2) - v = _usp_input_all_to_all(v, head_dim=2) + q, k, v = _usp_input_all_to_all_qkv(q, k, v) if ( _VARLEN_FA_ENABLED @@ -981,9 +980,7 @@ class USPAttention(nn.Module): k = k.contiguous() v = v.contiguous() else: - q = _usp_input_all_to_all(q, head_dim=2) - k = _usp_input_all_to_all(k, head_dim=2) - v = _usp_input_all_to_all(v, head_dim=2) + q, k, v = _usp_input_all_to_all_qkv(q, k, v) # Ring Attention within subgroups or local attention if get_ring_parallel_world_size() > 1: diff --git a/python/sglang/multimodal_gen/runtime/layers/usp.py b/python/sglang/multimodal_gen/runtime/layers/usp.py index f01ba5341..461dcad59 100644 --- a/python/sglang/multimodal_gen/runtime/layers/usp.py +++ b/python/sglang/multimodal_gen/runtime/layers/usp.py @@ -39,12 +39,45 @@ def _maybe_wait(tensor: torch.Tensor) -> torch.Tensor: return tensor -def _usp_all_to_all_single(x: torch.Tensor) -> torch.Tensor: +_A2A_STAGING_BUFFERS: dict[tuple, torch.Tensor] = {} + + +def _a2a_staging_buffer( + role: str, shape: tuple[int, ...], dtype: torch.dtype, device: torch.device +) -> torch.Tensor: + """Reusable staging buffer for a Ulysses collective. + + 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 + (role, shape, dtype) is exact and removes per-block allocator churn. + Bypassed under autograd and CUDA graph capture: a buffer first allocated + while capturing would live in the graph's private memory pool and must + not be shared with eager replays. + """ + if ( + torch.is_grad_enabled() + or torch.compiler.is_compiling() + or device.type != "cuda" + or torch.cuda.is_current_stream_capturing() + ): + return torch.empty(shape, dtype=dtype, device=device) + key = (role, tuple(shape), dtype, device.index) + buffer = _A2A_STAGING_BUFFERS.get(key) + if buffer is None: + buffer = torch.empty(shape, dtype=dtype, device=device) + _A2A_STAGING_BUFFERS[key] = buffer + return buffer + + +def _usp_all_to_all_single(x: torch.Tensor, role: str | None = None) -> torch.Tensor: ulysses_pg = get_sp_group().ulysses_group assert ulysses_pg is not None, "Ulysses process group is not initialized." x_shape = x.shape x = x.flatten().contiguous() - output = torch.empty_like(x) + if role is None: + output = torch.empty_like(x) + else: + output = _a2a_staging_buffer(role, x.shape, x.dtype, x.device) # USP calls this collective many times per denoising step and waits # immediately, so avoid the extra wrapper overhead of functional collectives. torch.distributed.all_to_all_single(output, x, group=ulysses_pg) @@ -270,7 +303,7 @@ def _usp_input_all_to_all(x: torch.Tensor, head_dim: int = 1) -> torch.Tensor: h_local, s_global = h_global // world_size, s_local * world_size x = x.permute(permute_order).contiguous() - x = _usp_all_to_all_single(x) + x = _usp_all_to_all_single(x, role="usp_input") x = x.reshape(world_size, h_local, b, s_local, d) # Reorder dims to place 'world_size' adjacent to 's_local' to merge them into 's_global' @@ -306,7 +339,18 @@ def _usp_input_all_to_all_packed_qkv( and q.stride(-1) == k.stride(-1) == v.stride(-1) == 1 and not torch.compiler.is_compiling() ): - packed = pack_qkv_destination_major(q, k, v, world_size) + packed = pack_qkv_destination_major( + q, + k, + v, + world_size, + out=_a2a_staging_buffer( + "usp_packed_qkv_src", + (world_size, s_local, h_local, 3 * head_size), + q.dtype, + q.device, + ), + ) else: packed = torch.empty( (world_size, s_local, h_local, 3 * head_size), @@ -319,12 +363,82 @@ def _usp_input_all_to_all_packed_qkv( ) packed[..., index * head_size : (index + 1) * head_size].copy_(head_shards) - packed = _usp_all_to_all_single(packed) + packed = _usp_all_to_all_single(packed, role="usp_packed_qkv_recv") packed = packed.reshape(s_local * world_size, h_local, 3 * head_size) q, k, v = packed.split(head_size, dim=-1) return q, k, v +def _can_use_packed_qkv_a2a_4d( + q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, world_size: int +) -> bool: + return ( + q.is_cuda + and q.ndim == 4 + and q.shape == k.shape == v.shape + and q.dtype == k.dtype == v.dtype + and q.dtype in (torch.float16, torch.bfloat16) + and q.is_contiguous() + and k.is_contiguous() + and v.is_contiguous() + and q.shape[2] % world_size == 0 + and not torch.compiler.is_compiling() + ) + + +def _usp_input_all_to_all_qkv( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Ulysses input exchange for Q/K/V with heads at dim=2. + + [b, s_local, h, d] x3 -> [b, s_global, h_local, d] x3. Q/K/V are packed + destination-major by one relayout kernel and exchanged in a single + collective instead of three; only data movement changes, so the result is + bit-identical to the unpacked path, which stays as the fallback for + ineligible inputs (CPU, GQA-mismatched shapes, non-contiguous layouts). + Adapted from the NVlabs Sana sol-engine branch (Apache-2.0). + """ + world_size = get_ulysses_parallel_world_size() + if world_size <= 1: + return q, k, v + if not _can_use_packed_qkv_a2a_4d(q, k, v, world_size): + return ( + _usp_input_all_to_all(q, head_dim=2), + _usp_input_all_to_all(k, head_dim=2), + _usp_input_all_to_all(v, head_dim=2), + ) + + b, s_local, h_global, d = q.shape + h_local = h_global // world_size + rows = b * s_local + packed = pack_qkv_destination_major( + q.view(rows, h_global, d), + k.view(rows, h_global, d), + v.view(rows, h_global, d), + world_size, + out=_a2a_staging_buffer( + "usp_packed_qkv_src", (world_size, rows, h_local, 3 * d), q.dtype, q.device + ), + ) + packed = _usp_all_to_all_single(packed, role="usp_packed_qkv_recv") + if b == 1: + # Received chunks are already sequence-major: rank j's rows arrive at + # offset j * s_local, so flattening the leading dims is a free view. + packed = packed.view(1, world_size * s_local, h_local, 3 * d) + else: + packed = ( + packed.view(world_size, b, s_local, h_local, 3 * d) + .permute(1, 0, 2, 3, 4) + .contiguous() + .view(b, world_size * s_local, h_local, 3 * d) + ) + q, k, v = packed.split(d, dim=-1) + # Copy out before the staging buffer is recycled by the next collective. + return q.contiguous(), k.contiguous(), v.contiguous() + + def _usp_input_all_to_all_varlen( x: torch.Tensor, seq_lens: list[int], head_dim: int = 1 ) -> torch.Tensor: @@ -455,7 +569,7 @@ def _usp_output_all_to_all(x: torch.Tensor, head_dim: int = 1) -> torch.Tensor: s_local, h_global = s_global // world_size, h_local * world_size x = x.permute(permute_order).contiguous() - x = _usp_all_to_all_single(x) + x = _usp_all_to_all_single(x, role="usp_output") x = x.reshape(world_size, s_local, b, h_local, d) # Reorder dims to place 'world_size' adjacent to 'h_local' to merge them into 'h_global' diff --git a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py index 0a1021d7f..8bb1d0857 100644 --- a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py +++ b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py @@ -223,7 +223,7 @@ def test_packed_qkv_exchange_preserves_rank_and_head_order(_): head_slice = slice(destination * 2, (destination + 1) * 2) return torch.cat((q[:, head_slice], k[:, head_slice], v[:, head_slice]), dim=-1) - def fake_all_to_all(actual): + def fake_all_to_all(actual, role=None): expected_input = torch.stack( [ packet(q_ranks[0], k_ranks[0], v_ranks[0], destination) diff --git a/python/sglang/multimodal_gen/test/unit/test_usp_packed_qkv_a2a.py b/python/sglang/multimodal_gen/test/unit/test_usp_packed_qkv_a2a.py new file mode 100644 index 000000000..a3ee3a309 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_usp_packed_qkv_a2a.py @@ -0,0 +1,66 @@ +"""The packed Ulysses Q/K/V input exchange must be bit-identical to the +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 +r-th chunk); the pack kernel and unpack views run unmodified on CUDA.""" + +import unittest +from unittest.mock import patch + +import torch + +from sglang.multimodal_gen.runtime.layers import usp as usp_mod + +_USP = "sglang.multimodal_gen.runtime.layers.usp" + + +@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA") +class TestPackedQKVInputA2A(unittest.TestCase): + def _run_all_ranks(self, fn, world): + sends, recvs = [], None + + def fake_a2a(x, role=None): + if recvs is None: # recording pass + sends.append(x.detach().clone()) + return torch.empty_like(x) + return recvs.pop(0).reshape(x.shape) + + with ( + patch(f"{_USP}._usp_all_to_all_single", fake_a2a), + patch(f"{_USP}.get_ulysses_parallel_world_size", return_value=world), + ): + for r in range(world): + fn(r) + recvs = [ + torch.cat([s.flatten().chunk(world)[r] for s in sends]) + for r in range(world) + ] + return [fn(r) for r in range(world)] # replay pass + + def test_packed_matches_unpacked_bitwise(self): + for world, b, s_global, h_global, d in ((4, 1, 128, 8, 64), (2, 2, 48, 6, 32)): + torch.manual_seed(1234) + s_local, h_local = s_global // world, h_global // world + full = [ + torch.randn( + b, s_global, h_global, d, dtype=torch.bfloat16, device="cuda" + ) + for _ in range(3) + ] + shards = [ + tuple(t[:, r * s_local : (r + 1) * s_local].contiguous() for t in full) + for r in range(world) + ] + packed = self._run_all_ranks( + lambda r: usp_mod._usp_input_all_to_all_qkv(*shards[r]), world + ) + for r in range(world): + for i in range(3): + spec = full[i][:, :, r * h_local : (r + 1) * h_local].contiguous() + self.assertTrue( + torch.equal(packed[r][i], spec), f"rank{r} qkv[{i}]" + ) + self.assertTrue(packed[r][i].is_contiguous()) + + +if __name__ == "__main__": + unittest.main()