Skip to content
Open
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
12 changes: 12 additions & 0 deletions megatron/core/fusions/fused_mla_yarn_rope_apply.py
Original file line number Diff line number Diff line change
Expand Up @@ -253,6 +253,12 @@ def forward(
assert sin.is_contiguous()
assert headdim == qk_head_dim + emb_dim
assert emb_dim % 4 == 0
if cos.shape[-1] != emb_dim or sin.shape[-1] != emb_dim:
raise ValueError(
f"cos/sin last dim must equal emb_dim={emb_dim} "
f"(got cos={cos.shape[-1]}, sin={sin.shape[-1]}); a narrower rotary cache "
f"(e.g. rotary_percent < 1) makes the fused MLA kernel read past the buffer."
)

grid = lambda META: (total_seqlen, triton.cdiv(nheads, META["BLOCK_H"]))
rotary_fwd_q_kernel[grid](
Expand Down Expand Up @@ -629,6 +635,12 @@ def forward(
assert cos.is_contiguous()
assert sin.is_contiguous()
assert emb_dim % 4 == 0
if cos.shape[-1] != emb_dim or sin.shape[-1] != emb_dim:
raise ValueError(
f"cos/sin last dim must equal emb_dim={emb_dim} "
f"(got cos={cos.shape[-1]}, sin={sin.shape[-1]}); a narrower rotary cache "
f"(e.g. rotary_percent < 1) makes the fused MLA kernel read past the buffer."
)

o_key = kv.new_empty(total_seqlen, nheads, emb_dim + k_dim)
o_value = kv.new_empty(total_seqlen, nheads, v_dim)
Expand Down
48 changes: 48 additions & 0 deletions tests/unit_tests/fusions/test_mla_yarn_rope_apply.py
Original file line number Diff line number Diff line change
Expand Up @@ -364,3 +364,51 @@ def test_mla_rotary_interleaved_with_apply_rope_fusion_emits_warning_and_uses_un
call_kw = unfused_spy.call_args[1]
assert call_kw["mla_rotary_interleaved"] is True
assert out.shape == t.shape


class TestFusedApplyMLARopeCosWidthGuard:
"""The fused MLA RoPE kernels read ``emb_dim`` cos/sin values per token and assume the
cache is ``emb_dim`` wide. A narrower cache (e.g. ``rotary_percent < 1`` shrinks it to
``int(emb_dim * rotary_percent)``) makes the kernel read past the buffer and return
garbage. These tests assert the guard rejects that up front, before the kernel launches,
so they do not require CUDA."""

@pytest.mark.skipif(
fused_apply_mla_rope_for_q is None, reason="fused MLA RoPE kernels unavailable"
)
def test_q_rejects_narrow_cos_sin(self):
qk_head_dim = 128
emb_dim = 64
num_heads = 4
seqlen = 8
batch_size = 2
narrow = emb_dim // 8 # mimics rotary_percent=0.125 -> int(64 * 0.125) = 8

q = torch.randn(seqlen, batch_size, num_heads, qk_head_dim + emb_dim)
cos = torch.randn(seqlen, 1, 1, narrow)
sin = torch.randn(seqlen, 1, 1, narrow)

with pytest.raises(ValueError, match="cos/sin last dim"):
fused_apply_mla_rope_for_q(q, cos, sin, qk_head_dim, emb_dim, cu_seqlens_q=None)

@pytest.mark.skipif(
fused_apply_mla_rope_for_kv is None, reason="fused MLA RoPE kernels unavailable"
)
def test_kv_rejects_narrow_cos_sin(self):
emb_dim = 64
k_dim = 128
v_dim = 128
num_heads = 4
seqlen = 8
batch_size = 2
narrow = emb_dim // 8

kv = torch.randn(seqlen, batch_size, num_heads, k_dim + v_dim)
k_pos_emb = torch.randn(seqlen, batch_size, 1, emb_dim)
cos = torch.randn(seqlen, 1, 1, narrow)
sin = torch.randn(seqlen, 1, 1, narrow)

with pytest.raises(ValueError, match="cos/sin last dim"):
fused_apply_mla_rope_for_kv(
kv, k_pos_emb, cos, sin, emb_dim, k_dim, v_dim, cu_seqlens_kv=None
)