[Diffusion] Reuse shared AlignedVector and tidy jit_kernel/diffusion (#29664)

Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
Xiaoyu Zhang
2026-06-30 11:38:22 +08:00
committed by GitHub
co-authored by Claude Opus 4.8
parent 8dbf04fc56
commit 3add35e26d
10 changed files with 116 additions and 159 deletions
@@ -15,6 +15,7 @@
#include <sgl_kernel/utils.h> // For RuntimeCheck, div_ceil #include <sgl_kernel/utils.h> // For RuntimeCheck, div_ceil
#include <sgl_kernel/utils.cuh> // For LaunchKernel #include <sgl_kernel/utils.cuh> // For LaunchKernel
#include <sgl_kernel/vec.cuh> // For device::AlignedVector
#include <cstdint> #include <cstdint>
@@ -41,10 +42,7 @@ __global__ void __launch_bounds__(kBlockSize) cat_pad_flat_kernel(
int64_t pad_d_left, int64_t pad_d_left,
int64_t pad_h_top, int64_t pad_h_top,
int64_t pad_w_left) { int64_t pad_w_left) {
union Pack { using Pack = device::AlignedVector<ET, kVec>;
ET elem[kVec];
uint4 raw;
};
const int64_t nthreads = static_cast<int64_t>(gridDim.x) * blockDim.x; const int64_t nthreads = static_cast<int64_t>(gridDim.x) * blockDim.x;
for (int64_t vid = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x; vid < total_vecs; vid += nthreads) { for (int64_t vid = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x; vid < total_vecs; vid += nthreads) {
@@ -81,7 +79,7 @@ __global__ void __launch_bounds__(kBlockSize) cat_pad_flat_kernel(
value = SGLANG_LDG(src + iw); value = SGLANG_LDG(src + iw);
} }
} }
pack.elem[i] = value; pack[i] = value;
if (++ow == out_w) { if (++ow == out_w) {
ow = 0; ow = 0;
@@ -110,7 +108,7 @@ __global__ void __launch_bounds__(kBlockSize) cat_pad_flat_kernel(
} }
} }
reinterpret_cast<uint4*>(out)[vid] = pack.raw; pack.store(out, vid);
} }
} }
@@ -17,6 +17,7 @@
#include <sgl_kernel/type.cuh> // For dtype_trait conversions #include <sgl_kernel/type.cuh> // For dtype_trait conversions
#include <sgl_kernel/utils.cuh> // For LaunchKernel and CUDA dtype aliases #include <sgl_kernel/utils.cuh> // For LaunchKernel and CUDA dtype aliases
#include <sgl_kernel/vec.cuh> // For device::AlignedVector
#include <cstdint> #include <cstdint>
@@ -102,13 +103,6 @@ __device__ __forceinline__ T residual_gate_value(T residual, T update, T gate) {
return dtype_trait<T>::from(to_float(residual) + to_float(product)); 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> template <typename T, int kVec>
__global__ void residual_gate_add_vec_kernel( __global__ void residual_gate_add_vec_kernel(
const T* __restrict__ residual, const T* __restrict__ residual,
@@ -116,18 +110,18 @@ __global__ void residual_gate_add_vec_kernel(
const T* __restrict__ gate, const T* __restrict__ gate,
T* __restrict__ out, T* __restrict__ out,
int64_t n_vec) { int64_t n_vec) {
using Vec = device::AlignedVector<T, kVec>;
const int64_t stride = static_cast<int64_t>(gridDim.x) * blockDim.x; 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) { 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]}; Vec r, u, g, o;
const Vec16<T> u{.raw = reinterpret_cast<const uint4*>(update)[v]}; r.load(residual, v);
const Vec16<T> g{.raw = reinterpret_cast<const uint4*>(gate)[v]}; u.load(update, v);
g.load(gate, v);
Vec16<T> o;
#pragma unroll #pragma unroll
for (int i = 0; i < kVec; ++i) { for (int i = 0; i < kVec; ++i) {
o.elems[i] = residual_gate_value(r.elems[i], u.elems[i], g.elems[i]); o[i] = residual_gate_value(r[i], u[i], g[i]);
} }
reinterpret_cast<uint4*>(out)[v] = o.raw; o.store(out, v);
} }
} }
@@ -139,12 +133,14 @@ __global__ void residual_gate_add_bcast_row_tile_kernel(
T* __restrict__ out, T* __restrict__ out,
int64_t rows, int64_t rows,
int64_t row_vec) { int64_t row_vec) {
using Vec = device::AlignedVector<T, kVec>;
const int64_t col_vec = static_cast<int64_t>(blockIdx.x) * kBcastColsVecPerBlock + threadIdx.x; const int64_t col_vec = static_cast<int64_t>(blockIdx.x) * kBcastColsVecPerBlock + threadIdx.x;
if (col_vec >= row_vec) { if (col_vec >= row_vec) {
return; return;
} }
const Vec16<T> g{.raw = SGLANG_LDG(reinterpret_cast<const uint4*>(gate) + col_vec)}; Vec g;
g.load(gate, col_vec);
// Grid-stride over row tiles so the launch stays valid even when the number // Grid-stride over row tiles so the launch stays valid even when the number
// of row tiles exceeds the gridDim.y hardware limit. // of row tiles exceeds the gridDim.y hardware limit.
@@ -156,15 +152,14 @@ __global__ void residual_gate_add_bcast_row_tile_kernel(
const int64_t row = row_base + row_offset; const int64_t row = row_base + row_offset;
if (row < rows) { if (row < rows) {
const int64_t v = row * row_vec + col_vec; const int64_t v = row * row_vec + col_vec;
const Vec16<T> r{.raw = reinterpret_cast<const uint4*>(residual)[v]}; Vec r, u, o;
const Vec16<T> u{.raw = reinterpret_cast<const uint4*>(update)[v]}; r.load(residual, v);
u.load(update, v);
Vec16<T> o;
#pragma unroll #pragma unroll
for (int i = 0; i < kVec; ++i) { for (int i = 0; i < kVec; ++i) {
o.elems[i] = residual_gate_value(r.elems[i], u.elems[i], g.elems[i]); o[i] = residual_gate_value(r[i], u[i], g[i]);
} }
reinterpret_cast<uint4*>(out)[v] = o.raw; o.store(out, v);
} }
} }
} }
@@ -1,9 +1,12 @@
#pragma once
#include <sgl_kernel/tensor.h> #include <sgl_kernel/tensor.h>
#include <sgl_kernel/utils.h> #include <sgl_kernel/utils.h>
#include <sgl_kernel/math.cuh> #include <sgl_kernel/math.cuh>
#include <sgl_kernel/type.cuh> #include <sgl_kernel/type.cuh>
#include <sgl_kernel/utils.cuh> #include <sgl_kernel/utils.cuh>
#include <sgl_kernel/vec.cuh> // For device::AlignedVector
#include <dlpack/dlpack.h> #include <dlpack/dlpack.h>
#include <tvm/ffi/container/tensor.h> #include <tvm/ffi/container/tensor.h>
@@ -14,8 +17,12 @@
#include <cuda_runtime.h> #include <cuda_runtime.h>
#include <type_traits> #include <type_traits>
namespace sglang_timestep_embedding {
namespace { namespace {
constexpr int kVec = 4; // 16B float vector store
template <bool kFlipSinToCos, typename TIn> template <bool kFlipSinToCos, typename TIn>
__global__ void timestep_embedding_kernel( __global__ void timestep_embedding_kernel(
const TIn* __restrict__ t_ptr, const TIn* __restrict__ t_ptr,
@@ -24,6 +31,8 @@ __global__ void timestep_embedding_kernel(
float neg_log_max_period, float neg_log_max_period,
float scale, float scale,
int batch_size) { int batch_size) {
using Vec = device::AlignedVector<float, kVec>;
int row_idx = static_cast<int>(blockIdx.x * blockDim.y + threadIdx.y); int row_idx = static_cast<int>(blockIdx.x * blockDim.y + threadIdx.y);
if (row_idx >= batch_size) { if (row_idx >= batch_size) {
return; return;
@@ -34,36 +43,29 @@ __global__ void timestep_embedding_kernel(
int half_dim = dim / 2; int half_dim = dim / 2;
int thread_offset = static_cast<int>(threadIdx.x); int thread_offset = static_cast<int>(threadIdx.x);
while (thread_offset * 4 < half_dim) { while (thread_offset * kVec < half_dim) {
float4* top_half; // !flip: output is [sin | cos]; flip: output is [cos | sin].
float4* bottom_half; float* cos_dst;
float* sin_dst;
if constexpr (!kFlipSinToCos) { if constexpr (!kFlipSinToCos) {
bottom_half = reinterpret_cast<float4*>(output_batch_base_ptr + thread_offset * 4); sin_dst = output_batch_base_ptr + thread_offset * kVec;
top_half = reinterpret_cast<float4*>(output_batch_base_ptr + half_dim + thread_offset * 4); cos_dst = output_batch_base_ptr + half_dim + thread_offset * kVec;
} else { } else {
top_half = reinterpret_cast<float4*>(output_batch_base_ptr + thread_offset * 4); cos_dst = output_batch_base_ptr + thread_offset * kVec;
bottom_half = reinterpret_cast<float4*>(output_batch_base_ptr + half_dim + thread_offset * 4); sin_dst = output_batch_base_ptr + half_dim + thread_offset * kVec;
} }
float4 vals; Vec cos_vec;
vals.x = scale * t_val * device::math::exp(neg_log_max_period * __int2float_rn(thread_offset * 4 + 0)); Vec sin_vec;
vals.y = scale * t_val * device::math::exp(neg_log_max_period * __int2float_rn(thread_offset * 4 + 1)); #pragma unroll
vals.z = scale * t_val * device::math::exp(neg_log_max_period * __int2float_rn(thread_offset * 4 + 2)); for (int i = 0; i < kVec; ++i) {
vals.w = scale * t_val * device::math::exp(neg_log_max_period * __int2float_rn(thread_offset * 4 + 3)); const float angle =
scale * t_val * device::math::exp(neg_log_max_period * __int2float_rn(thread_offset * kVec + i));
float4 cos_vals; cos_vec[i] = device::math::cos(angle);
cos_vals.x = device::math::cos(vals.x); sin_vec[i] = device::math::sin(angle);
cos_vals.y = device::math::cos(vals.y); }
cos_vals.z = device::math::cos(vals.z); cos_vec.store(cos_dst);
cos_vals.w = device::math::cos(vals.w); sin_vec.store(sin_dst);
*top_half = cos_vals;
float4 sin_vals;
sin_vals.x = device::math::sin(vals.x);
sin_vals.y = device::math::sin(vals.y);
sin_vals.z = device::math::sin(vals.z);
sin_vals.w = device::math::sin(vals.w);
*bottom_half = sin_vals;
thread_offset += static_cast<int>(blockDim.x); thread_offset += static_cast<int>(blockDim.x);
} }
@@ -118,6 +120,8 @@ inline void launch_timestep_embedding(
} }
} }
} // namespace
template <typename TIn> template <typename TIn>
void timestep_embedding( void timestep_embedding(
tvm::ffi::TensorView input, tvm::ffi::TensorView input,
@@ -147,4 +151,4 @@ void timestep_embedding(
launch_timestep_embedding<TIn>(input, output, dim, flip_sin_to_cos, downscale_freq_shift, scale, max_period); launch_timestep_embedding<TIn>(input, output, dim, flip_sin_to_cos, downscale_freq_shift, scale, max_period);
} }
} // namespace } // namespace sglang_timestep_embedding
@@ -10,50 +10,15 @@ from sglang.jit_kernel.diffusion.cutedsl.common.norm_fusion import (
broadcast_tensor_for_bsfd, broadcast_tensor_for_bsfd,
tensor_slice_for_bsfd, tensor_slice_for_bsfd,
) )
from sglang.jit_kernel.diffusion.cutedsl.utils import TORCH_TO_CUTE_DTYPE, WARP_SIZE from sglang.jit_kernel.diffusion.cutedsl.utils import (
WARP_SIZE,
to_cute_arg,
to_fake_cute_args,
)
_COMPILE_CACHE = {} _COMPILE_CACHE = {}
def to_cute_arg(
t,
*,
assume_aligned: Optional[int] = 32,
use_32bit_stride: bool = False,
enable_tvm_ffi: bool = True,
):
"""
Convert a Python value into a CuTeDSL value.
"""
if isinstance(t, torch.Tensor):
return cute.runtime.from_dlpack(
t,
assumed_align=assume_aligned,
use_32bit_stride=use_32bit_stride,
enable_tvm_ffi=enable_tvm_ffi,
)
if isinstance(t, int):
return cutlass.Int32(t)
if isinstance(t, float):
return cutlass.Float32(t)
return t
def to_fake_cute_args(t: torch.Tensor):
if isinstance(t, torch.Tensor):
# Only keep the last dim as compile-time value to maximum compiled kernel reuse
# e.g. (1,2,1536):(3027,1536,1) -> (?,?,1536):(?,?,1)
D = t.shape[-1]
dtype = TORCH_TO_CUTE_DTYPE[t.dtype]
shape = (*(cute.sym_int() for _ in range(t.ndim - 1)), D)
stride = (*(cute.sym_int(divisibility=D) for _ in range(t.ndim - 1)), 1)
fake_t = cute.runtime.make_fake_tensor(
dtype, shape, stride, memspace=cute.AddressSpace.gmem, assumed_align=32
)
return fake_t
return to_cute_arg(t)
class NormTanhMulAddNormScale: class NormTanhMulAddNormScale:
@classmethod @classmethod
def make_hash_key(cls, *inputs): def make_hash_key(cls, *inputs):
@@ -166,7 +131,7 @@ class NormTanhMulAddNormScale:
@cute.jit @cute.jit
def copy_if(src, dst): def copy_if(src, dst):
if cutlass.const_expr( if cutlass.const_expr(
isinstance(src, cute.Tensor) and isinstance(src, cute.Tensor) isinstance(src, cute.Tensor) and isinstance(dst, cute.Tensor)
): ):
cute.autovec_copy(src, dst) # LDG.128 cute.autovec_copy(src, dst) # LDG.128
@@ -10,50 +10,14 @@ from sglang.jit_kernel.diffusion.cutedsl.common.norm_fusion import (
broadcast_tensor_for_bsfd, broadcast_tensor_for_bsfd,
tensor_slice_for_bsfd, tensor_slice_for_bsfd,
) )
from sglang.jit_kernel.diffusion.cutedsl.utils import TORCH_TO_CUTE_DTYPE, WARP_SIZE from sglang.jit_kernel.diffusion.cutedsl.utils import (
WARP_SIZE,
to_fake_cute_args,
)
_COMPILE_CACHE = {} _COMPILE_CACHE = {}
def to_cute_arg(
t,
*,
assume_aligned: Optional[int] = 32,
use_32bit_stride: bool = False,
enable_tvm_ffi: bool = True,
):
"""
Convert a Python value into a CuTeDSL value.
"""
if isinstance(t, torch.Tensor):
return cute.runtime.from_dlpack(
t,
assumed_align=assume_aligned,
use_32bit_stride=use_32bit_stride,
enable_tvm_ffi=enable_tvm_ffi,
)
if isinstance(t, int):
return cutlass.Int32(t)
if isinstance(t, float):
return cutlass.Float32(t)
return t
def to_fake_cute_args(t: torch.Tensor):
if isinstance(t, torch.Tensor):
# Only keep the last dim as compile-time value to maximum compiled kernel reuse
# e.g. (1,2,1536):(3027,1536,1) -> (?,?,1536):(?,?,1)
D = t.shape[-1]
dtype = TORCH_TO_CUTE_DTYPE[t.dtype]
shape = (*(cute.sym_int() for _ in range(t.ndim - 1)), D)
stride = (*(cute.sym_int(divisibility=D) for _ in range(t.ndim - 1)), 1)
fake_t = cute.runtime.make_fake_tensor(
dtype, shape, stride, memspace=cute.AddressSpace.gmem, assumed_align=32
)
return fake_t
return to_cute_arg(t)
class ScaleResidualNormScaleShift: class ScaleResidualNormScaleShift:
@classmethod @classmethod
def make_hash_key(cls, *inputs): def make_hash_key(cls, *inputs):
@@ -234,7 +198,7 @@ def validate_x(t: torch.Tensor, B: int, S: int, D: int):
raise ValueError(f"Validate failed: not contiguous on dim D.") raise ValueError(f"Validate failed: not contiguous on dim D.")
def validate_weight_bias(t: Optional[torch.Tensor], B: int, S: int, D: int): def validate_weight_bias(t: Optional[torch.Tensor], D: int):
if t is None: if t is None:
return return
if t.dtype not in (torch.float16, torch.bfloat16, torch.float32): if t.dtype not in (torch.float16, torch.bfloat16, torch.float32):
@@ -312,8 +276,8 @@ def fused_norm_scale_shift(
# Tensor Validation # Tensor Validation
BSD = x.shape BSD = x.shape
validate_x(x, *BSD) validate_x(x, *BSD)
validate_weight_bias(weight, *BSD) validate_weight_bias(weight, BSD[-1])
validate_weight_bias(bias, *BSD) validate_weight_bias(bias, BSD[-1])
validate_scale_shift(scale, *BSD) validate_scale_shift(scale, *BSD)
validate_scale_shift(shift, *BSD) validate_scale_shift(shift, *BSD)
@@ -399,14 +363,13 @@ def fused_scale_residual_norm_scale_shift(
validate_x(x, *BSD) validate_x(x, *BSD)
validate_x(residual, *BSD) validate_x(residual, *BSD)
validate_gate(gate, *BSD) validate_gate(gate, *BSD)
validate_weight_bias(weight, *BSD) validate_weight_bias(weight, BSD[-1])
validate_weight_bias(bias, *BSD) validate_weight_bias(bias, BSD[-1])
validate_scale_shift(scale, *BSD) validate_scale_shift(scale, *BSD)
validate_scale_shift(shift, *BSD) validate_scale_shift(shift, *BSD)
if norm_type == "layer" or norm_type == "rms": if norm_type == "layer" or norm_type == "rms":
stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream) stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream)
# if norm_type == "layer" or norm_type == "rms":
D = x.shape[-1] D = x.shape[-1]
if D % 256 != 0 or D > 8192: if D % 256 != 0 or D > 8192:
raise ValueError( raise ValueError(
@@ -1,4 +1,7 @@
from typing import Optional
import cutlass import cutlass
import cutlass.cute as cute
import torch import torch
WARP_SIZE = 32 WARP_SIZE = 32
@@ -8,3 +11,42 @@ TORCH_TO_CUTE_DTYPE = {
torch.bfloat16: cutlass.BFloat16, torch.bfloat16: cutlass.BFloat16,
torch.float32: cutlass.Float32, torch.float32: cutlass.Float32,
} }
def to_cute_arg(
t,
*,
assume_aligned: Optional[int] = 32,
use_32bit_stride: bool = False,
enable_tvm_ffi: bool = True,
):
"""
Convert a Python value into a CuTeDSL value.
"""
if isinstance(t, torch.Tensor):
return cute.runtime.from_dlpack(
t,
assumed_align=assume_aligned,
use_32bit_stride=use_32bit_stride,
enable_tvm_ffi=enable_tvm_ffi,
)
if isinstance(t, int):
return cutlass.Int32(t)
if isinstance(t, float):
return cutlass.Float32(t)
return t
def to_fake_cute_args(t: torch.Tensor):
if isinstance(t, torch.Tensor):
# Only keep the last dim as compile-time value to maximum compiled kernel reuse
# e.g. (1,2,1536):(3027,1536,1) -> (?,?,1536):(?,?,1)
D = t.shape[-1]
dtype = TORCH_TO_CUTE_DTYPE[t.dtype]
shape = (*(cute.sym_int() for _ in range(t.ndim - 1)), D)
stride = (*(cute.sym_int(divisibility=D) for _ in range(t.ndim - 1)), 1)
fake_t = cute.runtime.make_fake_tensor(
dtype, shape, stride, memspace=cute.AddressSpace.gmem, assumed_align=32
)
return fake_t
return to_cute_arg(t)
@@ -38,7 +38,6 @@ def triton_autotune_configs():
for warp_count in [1, 2, 4, 8, 16, 32] for warp_count in [1, 2, 4, 8, 16, 32]
if warp_count * warp_size <= max_threads_per_block if warp_count * warp_size <= max_threads_per_block
] ]
# return [triton.Config({}, num_warps=8)]
# Copied from flash-attn # Copied from flash-attn
@@ -120,18 +120,6 @@ def prepare_rope_tables(
return rope_cos.contiguous(), rope_sin.contiguous() return rope_cos.contiguous(), rope_sin.contiguous()
def _precompute_inv_rms(
qkv: torch.Tensor, idx: int, C: int, eps: float = 1e-5
) -> torch.Tensor:
"""Compute 1/RMS for one component of QKV over the full C = H*D channel dim.
qkv: (B, N, 3, H, D); idx: 0=Q, 1=K, 2=V; C: H*D. Returns (B, N) float32.
"""
raw = qkv[:, :, idx].float() # (B, N, H, D)
sq_sum = (raw * raw).sum(dim=(-2, -1)) # (B, N)
return torch.rsqrt(sq_sum / C + eps)
# ===================================================================== # =====================================================================
# Fused single-pass Q+K inverse-RMS Triton kernel # Fused single-pass Q+K inverse-RMS Triton kernel
# ===================================================================== # =====================================================================
@@ -182,7 +170,7 @@ def fused_qk_inv_rms(
) -> tuple[torch.Tensor, torch.Tensor]: ) -> tuple[torch.Tensor, torch.Tensor]:
"""Single-pass Triton fused Q+K inverse-RMS. """Single-pass Triton fused Q+K inverse-RMS.
Replaces two ``_precompute_inv_rms`` scans with one launch that reads each Replaces two separate PyTorch RMS scans with one launch that reads each
``(b, n)`` row of ``qkv`` exactly once. ``(b, n)`` row of ``qkv`` exactly once.
qkv: (B, N, 3, H, D) contiguous. Returns (q_inv_rms, k_inv_rms), each (B, N) float32. qkv: (B, N, 3, H, D) contiguous. Returns (q_inv_rms, k_inv_rms), each (B, N) float32.
""" """
@@ -230,7 +218,7 @@ def fused_bigdn_func(
"""Bidirectional fused GDN. Returns ``(B, N, H, D)``. """Bidirectional fused GDN. Returns ``(B, N, H, D)``.
Thin entry point kept for call-site stability; delegates to Thin entry point kept for call-site stability; delegates to
:func:`fused_bigdn_bidi_chunkwise` from ``fused_gdn_chunkwise``. :func:`fused_bigdn_bidi_chunkwise` from ``sana_wm_gdn_chunkwise``.
""" """
from sglang.jit_kernel.diffusion.triton.sana_wm_gdn_chunkwise import ( from sglang.jit_kernel.diffusion.triton.sana_wm_gdn_chunkwise import (
fused_bigdn_bidi_chunkwise, fused_bigdn_bidi_chunkwise,
@@ -224,7 +224,6 @@ def _fused_scale_shift_4d_kernel(
scale_ptr, scale_ptr,
shift_ptr, shift_ptr,
scale_constant: tl.constexpr, # scale_constant is either 0 or 1. scale_constant: tl.constexpr, # scale_constant is either 0 or 1.
rows,
inner_dim, inner_dim,
seq_len, seq_len,
num_frames, num_frames,
@@ -386,7 +385,6 @@ def fuse_scale_shift_kernel(
scale_reshaped, scale_reshaped,
shift_reshaped, shift_reshaped,
scale_constant, scale_constant,
rows,
C, C,
L, L,
num_frames, num_frames,
@@ -18,7 +18,12 @@ def _jit_timestep_embedding_module(dtype: torch.dtype) -> Module:
"timestep_embedding", "timestep_embedding",
*args, *args,
cuda_files=["diffusion/timestep_embedding.cuh"], cuda_files=["diffusion/timestep_embedding.cuh"],
cuda_wrappers=[("timestep_embedding", f"timestep_embedding<{args}>")], cuda_wrappers=[
(
"timestep_embedding",
f"sglang_timestep_embedding::timestep_embedding<{args}>",
)
],
) )