diff --git a/python/sglang/jit_kernel/diffusion/triton/ltx2_rotary.py b/python/sglang/jit_kernel/diffusion/triton/ltx2_rotary.py index 6bf8791ce..ca2f4ed7e 100644 --- a/python/sglang/jit_kernel/diffusion/triton/ltx2_rotary.py +++ b/python/sglang/jit_kernel/diffusion/triton/ltx2_rotary.py @@ -19,23 +19,31 @@ def _ltx2_split_rotary_kernel( stride_sin_b: tl.constexpr, stride_sin_h: tl.constexpr, stride_sin_t: tl.constexpr, + BLOCK_HEADS: tl.constexpr, BLOCK_HALF: tl.constexpr, ): pid_bt = tl.program_id(0) - head = tl.program_id(1) + head_block = tl.program_id(1) batch = pid_bt // seq_len token = pid_bt - batch * seq_len + heads = head_block * BLOCK_HEADS + tl.arange(0, BLOCK_HEADS) offsets = tl.arange(0, BLOCK_HALF) - mask = offsets < half_dim + mask = (heads[:, None] < num_heads) & (offsets[None, :] < half_dim) - x_base = ((batch * seq_len + token) * num_heads + head) * head_dim - cos_base = batch * stride_cos_b + head * stride_cos_h + token * stride_cos_t - sin_base = batch * stride_sin_b + head * stride_sin_h + token * stride_sin_t + x_base = ((batch * seq_len + token) * num_heads + heads[:, None]) * head_dim + cos_base = ( + batch * stride_cos_b + heads[:, None] * stride_cos_h + token * stride_cos_t + ) + sin_base = ( + batch * stride_sin_b + heads[:, None] * stride_sin_h + token * stride_sin_t + ) - x_first = tl.load(x_ptr + x_base + offsets, mask=mask, other=0.0) - x_second = tl.load(x_ptr + x_base + half_dim + offsets, mask=mask, other=0.0) - cos = tl.load(cos_ptr + cos_base + offsets, mask=mask, other=0.0) - sin = tl.load(sin_ptr + sin_base + offsets, mask=mask, other=0.0) + x_first = tl.load(x_ptr + x_base + offsets[None, :], mask=mask, other=0.0) + x_second = tl.load( + x_ptr + x_base + half_dim + offsets[None, :], mask=mask, other=0.0 + ) + cos = tl.load(cos_ptr + cos_base + offsets[None, :], mask=mask, other=0.0) + sin = tl.load(sin_ptr + sin_base + offsets[None, :], mask=mask, other=0.0) # Match the original PyTorch order: x * cos is written as BF16 first, then # addcmul_ computes the sine product in FP32 before the final BF16 store. @@ -46,8 +54,8 @@ def _ltx2_split_rotary_kernel( x_first.to(tl.float32) * sin.to(tl.float32) ) - tl.store(out_ptr + x_base + offsets, out_first, mask=mask) - tl.store(out_ptr + x_base + half_dim + offsets, out_second, mask=mask) + tl.store(out_ptr + x_base + offsets[None, :], out_first, mask=mask) + tl.store(out_ptr + x_base + half_dim + offsets[None, :], out_second, mask=mask) def apply_ltx2_split_rotary_emb( @@ -69,7 +77,10 @@ def apply_ltx2_split_rotary_emb( out = torch.empty_like(x) block_half = triton.next_power_of_2(half_dim) - _ltx2_split_rotary_kernel[(batch * seq_len, num_heads)]( + block_heads = min(16, triton.next_power_of_2(num_heads)) + num_warps = min(8, max(1, block_heads)) + grid = (batch * seq_len, triton.cdiv(num_heads, block_heads)) + _ltx2_split_rotary_kernel[grid]( out, x, cos, @@ -84,7 +95,8 @@ def apply_ltx2_split_rotary_emb( sin.stride(0), sin.stride(1), sin.stride(2), + BLOCK_HEADS=block_heads, BLOCK_HALF=block_half, - num_warps=1, + num_warps=num_warps, ) return out