From 79d1286053f216a3d552e8804eb73b6fc2c8c6c7 Mon Sep 17 00:00:00 2001 From: Solaris-star <820622658@qq.com> Date: Mon, 20 Jul 2026 23:31:55 +0800 Subject: [PATCH] fix(triton): uncross Q/K halves in fused_rotary_emb out_q1 was computed from K and out_k0 from Q, so half of each tensor got the other tensor's rotary transform. Also use k_token_stride for off_k1 (was accidentally q_token_stride). Fixes #6428 Signed-off-by: Solaris-star <820622658@qq.com> --- colossalai/kernel/triton/fused_rotary_embedding.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/colossalai/kernel/triton/fused_rotary_embedding.py b/colossalai/kernel/triton/fused_rotary_embedding.py index cf2a70f7b64e..d4494a988b5b 100644 --- a/colossalai/kernel/triton/fused_rotary_embedding.py +++ b/colossalai/kernel/triton/fused_rotary_embedding.py @@ -59,7 +59,7 @@ def fused_rotary_emb( + dim_range0[None, None, :] * head_dim_stride ) off_k1 = ( - idx * q_token_stride + idx * k_token_stride + cur_head_range[None, :, None] * k_head_stride + dim_range1[None, None, :] * head_dim_stride ) @@ -88,10 +88,12 @@ def fused_rotary_emb( other=0.0, ) + # Standard rotary: [x0, x1] -> [x0*cos - x1*sin, x0*sin + x1*cos] + # Previous code cross-wired Q/K for out_q1 and out_k0 (see #6428). out_q0 = q_0 * cos - q_1 * sin - out_q1 = k_0 * sin + k_1 * cos + out_q1 = q_0 * sin + q_1 * cos - out_k0 = q_0 * cos - q_1 * sin + out_k0 = k_0 * cos - k_1 * sin out_k1 = k_0 * sin + k_1 * cos # concat tl.store(