Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
125 changes: 125 additions & 0 deletions tests/diffusion/layers/test_fused_qk_rope.py
Original file line number Diff line number Diff line change
@@ -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])
Original file line number Diff line number Diff line change
@@ -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
Loading
Loading