From dbb9d75ec4e420ac71c19e9f4514beec3b43a496 Mon Sep 17 00:00:00 2001 From: dongbo910220 <1275604947@qq.com> Date: Mon, 14 Sep 2026 06:37:44 +0800 Subject: [PATCH 1/2] [Perf] Fuse LongCat paired Q/K RoPE Signed-off-by: dongbo910220 <1275604947@qq.com> --- tests/diffusion/layers/test_fused_qk_rope.py | 125 +++++++++ .../test_longcat_image_transformer.py | 245 +++++++++++++++++ vllm_omni/diffusion/layers/fused_qk_rope.py | 251 ++++++++++++++++++ .../longcat_image_transformer.py | 199 ++++++++++++-- 4 files changed, 800 insertions(+), 20 deletions(-) create mode 100644 tests/diffusion/layers/test_fused_qk_rope.py create mode 100644 tests/diffusion/models/longcat_image/test_longcat_image_transformer.py create mode 100644 vllm_omni/diffusion/layers/fused_qk_rope.py diff --git a/tests/diffusion/layers/test_fused_qk_rope.py b/tests/diffusion/layers/test_fused_qk_rope.py new file mode 100644 index 00000000000..c37075ba665 --- /dev/null +++ b/tests/diffusion/layers/test_fused_qk_rope.py @@ -0,0 +1,125 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project + +import pytest +import torch +from diffusers.models.embeddings import apply_rotary_emb +from vllm.triton_utils import HAS_TRITON + +pytestmark = [pytest.mark.core_model, pytest.mark.diffusion] + + +def test_fused_qk_rope_fake_and_cpu_eligibility(): + from vllm_omni.diffusion.layers.fused_qk_rope import ( + _fused_qk_rope_fake, + fused_qk_rope, + fused_qk_rope_supported, + ) + + q = torch.randn(2, 3, 4, 8, dtype=torch.bfloat16) + k = torch.randn_like(q) + cos = torch.randn(3, 8) + sin = torch.randn_like(cos) + + fake_q, fake_k = _fused_qk_rope_fake(q, k, cos, sin) + assert fake_q.shape == q.shape and fake_q.dtype == q.dtype + assert fake_k.shape == k.shape and fake_k.dtype == k.dtype + assert not fused_qk_rope_supported(q, k, cos, sin) + with pytest.raises(ValueError, match="contiguous CUDA BF16"): + fused_qk_rope(q, k, cos, sin) + + +def _cuda_inputs(batch: int, sequence: int, heads: int = 24, head_dim: int = 128): + q = torch.randn(batch, sequence, heads, head_dim, device="cuda", dtype=torch.bfloat16) + k = torch.randn_like(q) + # Deliberately do not repeat adjacent table columns. Diffusers consumes + # full-width tables, so even and odd outputs may use different values. + cos = torch.randn(sequence, head_dim, device="cuda", dtype=torch.float32) + sin = torch.randn_like(cos) + return q, k, cos, sin + + +def _reference(q, k, cos, sin): + rotary_emb = (cos, sin) + return ( + apply_rotary_emb(q, rotary_emb, sequence_dim=1), + apply_rotary_emb(k, rotary_emb, sequence_dim=1), + ) + + +@pytest.mark.cuda +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +@pytest.mark.skipif(not HAS_TRITON, reason="Triton required") +@pytest.mark.parametrize("sequence", [512, 4096, 4608]) +def test_fused_qk_rope_is_bit_exact_at_longcat_production_shapes(sequence): + from vllm_omni.diffusion.layers.fused_qk_rope import _launch_fused_qk_rope, fused_qk_rope + + torch.manual_seed(sequence) + with torch.inference_mode(): + q, k, cos, sin = _cuda_inputs(1, sequence) + expected = _reference(q, k, cos, sin) + for actual in ( + _launch_fused_qk_rope(q, k, cos, sin), + fused_qk_rope(q, k, cos, sin), + ): + assert torch.equal(actual[0], expected[0]) + assert torch.equal(actual[1], expected[1]) + + +@pytest.mark.cuda +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +@pytest.mark.skipif(not HAS_TRITON, reason="Triton required") +def test_fused_qk_rope_is_bit_exact_for_batch_and_full_width_odd_even_tables(): + from vllm_omni.diffusion.layers.fused_qk_rope import fused_qk_rope + + torch.manual_seed(11) + with torch.inference_mode(): + q, k, cos, sin = _cuda_inputs(2, 17, heads=3) + assert not torch.equal(cos[:, ::2], cos[:, 1::2]) + assert not torch.equal(sin[:, ::2], sin[:, 1::2]) + expected = _reference(q, k, cos, sin) + actual = fused_qk_rope(q, k, cos, sin) + + assert torch.equal(actual[0], expected[0]) + assert torch.equal(actual[1], expected[1]) + + +@pytest.mark.cuda +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +@pytest.mark.skipif(not HAS_TRITON, reason="Triton required") +def test_fused_qk_rope_rejects_noncontiguous_and_non_fp32_tables(): + from vllm_omni.diffusion.layers.fused_qk_rope import fused_qk_rope_supported + + q, k, cos, sin = _cuda_inputs(1, 8, heads=2) + q_noncontiguous = torch.randn(1, 8, 2, 256, device="cuda", dtype=torch.bfloat16)[..., ::2] + k_noncontiguous = torch.randn_like(q_noncontiguous) + cos_noncontiguous = torch.randn(8, 256, device="cuda")[:, ::2] + sin_noncontiguous = torch.randn(8, 256, device="cuda")[:, ::2] + with torch.inference_mode(): + assert fused_qk_rope_supported(q, k, cos, sin) + assert not fused_qk_rope_supported(q_noncontiguous, k_noncontiguous, cos, sin) + assert not fused_qk_rope_supported(q, k, cos_noncontiguous, sin_noncontiguous) + assert not fused_qk_rope_supported(q, k, cos.to(torch.bfloat16), sin.to(torch.bfloat16)) + + q_wide = torch.randn(1, 8, 2, 256, device="cuda", dtype=torch.bfloat16) + k_wide = torch.randn_like(q_wide) + cos_wide = torch.randn(8, 256, device="cuda") + sin_wide = torch.randn_like(cos_wide) + assert not fused_qk_rope_supported(q_wide, k_wide, cos_wide, sin_wide) + + +@pytest.mark.cuda +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +@pytest.mark.skipif(not HAS_TRITON, reason="Triton required") +def test_fused_qk_rope_custom_op_has_fullgraph_fake(): + from vllm_omni.diffusion.layers.fused_qk_rope import fused_qk_rope + + torch.manual_seed(29) + with torch.inference_mode(): + q, k, cos, sin = _cuda_inputs(1, 19, heads=3) + expected = _reference(q, k, cos, sin) + compiled = torch.compile(fused_qk_rope, dynamic=True, fullgraph=True) + actual = compiled(q, k, cos, sin) + + assert torch.equal(actual[0], expected[0]) + assert torch.equal(actual[1], expected[1]) diff --git a/tests/diffusion/models/longcat_image/test_longcat_image_transformer.py b/tests/diffusion/models/longcat_image/test_longcat_image_transformer.py new file mode 100644 index 00000000000..8e105479fa7 --- /dev/null +++ b/tests/diffusion/models/longcat_image/test_longcat_image_transformer.py @@ -0,0 +1,245 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project + +from contextlib import nullcontext +from types import SimpleNamespace + +import pytest +import torch +from diffusers.models.embeddings import apply_rotary_emb + +pytestmark = [pytest.mark.core_model, pytest.mark.cpu, pytest.mark.diffusion] + + +@pytest.fixture(autouse=True) +def _clear_qk_rope_signature_caches(): + from vllm_omni.diffusion.models.longcat_image import longcat_image_transformer as longcat + + longcat._VERIFIED_QK_ROPE_SIGNATURES.clear() + longcat._FAILED_QK_ROPE_SIGNATURES.clear() + yield + longcat._VERIFIED_QK_ROPE_SIGNATURES.clear() + longcat._FAILED_QK_ROPE_SIGNATURES.clear() + + +def _inputs(sequence: int = 5, head_dim: int = 8): + q = torch.randn(2, sequence, 3, head_dim, dtype=torch.bfloat16) + k = torch.randn_like(q) + cos = torch.randn(sequence, head_dim) + sin = torch.randn_like(cos) + return q, k, (cos, sin) + + +def _reference(q, k, rotary_emb): + return ( + apply_rotary_emb(q, rotary_emb, sequence_dim=1), + apply_rotary_emb(k, rotary_emb, sequence_dim=1), + ) + + +def test_prepare_rotary_emb_resolves_threshold_once_only_when_enabled(monkeypatch): + from vllm_omni.diffusion.models.longcat_image import longcat_image_transformer as longcat + + calls = 0 + + def resolve(default): + nonlocal calls + calls += 1 + assert default == 512 + return 37 + + monkeypatch.setattr(longcat, "fused_qk_norm_rope_min_tokens", resolve) + txt_cos, txt_sin = torch.randn(2, 8), torch.randn(2, 8) + img_cos, img_sin = torch.randn(3, 8), torch.randn(3, 8) + native = longcat._prepare_rotary_emb(txt_cos, txt_sin, img_cos, img_sin, enable_fusion=False) + fused = longcat._prepare_rotary_emb(txt_cos, txt_sin, img_cos, img_sin, enable_fusion=True) + + assert calls == 1 + assert len(native) == 2 + assert len(fused) == 3 and fused[2] == 37 + assert torch.equal(native[0], fused[0]) + assert torch.equal(native[1], fused[1]) + + +@pytest.mark.parametrize( + ("mode", "sp_size"), + [ + ("unsupported", 1), + ("compile", 1), + ("grad", 1), + ("capture", 1), + ("sequence_parallel", 2), + ("below_threshold", 1), + ], +) +def test_qk_rope_non_eager_modes_use_exact_original_path(monkeypatch, mode, sp_size): + from vllm_omni.diffusion.models.longcat_image import longcat_image_transformer as longcat + + torch.manual_seed(1) + q, k, rotary_emb = _inputs() + expected = _reference(q, k, rotary_emb) + monkeypatch.setattr(torch.compiler, "is_compiling", lambda: mode == "compile") + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: mode == "capture") + monkeypatch.setattr(longcat, "fused_qk_rope_supported", lambda *args: mode != "unsupported") + monkeypatch.setattr( + longcat, + "fused_qk_rope", + lambda *args: (_ for _ in ()).throw(AssertionError("fallback called fused op")), + ) + + context = nullcontext() if mode == "grad" else torch.no_grad() + min_tokens = q.shape[0] * q.shape[1] + 1 if mode == "below_threshold" else 0 + with context: + actual = longcat._apply_qk_rope(q, k, rotary_emb, sp_size, min_tokens) + + assert torch.equal(actual[0], expected[0]) + assert torch.equal(actual[1], expected[1]) + + +def test_qk_rope_verifies_full_output_once_per_signature(monkeypatch): + from vllm_omni.diffusion.models.longcat_image import longcat_image_transformer as longcat + + torch.manual_seed(2) + monkeypatch.setattr(torch.compiler, "is_compiling", lambda: False) + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: False) + monkeypatch.setattr(longcat, "fused_qk_rope_supported", lambda *args: True) + + raw_reference = longcat._apply_qk_rope_reference + reference_calls = 0 + fused_calls = 0 + + def tracked_reference(*args): + nonlocal reference_calls + reference_calls += 1 + return raw_reference(*args) + + def exact_fused(query, key, cos, sin): + nonlocal fused_calls + fused_calls += 1 + return raw_reference(query, key, (cos, sin)) + + monkeypatch.setattr(longcat, "_apply_qk_rope_reference", tracked_reference) + monkeypatch.setattr(longcat, "fused_qk_rope", exact_fused) + + with torch.no_grad(): + q, k, rotary_emb = _inputs(sequence=5) + first = longcat._apply_qk_rope(q, k, rotary_emb, 1, 0) + second = longcat._apply_qk_rope(q.clone(), k.clone(), rotary_emb, 1, 0) + q_other, k_other, rotary_other = _inputs(sequence=7) + third = longcat._apply_qk_rope(q_other, k_other, rotary_other, 1, 0) + + assert torch.equal(first[0], _reference(q, k, rotary_emb)[0]) + assert torch.equal(second[1], _reference(q, k, rotary_emb)[1]) + assert torch.equal(third[0], _reference(q_other, k_other, rotary_other)[0]) + assert fused_calls == 3 + assert reference_calls == 2 + assert len(longcat._VERIFIED_QK_ROPE_SIGNATURES) == 2 + assert not longcat._FAILED_QK_ROPE_SIGNATURES + + +@pytest.mark.parametrize("failure", ["mismatch", "exception"]) +def test_qk_rope_first_use_failure_permanently_falls_back(monkeypatch, failure): + from vllm_omni.diffusion.models.longcat_image import longcat_image_transformer as longcat + + torch.manual_seed(3) + monkeypatch.setattr(torch.compiler, "is_compiling", lambda: False) + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: False) + monkeypatch.setattr(longcat, "fused_qk_rope_supported", lambda *args: True) + fused_calls = 0 + + def broken_fused(query, key, cos, sin): + nonlocal fused_calls + fused_calls += 1 + if failure == "exception": + raise RuntimeError("kernel failed") + return torch.zeros_like(query), torch.zeros_like(key) + + monkeypatch.setattr(longcat, "fused_qk_rope", broken_fused) + q, k, rotary_emb = _inputs() + expected = _reference(q, k, rotary_emb) + with torch.no_grad(): + first = longcat._apply_qk_rope(q, k, rotary_emb, 1, 0) + second = longcat._apply_qk_rope(q, k, rotary_emb, 1, 0) + + assert fused_calls == 1 + assert len(longcat._FAILED_QK_ROPE_SIGNATURES) == 1 + assert not longcat._VERIFIED_QK_ROPE_SIGNATURES + for actual in (first, second): + assert torch.equal(actual[0], expected[0]) + assert torch.equal(actual[1], expected[1]) + + +class _FakeQKV(torch.nn.Module): + def __init__(self, heads: int, head_dim: int) -> None: + super().__init__() + self.num_heads = heads + self.num_kv_heads = heads + self.width = heads * head_dim + + def forward(self, x): + return torch.cat((x, x + 1, x + 2), dim=-1), None + + +class _AddNorm(torch.nn.Module): + def __init__(self, value: float) -> None: + super().__init__() + self.value = value + + def forward(self, x): + return x + self.value + + +class _CaptureAttention(torch.nn.Module): + def forward(self, query, key, value, metadata=None): + del key, value, metadata + return query + + +class _PassthroughLinear(torch.nn.Module): + def forward(self, x): + return x, None + + +def test_dual_attention_keeps_native_norm_and_text_image_concat_order(monkeypatch): + from vllm_omni.diffusion.models.longcat_image import longcat_image_transformer as longcat + + attention = longcat.LongCatImageAttention.__new__(longcat.LongCatImageAttention) + torch.nn.Module.__init__(attention) + attention.parallel_config = SimpleNamespace(sequence_parallel_size=1) + attention.head_dim = 4 + attention.added_kv_proj_dim = 8 + attention.to_qkv = _FakeQKV(2, 4) + attention.add_kv_proj = _FakeQKV(2, 4) + attention.norm_q = _AddNorm(10) + attention.norm_k = _AddNorm(20) + attention.norm_added_q = _AddNorm(100) + attention.norm_added_k = _AddNorm(200) + attention.attn = _CaptureAttention() + attention.to_out = _PassthroughLinear() + attention.to_add_out = _PassthroughLinear() + + captured = [] + + def paired(query, key, rotary_emb, sp_size, min_tokens): + captured.append((query.clone(), key.clone(), rotary_emb, sp_size, min_tokens)) + return query, key + + monkeypatch.setattr(longcat, "_apply_qk_rope", paired) + image = torch.randn(1, 3, 8) + text = torch.randn(1, 2, 8) + rotary_pair = (torch.randn(5, 4), torch.randn(5, 4)) + rotary_emb = (*rotary_pair, 0) + image_out, text_out = attention(image, text, rotary_emb) + + image_q = image.unflatten(-1, (2, 4)) + text_q = text.unflatten(-1, (2, 4)) + expected_q = torch.cat((text_q + 100, image_q + 10), dim=1) + expected_k = torch.cat((text_q + 1 + 200, image_q + 1 + 20), dim=1) + assert len(captured) == 1 + assert torch.equal(captured[0][0], expected_q) + assert torch.equal(captured[0][1], expected_k) + assert captured[0][2][0] is rotary_pair[0] + assert captured[0][2][1] is rotary_pair[1] + assert captured[0][3:] == (1, 0) + assert image_out.shape == image.shape + assert text_out.shape == text.shape diff --git a/vllm_omni/diffusion/layers/fused_qk_rope.py b/vllm_omni/diffusion/layers/fused_qk_rope.py new file mode 100644 index 00000000000..95b24451b12 --- /dev/null +++ b/vllm_omni/diffusion/layers/fused_qk_rope.py @@ -0,0 +1,251 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project + +"""Bit-exact paired Q/K full-width interleaved RoPE for CUDA. + +The operator accepts already-normalized BF16 query and key tensors in BSND +layout and full-width FP32 cosine/sine tables. It deliberately performs two +round-to-nearest FP32 multiplies followed by one round-to-nearest FP32 add, +then rounds only once when storing BF16 output. This matches Diffusers' +``apply_rotary_emb`` arithmetic while combining its separate Q and K calls. +""" + +from __future__ import annotations + +import torch +from torch.library import Library +from vllm.platforms import current_platform +from vllm.triton_utils import HAS_TRITON, tl, triton +from vllm.utils.torch_utils import direct_register_custom_op + +_HEAD_DIM = 128 +_HEADS_PER_PROGRAM = 4 + + +if HAS_TRITON: + + @triton.jit + def _mul_rn_f32(x, y): + return tl.inline_asm_elementwise( + asm="mul.rn.f32 $0, $1, $2;", + constraints="=f,f,f", + args=[x, y], + dtype=tl.float32, + is_pure=True, + pack=1, + ) + + @triton.jit + def _add_rn_f32(x, y): + return tl.inline_asm_elementwise( + asm="add.rn.f32 $0, $1, $2;", + constraints="=f,f,f", + args=[x, y], + dtype=tl.float32, + is_pure=True, + pack=1, + ) + + @triton.jit + def _qk_rope_exact_kernel( + q_ptr, + k_ptr, + cos_ptr, + sin_ptr, + q_output_ptr, + k_output_ptr, + q_stride_token, + q_stride_head, + q_stride_dim, + k_stride_token, + k_stride_head, + k_stride_dim, + q_output_stride_token, + q_output_stride_head, + q_output_stride_dim, + k_output_stride_token, + k_output_stride_head, + k_output_stride_dim, + table_stride_sequence, + table_stride_dim, + sequence, + num_heads: tl.constexpr, + head_dim: tl.constexpr, + heads_per_program: tl.constexpr, + ): + token = tl.program_id(0) + group = tl.program_id(1) + is_key = tl.program_id(2) == 1 + heads = group * heads_per_program + tl.arange(0, heads_per_program) + dims = tl.arange(0, head_dim) + valid = heads[:, None] < num_heads + + q_offsets = token * q_stride_token + heads[:, None] * q_stride_head + dims[None, :] * q_stride_dim + k_offsets = token * k_stride_token + heads[:, None] * k_stride_head + dims[None, :] * k_stride_dim + input_ptrs = tl.where(is_key, k_ptr + k_offsets, q_ptr + q_offsets) + + pair_dims = dims ^ 1 + q_pair_offsets = token * q_stride_token + heads[:, None] * q_stride_head + pair_dims[None, :] * q_stride_dim + k_pair_offsets = token * k_stride_token + heads[:, None] * k_stride_head + pair_dims[None, :] * k_stride_dim + pair_ptrs = tl.where(is_key, k_ptr + k_pair_offsets, q_ptr + q_pair_offsets) + values = tl.load(input_ptrs, mask=valid, other=0.0).to(tl.float32) + pairs = tl.load(pair_ptrs, mask=valid, other=0.0).to(tl.float32) + signed_pairs = tl.where((dims[None, :] & 1) == 0, -pairs, pairs) + + sequence_index = token % sequence + table_offsets = sequence_index * table_stride_sequence + dims * table_stride_dim + cos = tl.load(cos_ptr + table_offsets).to(tl.float32)[None, :] + sin = tl.load(sin_ptr + table_offsets).to(tl.float32)[None, :] + + # Do not let the compiler contract this expression into an FMA. The + # reference has two independently rounded FP32 multiplies and one + # independently rounded FP32 add, followed by the final BF16 store. + first = _mul_rn_f32(values, cos) + second = _mul_rn_f32(signed_pairs, sin) + output = _add_rn_f32(first, second) + q_output_offsets = ( + token * q_output_stride_token + heads[:, None] * q_output_stride_head + dims[None, :] * q_output_stride_dim + ) + k_output_offsets = ( + token * k_output_stride_token + heads[:, None] * k_output_stride_head + dims[None, :] * k_output_stride_dim + ) + output_ptrs = tl.where( + is_key, + k_output_ptr + k_output_offsets, + q_output_ptr + q_output_offsets, + ) + tl.store(output_ptrs, output, mask=valid) + + +def fused_qk_rope_supported( + q: torch.Tensor, + k: torch.Tensor, + cos: torch.Tensor, + sin: torch.Tensor, +) -> bool: + """Return whether the exact CUDA kernel can consume these tensors.""" + + if not ( + HAS_TRITON + and torch.version.hip is None + and current_platform.is_cuda() + and not torch.is_grad_enabled() + and q.is_cuda + and k.is_cuda + and q.device == k.device + and q.dtype is torch.bfloat16 + and k.dtype is q.dtype + and q.ndim == 4 + and k.shape == q.shape + and q.numel() > 0 + and q.is_contiguous() + and k.is_contiguous() + ): + return False + + batch, sequence, _heads, head_dim = q.shape + del batch + if head_dim != _HEAD_DIM: + return False + if not ( + cos.shape == (sequence, head_dim) + and sin.shape == cos.shape + and cos.device == q.device + and sin.device == q.device + and cos.dtype is torch.float32 + and sin.dtype is torch.float32 + and cos.is_contiguous() + and sin.is_contiguous() + ): + return False + return True + + +def _launch_fused_qk_rope( + q: torch.Tensor, + k: torch.Tensor, + cos: torch.Tensor, + sin: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + q_output = torch.empty_like(q) + k_output = torch.empty_like(k) + tokens = q.shape[0] * q.shape[1] + head_groups = triton.cdiv(q.shape[2], _HEADS_PER_PROGRAM) + grid = (tokens, head_groups, 2) + with torch.accelerator.device_index(q.device.index): + _qk_rope_exact_kernel[grid]( + q, + k, + cos, + sin, + q_output, + k_output, + q.stride(1), + q.stride(2), + q.stride(3), + k.stride(1), + k.stride(2), + k.stride(3), + q_output.stride(1), + q_output.stride(2), + q_output.stride(3), + k_output.stride(1), + k_output.stride(2), + k_output.stride(3), + cos.stride(0), + cos.stride(1), + q.shape[1], + num_heads=q.shape[2], + head_dim=q.shape[-1], + heads_per_program=_HEADS_PER_PROGRAM, + num_warps=4, + ) + return q_output, k_output + + +def _fused_qk_rope_impl( + q: torch.Tensor, + k: torch.Tensor, + cos: torch.Tensor, + sin: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + if not fused_qk_rope_supported(q, k, cos, sin): + raise ValueError("fused_qk_rope received unsupported inputs") + return _launch_fused_qk_rope(q, k, cos, sin) + + +def _fused_qk_rope_fake( + q: torch.Tensor, + k: torch.Tensor, + cos: torch.Tensor, + sin: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + del cos, sin + return torch.empty_like(q), torch.empty_like(k) + + +_OMNI_OP_LIB = Library("vllm_omni", "FRAGMENT") +if not hasattr(torch.ops.vllm_omni, "fused_qk_rope"): + direct_register_custom_op( + op_name="fused_qk_rope", + op_func=_fused_qk_rope_impl, + fake_impl=_fused_qk_rope_fake, + mutates_args=[], + target_lib=_OMNI_OP_LIB, + ) + + +def fused_qk_rope( + q: torch.Tensor, + k: torch.Tensor, + cos: torch.Tensor, + sin: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + """Apply paired full-width adjacent/interleaved RoPE on CUDA.""" + + if not fused_qk_rope_supported(q, k, cos, sin): + raise ValueError("fused_qk_rope requires contiguous CUDA BF16 Q/K and contiguous full-width CUDA FP32 tables") + return torch.ops.vllm_omni.fused_qk_rope(q, k, cos, sin) + + +__all__ = ["fused_qk_rope", "fused_qk_rope_supported"] diff --git a/vllm_omni/diffusion/models/longcat_image/longcat_image_transformer.py b/vllm_omni/diffusion/models/longcat_image/longcat_image_transformer.py index 3cde2a868a3..2593baf94d7 100644 --- a/vllm_omni/diffusion/models/longcat_image/longcat_image_transformer.py +++ b/vllm_omni/diffusion/models/longcat_image/longcat_image_transformer.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project from collections.abc import Iterable from typing import Any @@ -16,6 +16,8 @@ from vllm.model_executor.layers.layernorm import RMSNorm from vllm.model_executor.layers.linear import ColumnParallelLinear, QKVParallelLinear, RowParallelLinear from vllm.model_executor.model_loader.weight_utils import default_weight_loader +from vllm.platforms import current_platform +from vllm.triton_utils import HAS_TRITON from vllm_omni.diffusion.attention.backends.abstract import AttentionMetadata from vllm_omni.diffusion.attention.layer import Attention @@ -26,10 +28,151 @@ SequenceParallelOutput, ) from vllm_omni.diffusion.forward_context import get_forward_context +from vllm_omni.diffusion.layers.fused_qk_norm_rope import fused_qk_norm_rope_min_tokens +from vllm_omni.diffusion.layers.fused_qk_rope import fused_qk_rope, fused_qk_rope_supported from vllm_omni.platforms import current_omni_platform logger = init_logger(__name__) +RotaryEmbedding = tuple[torch.Tensor, torch.Tensor] | tuple[torch.Tensor, torch.Tensor, int] + +# All LongCat production attention spans contain at least 512 tokens. Reuse +# the shared Q/K fusion threshold override so deployments can retain a +# hardware-specific crossover without adding another environment variable. +_FUSED_MIN_TOKENS = 512 +_FUSED_QK_ROPE = HAS_TRITON and current_platform.is_cuda() +_VERIFIED_QK_ROPE_SIGNATURES: set[tuple] = set() +_FAILED_QK_ROPE_SIGNATURES: set[tuple] = set() + + +def _fusion_enabled(sequence_parallel_size: int | None, *, enforce_eager: bool) -> bool: + return ( + _FUSED_QK_ROPE + and enforce_eager + and not torch.compiler.is_compiling() + and not torch.is_grad_enabled() + and not (sequence_parallel_size is not None and sequence_parallel_size > 1) + ) + + +def _prepare_rotary_emb( + txt_cos: torch.Tensor, + txt_sin: torch.Tensor, + img_cos: torch.Tensor, + img_sin: torch.Tensor, + *, + enable_fusion: bool, +) -> RotaryEmbedding: + """Join text/image tables and resolve the shared token gate once.""" + + joint = torch.cat((txt_cos, img_cos), dim=0), torch.cat((txt_sin, img_sin), dim=0) + if not enable_fusion: + return joint + return joint[0], joint[1], fused_qk_norm_rope_min_tokens(_FUSED_MIN_TOKENS) + + +def _apply_qk_rope_reference( + query: torch.Tensor, + key: torch.Tensor, + rotary_emb: tuple[torch.Tensor, torch.Tensor], +) -> tuple[torch.Tensor, torch.Tensor]: + return ( + apply_rotary_emb(query, rotary_emb, sequence_dim=1), + apply_rotary_emb(key, rotary_emb, sequence_dim=1), + ) + + +def _qk_rope_signature( + query: torch.Tensor, + key: torch.Tensor, + cos: torch.Tensor, + sin: torch.Tensor, +) -> tuple: + """Identify every layout whose full output needs one exactness check.""" + + return ( + query.device.type, + query.device.index, + query.dtype, + tuple(query.shape), + tuple(query.stride()), + tuple(key.shape), + tuple(key.stride()), + cos.dtype, + tuple(cos.shape), + tuple(cos.stride()), + tuple(sin.stride()), + ) + + +def _can_use_fused_qk_rope( + query: torch.Tensor, + key: torch.Tensor, + cos: torch.Tensor, + sin: torch.Tensor, + sequence_parallel_size: int | None, + min_tokens: int | None, +) -> bool: + """Keep compile, autograd, capture, SP, and unsupported inputs native.""" + + if ( + torch.compiler.is_compiling() + or torch.is_grad_enabled() + or (sequence_parallel_size is not None and sequence_parallel_size > 1) + or min_tokens is None + or not fused_qk_rope_supported(query, key, cos, sin) + or torch.cuda.is_current_stream_capturing() + ): + return False + return query.shape[0] * query.shape[1] >= min_tokens + + +def _apply_qk_rope( + query: torch.Tensor, + key: torch.Tensor, + rotary_emb: RotaryEmbedding | None, + sequence_parallel_size: int | None, + min_tokens: int | None, +) -> tuple[torch.Tensor, torch.Tensor]: + """Apply one paired RoPE launch, self-verifying each runtime signature.""" + + if rotary_emb is None: + return query, key + rotary_pair = rotary_emb[:2] + cos, sin = rotary_pair + if not _can_use_fused_qk_rope(query, key, cos, sin, sequence_parallel_size, min_tokens): + return _apply_qk_rope_reference(query, key, rotary_pair) + + signature = _qk_rope_signature(query, key, cos, sin) + if signature in _FAILED_QK_ROPE_SIGNATURES: + return _apply_qk_rope_reference(query, key, rotary_pair) + + try: + output = fused_qk_rope(query, key, cos, sin) + except Exception as exc: # noqa: BLE001 - optimized-path failures must fall back + _FAILED_QK_ROPE_SIGNATURES.add(signature) + logger.warning( + "Disabling LongCat paired Q/K RoPE fusion for signature %s after failure: %s", + signature, + exc, + ) + return _apply_qk_rope_reference(query, key, rotary_pair) + + if signature in _VERIFIED_QK_ROPE_SIGNATURES: + return output + + reference = _apply_qk_rope_reference(query, key, rotary_pair) + if torch.equal(output[0], reference[0]) and torch.equal(output[1], reference[1]): + _VERIFIED_QK_ROPE_SIGNATURES.add(signature) + return output + + _FAILED_QK_ROPE_SIGNATURES.add(signature) + logger.warning( + "Disabling LongCat paired Q/K RoPE fusion for signature %s after a bit-exactness mismatch", + signature, + ) + return reference + class FeedForward(nn.Module): def __init__(self, dim: int, dim_out: int | None = None, mult: int = 4, bias: bool = True): @@ -121,7 +264,7 @@ def _sp_attention_with_rope( text_key: torch.Tensor, text_value: torch.Tensor, text_seq_len: int, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None, + image_rotary_emb: RotaryEmbedding | None, ) -> torch.Tensor: """ Apply RoPE separately to text and image Q/K, then run SP attention with joint tensors. @@ -139,7 +282,7 @@ def _sp_attention_with_rope( Attention output with shape (B, txt_len + img_len/SP, H, D) """ if image_rotary_emb is not None: - freqs_cos, freqs_sin = image_rotary_emb + freqs_cos, freqs_sin = image_rotary_emb[:2] txt_rotary_emb = (freqs_cos[:text_seq_len], freqs_sin[:text_seq_len]) img_rotary_emb_split = (freqs_cos[text_seq_len:], freqs_sin[text_seq_len:]) # Apply RoPE to image Q/K @@ -165,7 +308,7 @@ def forward( self, hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor | None = None, - image_rotary_emb: torch.Tensor | None = None, + image_rotary_emb: RotaryEmbedding | None = None, **kwargs, ) -> torch.Tensor: """ @@ -196,6 +339,8 @@ def forward( query = self.norm_q(query) key = self.norm_k(key) + rotary_pair = image_rotary_emb[:2] if image_rotary_emb is not None else None + min_tokens = image_rotary_emb[2] if image_rotary_emb is not None and len(image_rotary_emb) > 2 else None if self.added_kv_proj_dim is not None: encoder_qkv, _ = self.add_kv_proj(encoder_hidden_states) @@ -222,7 +367,7 @@ def forward( text_key=encoder_key, text_value=encoder_value, text_seq_len=encoder_query.shape[1], - image_rotary_emb=image_rotary_emb, + image_rotary_emb=rotary_pair, ) else: # Non-SP Mode: Concat first, then apply RoPE to full sequence @@ -230,10 +375,15 @@ def forward( joint_key = torch.cat([encoder_key, key], dim=1) joint_value = torch.cat([encoder_value, value], dim=1) - if image_rotary_emb is not None: - # Apply RoPE to full (text + image) sequence - joint_query = apply_rotary_emb(joint_query, image_rotary_emb, sequence_dim=1) - joint_key = apply_rotary_emb(joint_key, image_rotary_emb, sequence_dim=1) + # Keep the native RMSNorm and text-then-image concatenation + # order, then combine only the two independent RoPE calls. + joint_query, joint_key = _apply_qk_rope( + joint_query, + joint_key, + rotary_pair, + sp_size, + min_tokens, + ) hidden_states = self.attn( joint_query, @@ -272,13 +422,17 @@ def forward( text_key=key[:, :text_seq_len], text_value=value[:, :text_seq_len], text_seq_len=text_seq_len, - image_rotary_emb=image_rotary_emb, + image_rotary_emb=rotary_pair, ) else: # Non-SP Mode: standard path - if image_rotary_emb is not None: - query = apply_rotary_emb(query, image_rotary_emb, sequence_dim=1) - key = apply_rotary_emb(key, image_rotary_emb, sequence_dim=1) + query, key = _apply_qk_rope( + query, + key, + rotary_pair, + sp_size, + min_tokens, + ) hidden_states = self.attn( query, @@ -345,7 +499,7 @@ def forward( hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor, temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, + image_rotary_emb: RotaryEmbedding | None = None, joint_attention_kwargs: dict[str, Any] | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb) @@ -536,7 +690,7 @@ def forward( hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor, temb: torch.Tensor, - image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None, + image_rotary_emb: RotaryEmbedding | None = None, joint_attention_kwargs: dict[str, Any] | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """ @@ -632,6 +786,7 @@ def __init__( # Store parallel config for SP support self.parallel_config = od_config.parallel_config + self.enforce_eager = bool(getattr(od_config, "enforce_eager", False)) self.pos_embed = LongCatImagePosEmbed(theta=10000, axes_dim=axes_dims_rope) self.rope_preparer = RoPEPreparer(self.pos_embed) @@ -701,11 +856,15 @@ def forward( # txt_cos/txt_sin (outputs 0, 1) remain replicated for dual-stream attention txt_cos, txt_sin, img_cos, img_sin = self.rope_preparer(txt_ids, img_ids) - # Reconstruct image_rotary_emb with chunked values - # Final shape: (txt_seq_len + img_seq_len // SP, head_dim) - image_rotary_emb = ( - torch.cat([txt_cos, img_cos], dim=0), - torch.cat([txt_sin, img_sin], dim=0), + # Preserve the ordinary full-width tables. A third scalar only marks + # the eligible eager SP=1 route and resolves the shared threshold once + # for all attention blocks in this transformer call. + image_rotary_emb = _prepare_rotary_emb( + txt_cos, + txt_sin, + img_cos, + img_sin, + enable_fusion=_fusion_enabled(sp_size, enforce_eager=self.enforce_eager), ) for block in self.transformer_blocks: From dec1be867737ddc574a79ba860b6bbc384c3d452 Mon Sep 17 00:00:00 2001 From: dongbo910220 <1275604947@qq.com> Date: Tue, 15 Sep 2026 21:58:41 +0800 Subject: [PATCH 2/2] fix: remove runtime LongCat RoPE verification Signed-off-by: dongbo910220 <1275604947@qq.com> --- .../test_longcat_image_transformer.py | 60 +++++++++---------- .../longcat_image_transformer.py | 46 +++++++------- 2 files changed, 50 insertions(+), 56 deletions(-) diff --git a/tests/diffusion/models/longcat_image/test_longcat_image_transformer.py b/tests/diffusion/models/longcat_image/test_longcat_image_transformer.py index 8e105479fa7..51cedbb8f00 100644 --- a/tests/diffusion/models/longcat_image/test_longcat_image_transformer.py +++ b/tests/diffusion/models/longcat_image/test_longcat_image_transformer.py @@ -15,10 +15,8 @@ def _clear_qk_rope_signature_caches(): from vllm_omni.diffusion.models.longcat_image import longcat_image_transformer as longcat - longcat._VERIFIED_QK_ROPE_SIGNATURES.clear() longcat._FAILED_QK_ROPE_SIGNATURES.clear() yield - longcat._VERIFIED_QK_ROPE_SIGNATURES.clear() longcat._FAILED_QK_ROPE_SIGNATURES.clear() @@ -96,49 +94,37 @@ def test_qk_rope_non_eager_modes_use_exact_original_path(monkeypatch, mode, sp_s assert torch.equal(actual[1], expected[1]) -def test_qk_rope_verifies_full_output_once_per_signature(monkeypatch): +def test_qk_rope_eligible_path_skips_runtime_parity_check(monkeypatch): from vllm_omni.diffusion.models.longcat_image import longcat_image_transformer as longcat torch.manual_seed(2) monkeypatch.setattr(torch.compiler, "is_compiling", lambda: False) monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: False) monkeypatch.setattr(longcat, "fused_qk_rope_supported", lambda *args: True) - - raw_reference = longcat._apply_qk_rope_reference - reference_calls = 0 fused_calls = 0 - def tracked_reference(*args): - nonlocal reference_calls - reference_calls += 1 - return raw_reference(*args) - - def exact_fused(query, key, cos, sin): + def fused(query, key, cos, sin): nonlocal fused_calls fused_calls += 1 - return raw_reference(query, key, (cos, sin)) + return query + 1, key + 2 - monkeypatch.setattr(longcat, "_apply_qk_rope_reference", tracked_reference) - monkeypatch.setattr(longcat, "fused_qk_rope", exact_fused) + monkeypatch.setattr( + longcat, + "_apply_qk_rope_reference", + lambda *args: (_ for _ in ()).throw(AssertionError("eligible path ran the native parity check")), + ) + monkeypatch.setattr(longcat, "fused_qk_rope", fused) with torch.no_grad(): - q, k, rotary_emb = _inputs(sequence=5) - first = longcat._apply_qk_rope(q, k, rotary_emb, 1, 0) - second = longcat._apply_qk_rope(q.clone(), k.clone(), rotary_emb, 1, 0) - q_other, k_other, rotary_other = _inputs(sequence=7) - third = longcat._apply_qk_rope(q_other, k_other, rotary_other, 1, 0) + q, k, rotary_emb = _inputs() + actual = longcat._apply_qk_rope(q, k, rotary_emb, 1, 0) - assert torch.equal(first[0], _reference(q, k, rotary_emb)[0]) - assert torch.equal(second[1], _reference(q, k, rotary_emb)[1]) - assert torch.equal(third[0], _reference(q_other, k_other, rotary_other)[0]) - assert fused_calls == 3 - assert reference_calls == 2 - assert len(longcat._VERIFIED_QK_ROPE_SIGNATURES) == 2 - assert not longcat._FAILED_QK_ROPE_SIGNATURES + assert fused_calls == 1 + assert torch.equal(actual[0], q + 1) + assert torch.equal(actual[1], k + 2) -@pytest.mark.parametrize("failure", ["mismatch", "exception"]) -def test_qk_rope_first_use_failure_permanently_falls_back(monkeypatch, failure): +def test_qk_rope_launch_failure_permanently_falls_back(monkeypatch): from vllm_omni.diffusion.models.longcat_image import longcat_image_transformer as longcat torch.manual_seed(3) @@ -150,9 +136,7 @@ def test_qk_rope_first_use_failure_permanently_falls_back(monkeypatch, failure): def broken_fused(query, key, cos, sin): nonlocal fused_calls fused_calls += 1 - if failure == "exception": - raise RuntimeError("kernel failed") - return torch.zeros_like(query), torch.zeros_like(key) + raise RuntimeError("kernel failed") monkeypatch.setattr(longcat, "fused_qk_rope", broken_fused) q, k, rotary_emb = _inputs() @@ -163,12 +147,22 @@ def broken_fused(query, key, cos, sin): assert fused_calls == 1 assert len(longcat._FAILED_QK_ROPE_SIGNATURES) == 1 - assert not longcat._VERIFIED_QK_ROPE_SIGNATURES for actual in (first, second): assert torch.equal(actual[0], expected[0]) assert torch.equal(actual[1], expected[1]) +def test_failed_qk_rope_signature_cache_is_bounded(): + from vllm_omni.diffusion.models.longcat_image import longcat_image_transformer as longcat + + for index in range(longcat._FAILED_QK_ROPE_SIGNATURES_MAX_SIZE + 1): + longcat._record_failed_qk_rope_signature((index,)) + + assert len(longcat._FAILED_QK_ROPE_SIGNATURES) == longcat._FAILED_QK_ROPE_SIGNATURES_MAX_SIZE + assert (0,) not in longcat._FAILED_QK_ROPE_SIGNATURES + assert (longcat._FAILED_QK_ROPE_SIGNATURES_MAX_SIZE,) in longcat._FAILED_QK_ROPE_SIGNATURES + + class _FakeQKV(torch.nn.Module): def __init__(self, heads: int, head_dim: int) -> None: super().__init__() diff --git a/vllm_omni/diffusion/models/longcat_image/longcat_image_transformer.py b/vllm_omni/diffusion/models/longcat_image/longcat_image_transformer.py index 2593baf94d7..a93d147047a 100644 --- a/vllm_omni/diffusion/models/longcat_image/longcat_image_transformer.py +++ b/vllm_omni/diffusion/models/longcat_image/longcat_image_transformer.py @@ -1,6 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project +from collections import OrderedDict from collections.abc import Iterable from typing import Any @@ -41,8 +42,8 @@ # hardware-specific crossover without adding another environment variable. _FUSED_MIN_TOKENS = 512 _FUSED_QK_ROPE = HAS_TRITON and current_platform.is_cuda() -_VERIFIED_QK_ROPE_SIGNATURES: set[tuple] = set() -_FAILED_QK_ROPE_SIGNATURES: set[tuple] = set() +_FAILED_QK_ROPE_SIGNATURES_MAX_SIZE = 128 +_FAILED_QK_ROPE_SIGNATURES: OrderedDict[tuple[object, ...], None] = OrderedDict() def _fusion_enabled(sequence_parallel_size: int | None, *, enforce_eager: bool) -> bool: @@ -87,8 +88,8 @@ def _qk_rope_signature( key: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, -) -> tuple: - """Identify every layout whose full output needs one exactness check.""" +) -> tuple[object, ...]: + """Identify a runtime layout for failure-cache fallback.""" return ( query.device.type, @@ -105,6 +106,20 @@ def _qk_rope_signature( ) +def _is_failed_qk_rope_signature(signature: tuple[object, ...]) -> bool: + if signature not in _FAILED_QK_ROPE_SIGNATURES: + return False + _FAILED_QK_ROPE_SIGNATURES.move_to_end(signature) + return True + + +def _record_failed_qk_rope_signature(signature: tuple[object, ...]) -> None: + _FAILED_QK_ROPE_SIGNATURES[signature] = None + _FAILED_QK_ROPE_SIGNATURES.move_to_end(signature) + while len(_FAILED_QK_ROPE_SIGNATURES) > _FAILED_QK_ROPE_SIGNATURES_MAX_SIZE: + _FAILED_QK_ROPE_SIGNATURES.popitem(last=False) + + def _can_use_fused_qk_rope( query: torch.Tensor, key: torch.Tensor, @@ -134,7 +149,7 @@ def _apply_qk_rope( sequence_parallel_size: int | None, min_tokens: int | None, ) -> tuple[torch.Tensor, torch.Tensor]: - """Apply one paired RoPE launch, self-verifying each runtime signature.""" + """Apply one paired RoPE launch, falling back after launch failures.""" if rotary_emb is None: return query, key @@ -144,13 +159,13 @@ def _apply_qk_rope( return _apply_qk_rope_reference(query, key, rotary_pair) signature = _qk_rope_signature(query, key, cos, sin) - if signature in _FAILED_QK_ROPE_SIGNATURES: + if _is_failed_qk_rope_signature(signature): return _apply_qk_rope_reference(query, key, rotary_pair) try: - output = fused_qk_rope(query, key, cos, sin) + return fused_qk_rope(query, key, cos, sin) except Exception as exc: # noqa: BLE001 - optimized-path failures must fall back - _FAILED_QK_ROPE_SIGNATURES.add(signature) + _record_failed_qk_rope_signature(signature) logger.warning( "Disabling LongCat paired Q/K RoPE fusion for signature %s after failure: %s", signature, @@ -158,21 +173,6 @@ def _apply_qk_rope( ) return _apply_qk_rope_reference(query, key, rotary_pair) - if signature in _VERIFIED_QK_ROPE_SIGNATURES: - return output - - reference = _apply_qk_rope_reference(query, key, rotary_pair) - if torch.equal(output[0], reference[0]) and torch.equal(output[1], reference[1]): - _VERIFIED_QK_ROPE_SIGNATURES.add(signature) - return output - - _FAILED_QK_ROPE_SIGNATURES.add(signature) - logger.warning( - "Disabling LongCat paired Q/K RoPE fusion for signature %s after a bit-exactness mismatch", - signature, - ) - return reference - class FeedForward(nn.Module): def __init__(self, dim: int, dim_out: int | None = None, mult: int = 4, bias: bool = True):