[Diffusion] Preserve BF16 rounding in Hopper LTX QKNorm and RoPE fusion (#38533)

This commit is contained in:
Xiaoyu Zhang
2026-09-11 23:02:54 +08:00
committed by GitHub
parent dd67a42634
commit 7f09fbcd25
2 changed files with 35 additions and 14 deletions
@@ -4,8 +4,8 @@
// Developed with MIT HAN Lab Kernel Design Agents: // Developed with MIT HAN Lab Kernel Design Agents:
// https://github.com/mit-han-lab/kernel-design-agents // https://github.com/mit-han-lab/kernel-design-agents
// //
// This mirrors the LTX2 eager oracle: RMSNorm and split RoPE both run in // The original FP32 path rounds only at the final attention input. The
// fp32, rounding to bf16 only once at the final attention input. // Hopper variant preserves the eager BF16 RMSNorm and cosine-product rounding.
#pragma once #pragma once
@@ -71,17 +71,29 @@ SGL_DEVICE float compute_rstd(
return *s_rstd; return *s_rstd;
} }
template <bool kRoundIntermediates>
SGL_DEVICE float norm_value(float x, float weight, float rstd) { SGL_DEVICE float norm_value(float x, float weight, float rstd) {
return weight * (rstd * x); const float value = weight * (rstd * x);
if constexpr (kRoundIntermediates) {
return __bfloat162float(__float2bfloat16_rn(value));
}
return value;
} }
template <bool kRoundIntermediates>
SGL_DEVICE void rope_pair(float x0, float x1, float cos, float sin, float& y0, float& y1) { SGL_DEVICE void rope_pair(float x0, float x1, float cos, float sin, float& y0, float& y1) {
const float p0 = x0 * cos; float p0 = x0 * cos;
const float p1 = x1 * cos; float p1 = x1 * cos;
if constexpr (kRoundIntermediates) {
// Hopper eager stores the BF16 cosine product before its addcmul update.
p0 = __bfloat162float(__float2bfloat16_rn(p0));
p1 = __bfloat162float(__float2bfloat16_rn(p1));
}
y0 = fmaf(-sin, x1, p0); y0 = fmaf(-sin, x1, p0);
y1 = fmaf(sin, x0, p1); y1 = fmaf(sin, x0, p1);
} }
template <bool kRoundIntermediates>
__global__ void ltx2_qknorm_split_rope_kernel( __global__ void ltx2_qknorm_split_rope_kernel(
const bf16_t* __restrict__ x, const bf16_t* __restrict__ x,
const bf16_t* __restrict__ cos, const bf16_t* __restrict__ cos,
@@ -119,19 +131,23 @@ __global__ void ltx2_qknorm_split_rope_kernel(
const int64_t offset = pair - head * half_dim; const int64_t offset = pair - head * half_dim;
const int64_t idx0 = head * head_dim + offset; const int64_t idx0 = head * head_dim + offset;
const int64_t idx1 = idx0 + half_dim; const int64_t idx1 = idx0 + half_dim;
const float n0 = norm_value(__bfloat162float(xrow[idx0]), __bfloat162float(weight[idx0]), rstd); const float n0 =
const float n1 = norm_value(__bfloat162float(xrow[idx1]), __bfloat162float(weight[idx1]), rstd); norm_value<kRoundIntermediates>(__bfloat162float(xrow[idx0]), __bfloat162float(weight[idx0]), rstd);
const float n1 =
norm_value<kRoundIntermediates>(__bfloat162float(xrow[idx1]), __bfloat162float(weight[idx1]), rstd);
const int64_t cos_offset = batch * stride_cos_b + head * stride_cos_h + token * stride_cos_t + offset; const int64_t cos_offset = batch * stride_cos_b + head * stride_cos_h + token * stride_cos_t + offset;
const int64_t sin_offset = batch * stride_sin_b + head * stride_sin_h + token * stride_sin_t + offset; const int64_t sin_offset = batch * stride_sin_b + head * stride_sin_h + token * stride_sin_t + offset;
float y0; float y0;
float y1; float y1;
rope_pair(n0, n1, __bfloat162float(cos[cos_offset]), __bfloat162float(sin[sin_offset]), y0, y1); rope_pair<kRoundIntermediates>(
n0, n1, __bfloat162float(cos[cos_offset]), __bfloat162float(sin[sin_offset]), y0, y1);
outrow[idx0] = __float2bfloat16_rn(y0); outrow[idx0] = __float2bfloat16_rn(y0);
outrow[idx1] = __float2bfloat16_rn(y1); outrow[idx1] = __float2bfloat16_rn(y1);
} }
} }
template <bool kRoundIntermediates>
inline void launch_one( inline void launch_one(
const tvm::ffi::TensorView& x, const tvm::ffi::TensorView& x,
const tvm::ffi::TensorView& cos, const tvm::ffi::TensorView& cos,
@@ -155,7 +171,7 @@ inline void launch_one(
} }
host::RuntimeCheck(num_rows <= static_cast<int64_t>(UINT32_MAX), "LTX2 QKNorm split-RoPE grid is too large"); host::RuntimeCheck(num_rows <= static_cast<int64_t>(UINT32_MAX), "LTX2 QKNorm split-RoPE grid is too large");
host::LaunchKernel(dim3(static_cast<uint32_t>(num_rows)), dim3(32, 4), device)( host::LaunchKernel(dim3(static_cast<uint32_t>(num_rows)), dim3(32, 4), device)(
ltx2_qknorm_split_rope_kernel, ltx2_qknorm_split_rope_kernel<kRoundIntermediates>,
reinterpret_cast<const bf16_t*>(data_ptr(x)), reinterpret_cast<const bf16_t*>(data_ptr(x)),
reinterpret_cast<const bf16_t*>(data_ptr(cos)), reinterpret_cast<const bf16_t*>(data_ptr(cos)),
reinterpret_cast<const bf16_t*>(data_ptr(sin)), reinterpret_cast<const bf16_t*>(data_ptr(sin)),
@@ -174,6 +190,7 @@ inline void launch_one(
} }
struct LTX2QKNormSplitRopeKernel { struct LTX2QKNormSplitRopeKernel {
template <bool kRoundIntermediates>
static void static void
run(tvm::ffi::TensorView q_out, run(tvm::ffi::TensorView q_out,
tvm::ffi::TensorView k_out, tvm::ffi::TensorView k_out,
@@ -233,7 +250,7 @@ struct LTX2QKNormSplitRopeKernel {
const int64_t batch_size = batch.unwrap(); const int64_t batch_size = batch.unwrap();
const DLDevice dl_device = device.unwrap(); const DLDevice dl_device = device.unwrap();
launch_one( launch_one<kRoundIntermediates>(
q, q,
q_cos, q_cos,
q_sin, q_sin,
@@ -251,7 +268,7 @@ struct LTX2QKNormSplitRopeKernel {
q_sin.stride(1), q_sin.stride(1),
q_sin.stride(2), q_sin.stride(2),
dl_device); dl_device);
launch_one( launch_one<kRoundIntermediates>(
k, k,
k_cos, k_cos,
k_sin, k_sin,
@@ -13,14 +13,15 @@ if TYPE_CHECKING:
@cache_once @cache_once
def _jit_ltx2_qknorm_split_rope_module() -> Module: def _jit_ltx2_qknorm_split_rope_module(round_intermediates: bool = False) -> Module:
return load_jit( return load_jit(
"diffusion_ltx2_qknorm_split_rope", "diffusion_ltx2_qknorm_split_rope",
cuda_files=[_cuda_source("diffusion/ltx2_qknorm_split_rope.cuh")], cuda_files=[_cuda_source("diffusion/ltx2_qknorm_split_rope.cuh")],
cuda_wrappers=[ cuda_wrappers=[
( (
"ltx2_qknorm_split_rope_pair", "ltx2_qknorm_split_rope_pair",
"ltx2_qknorm_split_rope::LTX2QKNormSplitRopeKernel::run", "ltx2_qknorm_split_rope::LTX2QKNormSplitRopeKernel::run"
f"<{str(round_intermediates).lower()}>",
) )
], ],
) )
@@ -38,6 +39,7 @@ def _fake_impl(
eps: float, eps: float,
num_heads: int, num_heads: int,
head_dim: int, head_dim: int,
round_intermediates: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]: ) -> tuple[torch.Tensor, torch.Tensor]:
return torch.empty_like(q, dtype=torch.bfloat16), torch.empty_like( return torch.empty_like(q, dtype=torch.bfloat16), torch.empty_like(
k, dtype=torch.bfloat16 k, dtype=torch.bfloat16
@@ -61,10 +63,11 @@ def _ltx2_qknorm_split_rope_custom_op(
eps: float, eps: float,
num_heads: int, num_heads: int,
head_dim: int, head_dim: int,
round_intermediates: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]: ) -> tuple[torch.Tensor, torch.Tensor]:
q_out = torch.empty_like(q, dtype=torch.bfloat16) q_out = torch.empty_like(q, dtype=torch.bfloat16)
k_out = torch.empty_like(k, dtype=torch.bfloat16) k_out = torch.empty_like(k, dtype=torch.bfloat16)
module = _jit_ltx2_qknorm_split_rope_module() module = _jit_ltx2_qknorm_split_rope_module(round_intermediates)
module.ltx2_qknorm_split_rope_pair( module.ltx2_qknorm_split_rope_pair(
q_out, q_out,
k_out, k_out,
@@ -215,4 +218,5 @@ def ltx2_qknorm_split_rope_cuda(
float(eps), float(eps),
int(num_heads), int(num_heads),
int(head_dim), int(head_dim),
round_intermediates=allow_sm90 and _is_sm90(q),
) )