[KDA-Pilot] Add diffusion residual-gate CUDA fast path for LTX2 (#29361)

Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
Xiaoyu Zhang
2026-06-27 12:59:41 +08:00
committed by GitHub
co-authored by Claude Opus 4.8
parent cd6dedf972
commit 495f13fa12
6 changed files with 672 additions and 10 deletions
@@ -8,6 +8,9 @@
//
// All other public-op inputs fall back to the existing CuTe-DSL implementation
// from the Python dispatcher.
//
// Developed with MIT HAN Lab Kernel Design Agents:
// https://github.com/mit-han-lab/kernel-design-agents
#pragma once
@@ -0,0 +1,322 @@
// CUDA fast path for diffusion residual-gate elementwise updates.
//
// Implements:
// out = residual + update * gate
//
// The production shapes come from LTX-2.3 HQ residual/gate updates. This is
// intentionally narrow: contiguous residual/update/out tensors, with either a
// full contiguous gate or a row-broadcast [1, 1, D] gate.
//
// Developed with MIT HAN Lab Kernel Design Agents:
// https://github.com/mit-han-lab/kernel-design-agents
#pragma once
#include <sgl_kernel/tensor.h> // For host dtype helpers and TensorView metadata
#include <sgl_kernel/utils.h> // For RuntimeCheck and div_ceil
#include <sgl_kernel/type.cuh> // For dtype_trait conversions
#include <sgl_kernel/utils.cuh> // For LaunchKernel and CUDA dtype aliases
#include <cstdint>
namespace sglang_residual_gate_add {
namespace {
constexpr int kBlockSize = 256;
constexpr int kBcastRowsPerBlock = 4;
constexpr int kBcastColsVecPerBlock = 256;
constexpr int64_t kMaxGrid = 65535;
enum class GateMode : int { kFull = 0, kBcastRow = 1 };
inline const char* data_ptr(const tvm::ffi::TensorView& t) {
return static_cast<const char*>(t.data_ptr()) + t.byte_offset();
}
inline char* mutable_data_ptr(const tvm::ffi::TensorView& t) {
return static_cast<char*>(t.data_ptr()) + t.byte_offset();
}
inline bool aligned16(const void* p) {
return (reinterpret_cast<uintptr_t>(p) & 0xF) == 0;
}
inline int64_t numel(const tvm::ffi::TensorView& t) {
int64_t n = 1;
for (int i = 0; i < t.ndim(); ++i) {
n *= t.size(i);
}
return n;
}
inline int64_t grid_for(int64_t total) {
int64_t grid = host::div_ceil(total, static_cast<int64_t>(kBlockSize));
if (grid < 1) {
grid = 1;
}
if (grid > kMaxGrid) {
grid = kMaxGrid;
}
return grid;
}
inline bool is_dense_contiguous(const tvm::ffi::TensorView& t) {
int64_t expected = 1;
for (int i = t.ndim() - 1; i >= 0; --i) {
if (t.size(i) == 1) {
continue;
}
if (t.stride(i) != expected) {
return false;
}
expected *= t.size(i);
}
return true;
}
template <typename T>
inline void check_dtype(const tvm::ffi::TensorView& t) {
host::RuntimeCheck(host::is_type<T>(t.dtype()), "unexpected dtype for residual_gate_add");
}
template <typename T>
__device__ __forceinline__ float to_float(T v) {
return static_cast<float>(v);
}
template <>
__device__ __forceinline__ float to_float<fp16_t>(fp16_t v) {
return __half2float(v);
}
template <>
__device__ __forceinline__ float to_float<bf16_t>(bf16_t v) {
return __bfloat162float(v);
}
template <typename T>
__device__ __forceinline__ T residual_gate_value(T residual, T update, T gate) {
const T product = dtype_trait<T>::from(to_float(update) * to_float(gate));
return dtype_trait<T>::from(to_float(residual) + to_float(product));
}
template <typename T>
union Vec16 {
static constexpr int kElems = 16 / sizeof(T);
uint4 raw;
T elems[kElems];
};
template <typename T, int kVec>
__global__ void residual_gate_add_vec_kernel(
const T* __restrict__ residual,
const T* __restrict__ update,
const T* __restrict__ gate,
T* __restrict__ out,
int64_t n_vec) {
const int64_t stride = static_cast<int64_t>(gridDim.x) * blockDim.x;
for (int64_t v = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x; v < n_vec; v += stride) {
const Vec16<T> r{.raw = reinterpret_cast<const uint4*>(residual)[v]};
const Vec16<T> u{.raw = reinterpret_cast<const uint4*>(update)[v]};
const Vec16<T> g{.raw = reinterpret_cast<const uint4*>(gate)[v]};
Vec16<T> o;
#pragma unroll
for (int i = 0; i < kVec; ++i) {
o.elems[i] = residual_gate_value(r.elems[i], u.elems[i], g.elems[i]);
}
reinterpret_cast<uint4*>(out)[v] = o.raw;
}
}
template <typename T, int kVec>
__global__ void residual_gate_add_bcast_row_tile_kernel(
const T* __restrict__ residual,
const T* __restrict__ update,
const T* __restrict__ gate,
T* __restrict__ out,
int64_t rows,
int64_t row_vec) {
const int64_t col_vec = static_cast<int64_t>(blockIdx.x) * kBcastColsVecPerBlock + threadIdx.x;
if (col_vec >= row_vec) {
return;
}
const Vec16<T> g{.raw = SGLANG_LDG(reinterpret_cast<const uint4*>(gate) + col_vec)};
// Grid-stride over row tiles so the launch stays valid even when the number
// of row tiles exceeds the gridDim.y hardware limit.
const int64_t row_tile_stride = static_cast<int64_t>(gridDim.y) * kBcastRowsPerBlock;
for (int64_t row_base = static_cast<int64_t>(blockIdx.y) * kBcastRowsPerBlock; row_base < rows;
row_base += row_tile_stride) {
#pragma unroll
for (int row_offset = 0; row_offset < kBcastRowsPerBlock; ++row_offset) {
const int64_t row = row_base + row_offset;
if (row < rows) {
const int64_t v = row * row_vec + col_vec;
const Vec16<T> r{.raw = reinterpret_cast<const uint4*>(residual)[v]};
const Vec16<T> u{.raw = reinterpret_cast<const uint4*>(update)[v]};
Vec16<T> o;
#pragma unroll
for (int i = 0; i < kVec; ++i) {
o.elems[i] = residual_gate_value(r.elems[i], u.elems[i], g.elems[i]);
}
reinterpret_cast<uint4*>(out)[v] = o.raw;
}
}
}
}
template <typename T, GateMode kGate>
__global__ void residual_gate_add_scalar_kernel(
const T* __restrict__ residual,
const T* __restrict__ update,
const T* __restrict__ gate,
T* __restrict__ out,
int64_t begin,
int64_t total,
int64_t D) {
const int64_t stride = static_cast<int64_t>(gridDim.x) * blockDim.x;
for (int64_t i = begin + static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x; i < total; i += stride) {
const T gate_value = kGate == GateMode::kFull ? gate[i] : SGLANG_LDG(gate + (i % D));
out[i] = residual_gate_value(residual[i], update[i], gate_value);
}
}
template <typename T>
inline void launch_residual_gate_add(
const tvm::ffi::TensorView& out,
const tvm::ffi::TensorView& residual,
const tvm::ffi::TensorView& update,
const tvm::ffi::TensorView& gate,
GateMode mode) {
const int64_t total = numel(residual);
if (total == 0) {
return;
}
const int64_t D = residual.size(residual.ndim() - 1);
const T* residual_ptr = reinterpret_cast<const T*>(data_ptr(residual));
const T* update_ptr = reinterpret_cast<const T*>(data_ptr(update));
const T* gate_ptr = reinterpret_cast<const T*>(data_ptr(gate));
T* out_ptr = reinterpret_cast<T*>(mutable_data_ptr(out));
constexpr int kVec = 16 / sizeof(T);
const bool vec_ok = aligned16(residual_ptr) && aligned16(update_ptr) && aligned16(gate_ptr) && aligned16(out_ptr) &&
(D % kVec == 0) && (mode == GateMode::kBcastRow || total % kVec == 0);
int64_t done = 0;
if (vec_ok) {
const int64_t n_vec = total / kVec;
const int64_t row_vec = D / kVec;
if (mode == GateMode::kFull) {
host::LaunchKernel(static_cast<uint32_t>(grid_for(n_vec)), kBlockSize, out.device())(
residual_gate_add_vec_kernel<T, kVec>, residual_ptr, update_ptr, gate_ptr, out_ptr, n_vec);
} else {
const int64_t rows = total / D;
const int64_t col_blocks = host::div_ceil(row_vec, static_cast<int64_t>(kBcastColsVecPerBlock));
const int64_t row_tiles = host::div_ceil(rows, static_cast<int64_t>(kBcastRowsPerBlock));
const int64_t row_blocks = row_tiles > kMaxGrid ? kMaxGrid : row_tiles;
host::LaunchKernel(
dim3(static_cast<uint32_t>(col_blocks), static_cast<uint32_t>(row_blocks)),
dim3(kBcastColsVecPerBlock),
out.device())(
residual_gate_add_bcast_row_tile_kernel<T, kVec>, residual_ptr, update_ptr, gate_ptr, out_ptr, rows, row_vec);
}
done = n_vec * kVec;
}
if (done < total) {
if (mode == GateMode::kFull) {
host::LaunchKernel(static_cast<uint32_t>(grid_for(total - done)), kBlockSize, out.device())(
residual_gate_add_scalar_kernel<T, GateMode::kFull>,
residual_ptr,
update_ptr,
gate_ptr,
out_ptr,
done,
total,
D);
} else {
host::LaunchKernel(static_cast<uint32_t>(grid_for(total - done)), kBlockSize, out.device())(
residual_gate_add_scalar_kernel<T, GateMode::kBcastRow>,
residual_ptr,
update_ptr,
gate_ptr,
out_ptr,
done,
total,
D);
}
}
}
template <typename T>
inline GateMode validate_residual_gate_add(
const tvm::ffi::TensorView& out,
const tvm::ffi::TensorView& residual,
const tvm::ffi::TensorView& update,
const tvm::ffi::TensorView& gate) {
check_dtype<T>(out);
check_dtype<T>(residual);
check_dtype<T>(update);
check_dtype<T>(gate);
host::RuntimeCheck(residual.device().device_type == kDLCUDA, "residual must be CUDA");
host::RuntimeCheck(update.device().device_type == kDLCUDA, "update must be CUDA");
host::RuntimeCheck(gate.device().device_type == kDLCUDA, "gate must be CUDA");
host::RuntimeCheck(out.device().device_type == kDLCUDA, "out must be CUDA");
host::RuntimeCheck(
residual.device().device_id == update.device().device_id &&
residual.device().device_id == gate.device().device_id &&
residual.device().device_id == out.device().device_id,
"residual/update/gate/out must be on the same CUDA device");
host::RuntimeCheck(residual.ndim() >= 2, "residual must be at least 2D");
host::RuntimeCheck(update.ndim() == residual.ndim(), "update rank must match residual");
host::RuntimeCheck(out.ndim() == residual.ndim(), "out rank must match residual");
for (int i = 0; i < residual.ndim(); ++i) {
host::RuntimeCheck(update.size(i) == residual.size(i), "update shape must match residual");
host::RuntimeCheck(out.size(i) == residual.size(i), "out shape must match residual");
}
host::RuntimeCheck(is_dense_contiguous(residual), "residual must be contiguous");
host::RuntimeCheck(is_dense_contiguous(update), "update must be contiguous");
host::RuntimeCheck(is_dense_contiguous(out), "out must be contiguous");
host::RuntimeCheck(is_dense_contiguous(gate), "gate must be contiguous");
host::RuntimeCheck(data_ptr(out) != data_ptr(residual), "out must not alias residual");
host::RuntimeCheck(data_ptr(out) != data_ptr(update), "out must not alias update");
host::RuntimeCheck(data_ptr(out) != data_ptr(gate), "out must not alias gate");
const int D_dim = residual.ndim() - 1;
const int row_dim = residual.ndim() - 2;
host::RuntimeCheck(gate.ndim() == residual.ndim(), "gate rank must match residual");
host::RuntimeCheck(gate.size(D_dim) == residual.size(D_dim), "gate last dim must match residual");
bool full_gate = true;
for (int i = 0; i < residual.ndim(); ++i) {
full_gate = full_gate && gate.size(i) == residual.size(i);
}
if (full_gate) {
return GateMode::kFull;
}
host::RuntimeCheck(gate.size(row_dim) == 1, "broadcast gate row dim must be 1");
for (int i = 0; i < D_dim; ++i) {
host::RuntimeCheck(gate.size(i) == 1, "broadcast gate leading dims must be 1");
}
return GateMode::kBcastRow;
}
} // namespace
template <typename T>
struct ResidualGateAddKernel {
static void
run(tvm::ffi::TensorView out, tvm::ffi::TensorView residual, tvm::ffi::TensorView update, tvm::ffi::TensorView gate) {
const GateMode mode = validate_residual_gate_add<T>(out, residual, update, gate);
launch_residual_gate_add<T>(out, residual, update, gate, mode);
}
};
} // namespace sglang_residual_gate_add
@@ -0,0 +1,85 @@
from __future__ import annotations
from typing import TYPE_CHECKING
import torch
from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args
from sglang.srt.utils.custom_op import register_custom_op
if TYPE_CHECKING:
from tvm_ffi.module import Module
_SUPPORTED_DTYPES = (torch.float16, torch.bfloat16, torch.float32)
@cache_once
def _jit_residual_gate_add_module(dtype: torch.dtype) -> Module:
args = make_cpp_args(dtype)
return load_jit(
"diffusion_residual_gate_add",
*args,
cuda_files=["diffusion/residual_gate_add.cuh"],
cuda_wrappers=[
(
"residual_gate_add",
"sglang_residual_gate_add::" f"ResidualGateAddKernel<{args}>::run",
),
],
)
def _fake_impl(
residual: torch.Tensor, update: torch.Tensor, gate: torch.Tensor
) -> torch.Tensor:
return torch.empty_like(residual)
@register_custom_op(
op_name="diffusion_residual_gate_add",
mutates_args=[],
fake_impl=_fake_impl,
)
def _residual_gate_add_custom_op(
residual: torch.Tensor, update: torch.Tensor, gate: torch.Tensor
) -> torch.Tensor:
out = torch.empty_like(residual)
module = _jit_residual_gate_add_module(residual.dtype)
module.residual_gate_add(out, residual, update, gate)
return out
def _is_row_broadcast_gate(residual: torch.Tensor, gate: torch.Tensor) -> bool:
if gate.dim() != residual.dim() or gate.shape[-1] != residual.shape[-1]:
return False
row_dim = gate.dim() - 2
return gate.shape[row_dim] == 1 and all(size == 1 for size in gate.shape[:-1])
def can_use_residual_gate_add_cuda(
residual: torch.Tensor, update: torch.Tensor, gate: torch.Tensor
) -> bool:
return (
residual.dtype in _SUPPORTED_DTYPES
and residual.dtype == update.dtype
and residual.dtype == gate.dtype
and residual.is_cuda
and update.is_cuda
and gate.is_cuda
and residual.device == update.device == gate.device
and residual.dim() >= 2
and update.shape == residual.shape
and (gate.shape == residual.shape or _is_row_broadcast_gate(residual, gate))
and residual.is_contiguous()
and update.is_contiguous()
and gate.is_contiguous()
)
def residual_gate_add_cuda(
residual: torch.Tensor, update: torch.Tensor, gate: torch.Tensor
) -> torch.Tensor:
if not can_use_residual_gate_add_cuda(residual, update, gate):
raise RuntimeError("unsupported input for residual_gate_add CUDA")
return _residual_gate_add_custom_op(residual, update, gate)
@@ -10,6 +10,10 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
from sglang.jit_kernel.diffusion.residual_gate_add import (
can_use_residual_gate_add_cuda,
residual_gate_add_cuda,
)
from sglang.multimodal_gen.configs.models.dits.ltx_2 import LTX2ArchConfig, LTX2Config
from sglang.multimodal_gen.runtime.distributed import (
get_sp_parallel_rank,
@@ -48,6 +52,29 @@ logger = init_logger(__name__)
ADALN_NUM_BASE_PARAMS = 6
ADALN_NUM_CROSS_ATTN_PARAMS = 3
_LTX2_RESIDUAL_GATE_CUDA_DISABLED = False
def _ltx2_residual_gate_add(
residual: torch.Tensor,
update: torch.Tensor,
gate: torch.Tensor,
) -> torch.Tensor:
global _LTX2_RESIDUAL_GATE_CUDA_DISABLED
if not _LTX2_RESIDUAL_GATE_CUDA_DISABLED and can_use_residual_gate_add_cuda(
residual, update, gate
):
try:
return residual_gate_add_cuda(residual, update, gate)
except Exception as exc:
if torch.compiler.is_compiling():
raise
logger.warning_once(f"Disabling LTX2 residual-gate CUDA fast path: {exc}")
_LTX2_RESIDUAL_GATE_CUDA_DISABLED = True
return residual + update * gate
_LTX2_FUSED_ADA_VALUES_RUNTIME_DISABLED = False
@@ -1108,7 +1135,9 @@ class LTX2TransformerBlock(nn.Module):
gather_context_kv_for_sp=audio_replicated_for_sp,
context_replicated_prefix_len=video_memory_prefix_len,
)
hidden_states = hidden_states + attn_hidden_states * vgate_msa
hidden_states = _ltx2_residual_gate_add(
hidden_states, attn_hidden_states, vgate_msa
)
if audio_ada_values is None:
ashift_msa, ascale_msa, agate_msa = self.get_ada_values(
@@ -1128,7 +1157,9 @@ class LTX2TransformerBlock(nn.Module):
all_perturbed=skip_audio_self_attn,
skip_sequence_parallel_override=audio_replicated_for_sp,
)
audio_hidden_states = audio_hidden_states + attn_audio_hidden_states * agate_msa
audio_hidden_states = _ltx2_residual_gate_add(
audio_hidden_states, attn_audio_hidden_states, agate_msa
)
# 2. Prompt Cross-Attention
if self.cross_attention_adaln:
# LTX2.3
@@ -1156,7 +1187,9 @@ class LTX2TransformerBlock(nn.Module):
context=mod_encoder_hidden_states,
mask=encoder_attention_mask,
)
hidden_states = hidden_states + attn_hidden_states * vgate_q
hidden_states = _ltx2_residual_gate_add(
hidden_states, attn_hidden_states, vgate_q
)
if audio_ada_values is None:
ashift_q, ascale_q, agate_q = self.get_ada_values(
@@ -1182,8 +1215,8 @@ class LTX2TransformerBlock(nn.Module):
context=mod_audio_encoder_hidden_states,
mask=audio_encoder_attention_mask,
)
audio_hidden_states = (
audio_hidden_states + attn_audio_hidden_states * agate_q
audio_hidden_states = _ltx2_residual_gate_add(
audio_hidden_states, attn_audio_hidden_states, agate_q
)
else:
norm_hidden_states = self.rms_norm(hidden_states, self.norm_eps)
@@ -1284,7 +1317,9 @@ class LTX2TransformerBlock(nn.Module):
a2v_attn_hidden_states = (
a2v_attn_hidden_states * a2v_cross_attn_perturbation_mask
)
hidden_states = hidden_states + a2v_gate * a2v_attn_hidden_states
hidden_states = _ltx2_residual_gate_add(
hidden_states, a2v_attn_hidden_states, a2v_gate
)
# V2A
mod_norm_hidden_states = (
@@ -1308,8 +1343,8 @@ class LTX2TransformerBlock(nn.Module):
v2a_attn_hidden_states = (
v2a_attn_hidden_states * v2a_cross_attn_perturbation_mask
)
audio_hidden_states = (
audio_hidden_states + v2a_gate * v2a_attn_hidden_states
audio_hidden_states = _ltx2_residual_gate_add(
audio_hidden_states, v2a_attn_hidden_states, v2a_gate
)
# 4. Feedforward
if video_ada_values is None:
@@ -1322,7 +1357,7 @@ class LTX2TransformerBlock(nn.Module):
self.rms_norm(hidden_states, self.norm_eps) * (1 + vscale_mlp) + vshift_mlp
)
ff_output = self.ff(norm_hidden_states)
hidden_states = hidden_states + ff_output * vgate_mlp
hidden_states = _ltx2_residual_gate_add(hidden_states, ff_output, vgate_mlp)
if audio_ada_values is None:
ashift_mlp, ascale_mlp, agate_mlp = self.get_ada_values(
@@ -1335,7 +1370,9 @@ class LTX2TransformerBlock(nn.Module):
+ ashift_mlp
)
audio_ff_output = self.audio_ff(norm_audio_hidden_states)
audio_hidden_states = audio_hidden_states + audio_ff_output * agate_mlp
audio_hidden_states = _ltx2_residual_gate_add(
audio_hidden_states, audio_ff_output, agate_mlp
)
return hidden_states, audio_hidden_states