[Diffusion] Apply qknorm to flux2 and apply lightx2v rms_norm_one_pass kernel(without residual) (#17305)

Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
This commit is contained in:
Xiaoyu Zhang
2026-01-19 21:25:33 +08:00
committed by GitHub
co-authored by gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
parent f374623fa9
commit cc410a1088
3 changed files with 99 additions and 6 deletions
@@ -1106,3 +1106,58 @@ def rms_norm_fn(
out,
residual_out,
)
# Adapted from https://github.com/ModelTC/LightX2V/blob/main/lightx2v/common/ops/norm/triton_ops.py#L905-L956
@triton.jit
def _rms_norm_tiled_onepass(
y_ptr,
x_ptr,
w_ptr,
SEQ: tl.constexpr,
DIM: tl.constexpr,
EPS: tl.constexpr,
BLOCK_SIZE_SEQ: tl.constexpr,
BLOCK_SIZE_DIM: tl.constexpr,
):
seq_blk_id = tl.program_id(0)
seq_id = seq_blk_id * BLOCK_SIZE_SEQ
seq_offset = seq_id + tl.arange(0, BLOCK_SIZE_SEQ)[:, None]
s_mask = seq_offset < SEQ
d_offset = tl.arange(0, BLOCK_SIZE_DIM)[None, :]
d_mask = d_offset < DIM
y_blk = y_ptr + seq_offset * DIM + d_offset
x_blk = x_ptr + seq_offset * DIM + d_offset
mask = s_mask & d_mask
x = tl.load(x_blk, mask=mask, other=0.0).to(tl.float32)
mean_square = tl.sum(x * x, axis=1, keep_dims=True) / DIM
rstd = tl.math.rsqrt(mean_square + EPS)
w = tl.load(w_ptr + d_offset, mask=d_mask)
tl.store(y_blk, x * rstd * w, mask=mask)
def triton_one_pass_rms_norm(x: torch.Tensor, w: torch.Tensor, eps: float = 1e-6):
shape = x.shape
x = x.contiguous()
y = torch.empty_like(x)
x_view = x.reshape(-1, shape[-1])
y_view = y.reshape(-1, shape[-1])
S, D = x_view.shape
BLOCK_SIZE_SEQ = min(16, triton.next_power_of_2(max(1, S // 512)))
grid = (triton.cdiv(S, BLOCK_SIZE_SEQ),)
with torch.cuda.device(x.device):
torch.library.wrap_triton(_rms_norm_tiled_onepass)[grid](
y_view,
x_view,
w,
S,
D,
eps,
BLOCK_SIZE_DIM=triton.next_power_of_2(D),
BLOCK_SIZE_SEQ=BLOCK_SIZE_SEQ,
)
return y