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..51cedbb8f00 --- /dev/null +++ b/tests/diffusion/models/longcat_image/test_longcat_image_transformer.py @@ -0,0 +1,239 @@ +# 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._FAILED_QK_ROPE_SIGNATURES.clear() + yield + 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_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) + fused_calls = 0 + + def fused(query, key, cos, sin): + nonlocal fused_calls + fused_calls += 1 + return query + 1, key + 2 + + 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() + actual = longcat._apply_qk_rope(q, k, rotary_emb, 1, 0) + + assert fused_calls == 1 + assert torch.equal(actual[0], q + 1) + assert torch.equal(actual[1], k + 2) + + +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) + 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 + raise RuntimeError("kernel failed") + + 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 + 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__() + 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..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 project +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project +from collections import OrderedDict from collections.abc import Iterable from typing import Any @@ -16,6 +17,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 +29,150 @@ 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() +_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: + 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[object, ...]: + """Identify a runtime layout for failure-cache fallback.""" + + 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 _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, + 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, falling back after launch failures.""" + + 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 _is_failed_qk_rope_signature(signature): + return _apply_qk_rope_reference(query, key, rotary_pair) + + try: + return fused_qk_rope(query, key, cos, sin) + except Exception as exc: # noqa: BLE001 - optimized-path failures must fall back + _record_failed_qk_rope_signature(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) + 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: