[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:
co-authored by
Claude Opus 4.8
parent
cd6dedf972
commit
495f13fa12
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user