From 6f2b51ade1b1d072ee7e8d2727f1dbaec0f496ae Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <35585791+BBuf@users.noreply.github.com> Date: Thu, 26 Mar 2026 08:59:25 +0800 Subject: [PATCH] [Diffusion] Optimize diffusion Triton rotary embedding by processing multiple heads per token (#21387) --- .../jit_kernel/diffusion/triton/rotary.py | 72 ++++++++++++------- 1 file changed, 46 insertions(+), 26 deletions(-) diff --git a/python/sglang/jit_kernel/diffusion/triton/rotary.py b/python/sglang/jit_kernel/diffusion/triton/rotary.py index 02665fbaf..16dc6f61e 100644 --- a/python/sglang/jit_kernel/diffusion/triton/rotary.py +++ b/python/sglang/jit_kernel/diffusion/triton/rotary.py @@ -7,12 +7,13 @@ from sglang.multimodal_gen.runtime.platforms import current_platform @triton.autotune( configs=[ - triton.Config({"BLOCK_HS_HALF": 32}, num_warps=2), - triton.Config({"BLOCK_HS_HALF": 64}, num_warps=4), - triton.Config({"BLOCK_HS_HALF": 128}, num_warps=4), - triton.Config({"BLOCK_HS_HALF": 256}, num_warps=8), + triton.Config({"BLOCK_HEADS": 1, "BLOCK_HS_HALF": 32}, num_warps=2), + triton.Config({"BLOCK_HEADS": 2, "BLOCK_HS_HALF": 32}, num_warps=2), + triton.Config({"BLOCK_HEADS": 4, "BLOCK_HS_HALF": 32}, num_warps=4), + triton.Config({"BLOCK_HEADS": 4, "BLOCK_HS_HALF": 64}, num_warps=4), + triton.Config({"BLOCK_HEADS": 8, "BLOCK_HS_HALF": 64}, num_warps=8), ], - key=["head_size"], + key=["num_heads", "head_size"], ) @triton.jit def _rotary_embedding_kernel( @@ -23,44 +24,61 @@ def _rotary_embedding_kernel( num_heads, head_size, num_tokens, - stride_x_row, + stride_out_bt, + stride_out_head, + stride_x_bt, + stride_x_head, stride_cos_row, stride_sin_row, + BLOCK_HEADS: tl.constexpr, BLOCK_HS_HALF: tl.constexpr, ): - row_idx = tl.program_id(0) - token_idx = (row_idx // num_heads) % num_tokens + bt_idx = tl.program_id(0) + head_block_idx = tl.program_id(1) + token_idx = bt_idx % num_tokens - x_row_ptr = x_ptr + row_idx * stride_x_row cos_row_ptr = cos_ptr + token_idx * stride_cos_row sin_row_ptr = sin_ptr + token_idx * stride_sin_row - output_row_ptr = output_ptr + row_idx * stride_x_row + head_offsets = head_block_idx * BLOCK_HEADS + tl.arange(0, BLOCK_HEADS) + head_mask = head_offsets < num_heads - # half size for x1 and x2 head_size_half = head_size // 2 + x_row_ptrs = x_ptr + bt_idx * stride_x_bt + head_offsets[:, None] * stride_x_head + output_row_ptrs = ( + output_ptr + bt_idx * stride_out_bt + head_offsets[:, None] * stride_out_head + ) for block_start in range(0, head_size_half, BLOCK_HS_HALF): offsets_half = block_start + tl.arange(0, BLOCK_HS_HALF) - mask = offsets_half < head_size_half + half_mask = offsets_half < head_size_half + mask = head_mask[:, None] & half_mask[None, :] - cos_vals = tl.load(cos_row_ptr + offsets_half, mask=mask, other=0.0) - sin_vals = tl.load(sin_row_ptr + offsets_half, mask=mask, other=0.0) + cos_vals = tl.load(cos_row_ptr + offsets_half, mask=half_mask, other=0.0) + sin_vals = tl.load(sin_row_ptr + offsets_half, mask=half_mask, other=0.0) offsets_x1 = 2 * offsets_half offsets_x2 = 2 * offsets_half + 1 - x1_vals = tl.load(x_row_ptr + offsets_x1, mask=mask, other=0.0) - x2_vals = tl.load(x_row_ptr + offsets_x2, mask=mask, other=0.0) + x1_vals = tl.load(x_row_ptrs + offsets_x1[None, :], mask=mask, other=0.0) + x2_vals = tl.load(x_row_ptrs + offsets_x2[None, :], mask=mask, other=0.0) x1_fp32 = x1_vals.to(tl.float32) x2_fp32 = x2_vals.to(tl.float32) - cos_fp32 = cos_vals.to(tl.float32) - sin_fp32 = sin_vals.to(tl.float32) + cos_fp32 = cos_vals.to(tl.float32)[None, :] + sin_fp32 = sin_vals.to(tl.float32)[None, :] o1_vals = tl.fma(-x2_fp32, sin_fp32, x1_fp32 * cos_fp32) o2_vals = tl.fma(x1_fp32, sin_fp32, x2_fp32 * cos_fp32) - tl.store(output_row_ptr + offsets_x1, o1_vals.to(x1_vals.dtype), mask=mask) - tl.store(output_row_ptr + offsets_x2, o2_vals.to(x2_vals.dtype), mask=mask) + tl.store( + output_row_ptrs + offsets_x1[None, :], + o1_vals.to(x1_vals.dtype), + mask=mask, + ) + tl.store( + output_row_ptrs + offsets_x2[None, :], + o2_vals.to(x2_vals.dtype), + mask=mask, + ) def apply_rotary_embedding( @@ -76,11 +94,8 @@ def apply_rotary_embedding( assert head_size % 2 == 0, "head_size must be divisible by 2" - x_reshaped = x.view(-1, head_size) - output_reshaped = output.view(-1, head_size) - - # num_tokens per head, 1 token per block - grid = (bsz * num_tokens * num_heads,) + x_reshaped = x.view(bsz * num_tokens, num_heads, head_size) + output_reshaped = output.view(bsz * num_tokens, num_heads, head_size) if interleaved and cos.shape[-1] == head_size: cos = cos[..., ::2].contiguous() @@ -89,7 +104,9 @@ def apply_rotary_embedding( cos = cos.contiguous() sin = sin.contiguous() - _rotary_embedding_kernel[grid]( + _rotary_embedding_kernel[ + lambda META: (bsz * num_tokens, triton.cdiv(num_heads, META["BLOCK_HEADS"])) + ]( output_reshaped, x_reshaped, cos, @@ -97,7 +114,10 @@ def apply_rotary_embedding( num_heads, head_size, num_tokens, + output_reshaped.stride(0), + output_reshaped.stride(1), x_reshaped.stride(0), + x_reshaped.stride(1), cos.stride(0), sin.stride(0), )