Skip to content
2 changes: 1 addition & 1 deletion src/liger_kernel/ops/backends/_ascend/ops/qwen2vl_mrope.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,7 @@ def _triton_qwen2vl_mrope_npu(
actual_rows = tl.minimum(rows_per_program, total_rows - start_row)

for row_offset in tl.range(0, actual_rows):
pid = start_row + row_offset
pid = (start_row + row_offset).to(tl.int64)

t_end = mrope_section_t
h_end = t_end + mrope_section_h
Expand Down
2 changes: 1 addition & 1 deletion src/liger_kernel/ops/llama4_rope.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,7 @@ def _llama4_rope_kernel(
Grid: (batch*seq, head)
"""
# 2D grid
pid_bs = tl.program_id(0) # over batch*seq
pid_bs = tl.program_id(0).to(tl.int64) # over batch*seq
pid_h = tl.program_id(1) # over heads

batch_idx = pid_bs // seq_len
Expand Down
2 changes: 1 addition & 1 deletion src/liger_kernel/ops/qwen2vl_mrope.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,7 @@ def _triton_qwen2vl_mrope(
BLOCK_SIZE: tl.constexpr,
BACKWARD_PASS: tl.constexpr = False,
):
pid = tl.program_id(0)
pid = tl.program_id(0).to(tl.int64)

# locate start address
q_ptr = q_ptr + pid * (n_qh * hd)
Expand Down
70 changes: 70 additions & 0 deletions test/transformers/test_llama4_rope.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@

from liger_kernel.ops import LigerLlama4RopeFunction
from liger_kernel.transformers.llama4_rope import liger_llama4_text_rotary_pos_emb
from liger_kernel.utils import get_total_gpu_memory
from liger_kernel.utils import infer_device

try:
Expand Down Expand Up @@ -147,3 +148,72 @@ def test_functional_correctness(bsz, seq_len, num_q_heads, num_kv_heads, head_di

assert torch.allclose(q1_grad, q2_grad, atol=atol, rtol=rtol)
assert torch.allclose(k1_grad, k2_grad, atol=atol, rtol=rtol)


# The kernel indexes the flattened batch*seq dimension as
# `base_offset * q_row_stride`, whose largest value is
# `(bsz * seq_len - 1) * n_q_heads * head_dim`. That product is one row short of
# `q.numel()`, so it only exceeds int32 once q itself holds more than 2**31
# elements. There is no cheaper shape: the bound is the element count.
_HEAD_DIM = 128
_N_Q_HEADS = 64
_ROW = _N_Q_HEADS * _HEAD_DIM
_SEQ_LEN = (2**31 - 1) // _ROW + 2 # smallest seq_len whose last offset wraps int32


@pytest.mark.skipif(not IS_LLAMA4_AVAILABLE, reason="Llama4 is not available in transformers.")
@pytest.mark.skipif(not supports_bfloat16(), reason="bfloat16 not supported on this GPU")
@pytest.mark.parametrize(
"bsz, seq_len, num_q_heads, num_kv_heads, head_dim",
[
pytest.param(
1,
_SEQ_LEN,
_N_Q_HEADS,
1,
_HEAD_DIM,
marks=pytest.mark.skipif(
infer_device() == "cpu" or get_total_gpu_memory() < 20,
reason="This test requires a GPU with at least 20GB of memory",
),
),
],
)
def test_row_offset_does_not_wrap_int32(bsz, seq_len, num_q_heads, num_kv_heads, head_dim):
"""Regression for #1335: the last token must still be rotated, not read from a wrapped address."""
dtype = torch.bfloat16
n_kv_heads = num_kv_heads

q = torch.zeros((bsz, seq_len, num_q_heads, head_dim), device=device, dtype=dtype)
k = torch.zeros((bsz, seq_len, n_kv_heads, head_dim), device=device, dtype=dtype)

# Only the final token carries a signal. It sits at the offset that wraps.
q[0, -1].fill_(1.0)
k[0, -1].fill_(1.0)

config = Llama4TextConfig(
hidden_size=num_q_heads * head_dim,
num_attention_heads=num_q_heads,
num_key_value_heads=n_kv_heads,
head_dim=head_dim,
max_position_embeddings=seq_len,
rope_theta=10000.0,
rope_scaling=None,
)
rotary_emb = Llama4TextRotaryEmbedding(config=config, device=device)
pos_ids = torch.arange(seq_len, device=device).unsqueeze(0)
freqs_cis = rotary_emb(q, pos_ids)

q_out, k_out = LigerLlama4RopeFunction.apply(q, k, freqs_cis)

# A wrapped index leaves the last row untouched or fills it from elsewhere.
assert torch.isfinite(q_out[0, -1]).all()
assert torch.isfinite(k_out[0, -1]).all()
assert q_out[0, -1].abs().sum() > 0
assert k_out[0, -1].abs().sum() > 0

ref_q, ref_k = apply_rotary_emb(
q[:, -1:].float(), k[:, -1:].float(), freqs_cis[-1:].unsqueeze(0)
)
assert torch.allclose(q_out[:, -1:].float(), ref_q, atol=1e-1, rtol=1e-5)
assert torch.allclose(k_out[:, -1:].float(), ref_k, atol=1e-1, rtol=1e-5)