[codex] Optimize LTX2 split rotary kernel (#24732)
This commit is contained in:
@@ -19,23 +19,31 @@ def _ltx2_split_rotary_kernel(
|
|||||||
stride_sin_b: tl.constexpr,
|
stride_sin_b: tl.constexpr,
|
||||||
stride_sin_h: tl.constexpr,
|
stride_sin_h: tl.constexpr,
|
||||||
stride_sin_t: tl.constexpr,
|
stride_sin_t: tl.constexpr,
|
||||||
|
BLOCK_HEADS: tl.constexpr,
|
||||||
BLOCK_HALF: tl.constexpr,
|
BLOCK_HALF: tl.constexpr,
|
||||||
):
|
):
|
||||||
pid_bt = tl.program_id(0)
|
pid_bt = tl.program_id(0)
|
||||||
head = tl.program_id(1)
|
head_block = tl.program_id(1)
|
||||||
batch = pid_bt // seq_len
|
batch = pid_bt // seq_len
|
||||||
token = pid_bt - batch * seq_len
|
token = pid_bt - batch * seq_len
|
||||||
|
heads = head_block * BLOCK_HEADS + tl.arange(0, BLOCK_HEADS)
|
||||||
offsets = tl.arange(0, BLOCK_HALF)
|
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
|
x_base = ((batch * seq_len + token) * num_heads + heads[:, None]) * head_dim
|
||||||
cos_base = batch * stride_cos_b + head * stride_cos_h + token * stride_cos_t
|
cos_base = (
|
||||||
sin_base = batch * stride_sin_b + head * stride_sin_h + token * stride_sin_t
|
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_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, mask=mask, other=0.0)
|
x_second = tl.load(
|
||||||
cos = tl.load(cos_ptr + cos_base + offsets, mask=mask, other=0.0)
|
x_ptr + x_base + half_dim + offsets[None, :], mask=mask, other=0.0
|
||||||
sin = tl.load(sin_ptr + sin_base + offsets, 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
|
# 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.
|
# 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)
|
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 + offsets[None, :], out_first, mask=mask)
|
||||||
tl.store(out_ptr + x_base + half_dim + offsets, out_second, mask=mask)
|
tl.store(out_ptr + x_base + half_dim + offsets[None, :], out_second, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
def apply_ltx2_split_rotary_emb(
|
def apply_ltx2_split_rotary_emb(
|
||||||
@@ -69,7 +77,10 @@ def apply_ltx2_split_rotary_emb(
|
|||||||
|
|
||||||
out = torch.empty_like(x)
|
out = torch.empty_like(x)
|
||||||
block_half = triton.next_power_of_2(half_dim)
|
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,
|
out,
|
||||||
x,
|
x,
|
||||||
cos,
|
cos,
|
||||||
@@ -84,7 +95,8 @@ def apply_ltx2_split_rotary_emb(
|
|||||||
sin.stride(0),
|
sin.stride(0),
|
||||||
sin.stride(1),
|
sin.stride(1),
|
||||||
sin.stride(2),
|
sin.stride(2),
|
||||||
|
BLOCK_HEADS=block_heads,
|
||||||
BLOCK_HALF=block_half,
|
BLOCK_HALF=block_half,
|
||||||
num_warps=1,
|
num_warps=num_warps,
|
||||||
)
|
)
|
||||||
return out
|
return out
|
||||||
|
|||||||
Reference in New Issue
Block a user