From 6ae7d1d96805deb62b2f32ac88ace43d0c01027e Mon Sep 17 00:00:00 2001 From: FNU AKSHANSH <105249360+akshansh47@users.noreply.github.com> Date: Wed, 5 Aug 2026 10:11:04 -0700 Subject: [PATCH] Cast tl.program_id(0) to int64 in llama4_rope and qwen2vl_mrope Same int32 overflow hazard fixed for rope.py / rms_norm.py in #804. llama4_rope and qwen2vl_mrope missed the cast; Ascend llama4_rope already has it. Fixes #1335. Co-authored-by: Cursor --- src/liger_kernel/ops/llama4_rope.py | 2 +- src/liger_kernel/ops/qwen2vl_mrope.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/src/liger_kernel/ops/llama4_rope.py b/src/liger_kernel/ops/llama4_rope.py index 9167d69f3..02a8647be 100644 --- a/src/liger_kernel/ops/llama4_rope.py +++ b/src/liger_kernel/ops/llama4_rope.py @@ -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 diff --git a/src/liger_kernel/ops/qwen2vl_mrope.py b/src/liger_kernel/ops/qwen2vl_mrope.py index fbd120f96..b431c8976 100644 --- a/src/liger_kernel/ops/qwen2vl_mrope.py +++ b/src/liger_kernel/ops/qwen2vl_mrope.py @@ -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)