feat: add Wan2.2 S2V (audio+img-to-video) support (#1925)

Co-authored-by: leejet <leejet714@gmail.com>
This commit is contained in:
George
2026-09-13 23:33:50 +08:00
committed by GitHub
co-authored by leejet
parent 4a7da26b73
commit 0bd72f075a
32 changed files with 1533 additions and 57 deletions
+2
View File
@@ -69,6 +69,8 @@ struct AnimaDiffusionExtra {
struct WanDiffusionExtra {
const sd::Tensor<float>* vace_context = nullptr;
float vace_strength = 1.f;
// S2V audio, sd::Tensor layout: [dim, T_latent*4, layers].
const sd::Tensor<float>* audio_embed = nullptr;
};
struct HiDreamO1DiffusionExtra {
+158 -32
View File
@@ -1,6 +1,7 @@
#ifndef __SD_MODEL_DIFFUSION_WAN_HPP__
#define __SD_MODEL_DIFFUSION_WAN_HPP__
#include <algorithm>
#include <cinttypes>
#include <map>
#include <memory>
@@ -33,11 +34,16 @@ namespace WAN {
int vace_layers = 0;
int64_t vace_in_dim = 96;
std::map<int, int> vace_layers_mapping = {};
bool qk_norm = true;
bool cross_attn_norm = true;
float eps = 1e-6f;
int64_t flf_pos_embed_token_number = 0;
int theta = 10000;
int64_t audio_dim = 1024;
int num_audio_token = 4; // excludes the learned padding token
std::vector<int> audio_inject_layers = {};
std::map<int, int> audio_inject_mapping = {}; // block index -> injector index
std::string adain_mode = "attn_norm";
bool qk_norm = true;
bool cross_attn_norm = true;
float eps = 1e-6f;
int64_t flf_pos_embed_token_number = 0;
int theta = 10000;
// wan2.1 1.3B: 1536/12, wan2.1/2.2 14B: 5120/40, wan2.2 5B: 3074/24
std::vector<int> axes_dim = {44, 42, 42};
int64_t axes_dim_sum = 128;
@@ -74,6 +80,10 @@ namespace WAN {
if (name.find("img_emb") != std::string::npos) {
config.model_type = "i2v";
}
if (name.find("audio_injector") != std::string::npos || name.find("casual_audio_encoder") != std::string::npos) {
config.model_type = "s2v";
config.audio_inject_layers = {0, 4, 8, 12, 16, 20, 24, 27, 30, 33, 36, 39};
}
if (name.find("img_emb.emb_pos") != std::string::npos) {
config.flf_pos_embed_token_number = 514;
}
@@ -265,6 +275,13 @@ namespace WAN {
}
};
} // namespace WAN
// Audio injection reuses WanT2VCrossAttention defined above.
#include "model/diffusion/wan_audio.hpp"
namespace WAN {
static ggml_tensor* modulate_add(ggml_context* ctx, ggml_tensor* x, ggml_tensor* e) {
// x: [N, n_token, dim]
// e: [N, 1, dim] or [N, T, 1, dim]
@@ -532,6 +549,13 @@ namespace WAN {
protected:
WanConfig config;
void init_params(ggml_context* ctx, const String2TensorStorage& tensor_storage_map = {}, const std::string prefix = "") override {
if (config.model_type == "s2v") {
enum ggml_type wtype = GGML_TYPE_F32; // elementwise add vs F32 activations
params["trainable_cond_mask.weight"] = ggml_new_tensor_2d(ctx, wtype, config.dim, 3);
}
}
public:
Wan() {}
Wan(WanConfig config)
@@ -554,7 +578,7 @@ namespace WAN {
// blocks
for (int i = 0; i < config.num_layers; i++) {
auto block = std::shared_ptr<GGMLBlock>(new WanAttentionBlock(config.model_type == "t2v",
auto block = std::shared_ptr<GGMLBlock>(new WanAttentionBlock(config.model_type != "i2v",
config.dim,
config.ffn_dim,
config.num_heads,
@@ -595,6 +619,14 @@ namespace WAN {
blocks["vace_patch_embedding"] = std::shared_ptr<GGMLBlock>(new Conv3d(config.vace_in_dim, config.dim, config.patch_size, config.patch_size));
}
if (config.model_type == "s2v") {
blocks["casual_audio_encoder"] = std::make_shared<WanCausalAudioEncoder>(config.audio_dim, config.dim, config.num_audio_token);
blocks["audio_injector"] = std::make_shared<WanAudioInjector>(config.dim, config.num_heads, (int)config.audio_inject_layers.size(), config.qk_norm, config.eps);
for (size_t i = 0; i < config.audio_inject_layers.size(); i++) {
config.audio_inject_mapping[config.audio_inject_layers[i]] = (int)i;
}
}
}
ggml_tensor* pad_to_patch_size(GGMLRunnerContext* ctx,
@@ -642,18 +674,24 @@ namespace WAN {
ggml_tensor* timestep,
ggml_tensor* context,
ggml_tensor* pe,
ggml_tensor* clip_fea = nullptr,
ggml_tensor* vace_context = nullptr,
float vace_strength = 1.f,
int64_t N = 1) {
ggml_tensor* clip_fea = nullptr,
ggml_tensor* vace_context = nullptr,
float vace_strength = 1.f,
int64_t N = 1,
ggml_tensor* audio_embed = nullptr,
ggml_tensor* reference_latent = nullptr) {
// x: [N*C, T, H, W], C => in_dim
// vace_context: [N*vace_in_dim, T, H, W]
// timestep: [N,] or [T]
// context: [N, L, text_dim]
// return: [N, t_len*h_len*w_len, out_dim*pt*ph*pw]
// audio_embed: [layers, T*4, audio_dim]
// reference_latent: [N*C, T_ref, H, W]
// return: [N, (t_len [+ t_ref_len]) * h_len*w_len, out_dim*pt*ph*pw]
GGML_ASSERT(N == 1);
int64_t T = x->ne[2];
auto patch_embedding = std::dynamic_pointer_cast<Conv3d>(blocks["patch_embedding"]);
auto text_embedding_0 = std::dynamic_pointer_cast<Linear>(blocks["text_embedding.0"]);
@@ -670,6 +708,40 @@ namespace WAN {
x = ggml_reshape_3d(ctx->ggml_ctx, x, x->ne[0] * x->ne[1] * x->ne[2], x->ne[3] / N, N); // [N, dim, t_len*h_len*w_len]
x = ggml_ext_cont(ctx->ggml_ctx, ggml_ext_torch_permute(ctx->ggml_ctx, x, 1, 0, 2, 3)); // [N, t_len*h_len*w_len, dim]
ggml_tensor* audio_local = nullptr;
ggml_tensor* audio_global = nullptr;
int64_t seq_len = x->ne[1];
int64_t t_ref_len = 0;
if (config.model_type == "s2v") {
if (audio_embed != nullptr) {
GGML_ASSERT(audio_embed->ne[1] == T * 4);
auto audio_encoder = std::dynamic_pointer_cast<WanCausalAudioEncoder>(blocks["casual_audio_encoder"]);
auto audio_emb = audio_encoder->forward(ctx, audio_embed);
audio_local = audio_emb.first;
audio_global = audio_emb.second;
GGML_ASSERT(audio_local->ne[2] == T);
}
// video tokens get cond_mask[0], reference tokens cond_mask[1]
auto cond_mask = params["trainable_cond_mask.weight"];
auto cm0 = ggml_reshape_3d(ctx->ggml_ctx, ggml_ext_slice(ctx->ggml_ctx, cond_mask, 1, 0, 1), config.dim, 1, 1);
x = ggml_add(ctx->ggml_ctx, x, cm0);
if (reference_latent != nullptr) {
t_ref_len = reference_latent->ne[2];
auto ref = patch_embedding->forward(ctx, reference_latent);
ref = ggml_reshape_3d(ctx->ggml_ctx, ref, ref->ne[0] * ref->ne[1] * ref->ne[2], ref->ne[3] / N, N);
ref = ggml_ext_cont(ctx->ggml_ctx, ggml_ext_torch_permute(ctx->ggml_ctx, ref, 1, 0, 2, 3)); // [N, t_ref*h_len*w_len, dim]
auto cm1 = ggml_reshape_3d(ctx->ggml_ctx, ggml_ext_slice(ctx->ggml_ctx, cond_mask, 1, 1, 2), config.dim, 1, 1);
ref = ggml_add(ctx->ggml_ctx, ref, cm1);
x = ggml_concat(ctx->ggml_ctx, x, ref, 1);
// Reference tokens use timestep 0.
GGML_ASSERT(timestep->ne[0] == T);
timestep = ggml_ext_pad(ctx->ggml_ctx, timestep, (int)t_ref_len, 0, 0, 0);
}
}
// time_embedding
auto e = ggml_ext_timestep_embedding(ctx->ggml_ctx, timestep, config.freq_dim);
e = time_embedding_0->forward(ctx, e);
@@ -714,6 +786,11 @@ namespace WAN {
auto x_orig = x;
std::shared_ptr<WanAudioInjector> audio_injector;
if (audio_local != nullptr) {
audio_injector = std::dynamic_pointer_cast<WanAudioInjector>(blocks["audio_injector"]);
}
for (int i = 0; i < config.num_layers; i++) {
auto block = std::dynamic_pointer_cast<WanAttentionBlock>(blocks["blocks." + std::to_string(i)]);
@@ -731,6 +808,13 @@ namespace WAN {
c_skip = ggml_ext_scale(ctx->ggml_ctx, c_skip, vace_strength);
x = ggml_add(ctx->ggml_ctx, x, c_skip);
}
if (audio_injector != nullptr) {
auto inject_iter = config.audio_inject_mapping.find(i);
if (inject_iter != config.audio_inject_mapping.end()) {
x = audio_injector->forward(ctx, x, seq_len, T, inject_iter->second, audio_local, audio_global);
}
}
sd::ggml_graph_cut::mark_graph_cut(x, "wan.blocks." + std::to_string(i), "x");
if (c != nullptr) {
sd::ggml_graph_cut::mark_graph_cut(c, "wan.blocks." + std::to_string(i), "c");
@@ -747,11 +831,13 @@ namespace WAN {
ggml_tensor* timestep,
ggml_tensor* context,
ggml_tensor* pe,
ggml_tensor* clip_fea = nullptr,
ggml_tensor* time_dim_concat = nullptr,
ggml_tensor* vace_context = nullptr,
float vace_strength = 1.f,
int64_t N = 1) {
ggml_tensor* clip_fea = nullptr,
ggml_tensor* time_dim_concat = nullptr,
ggml_tensor* vace_context = nullptr,
float vace_strength = 1.f,
int64_t N = 1,
ggml_tensor* audio_embed = nullptr,
ggml_tensor* reference_latent = nullptr) {
// Forward pass of DiT.
// x: [N*C, T, H, W]
// timestep: [N,]
@@ -779,7 +865,12 @@ namespace WAN {
t_len = ((x->ne[2] + (std::get<0>(config.patch_size) / 2)) / std::get<0>(config.patch_size));
}
auto out = forward_orig(ctx, x, timestep, context, pe, clip_fea, vace_context, vace_strength, N); // [N, t_len*h_len*w_len, pt*ph*pw*C]
auto out = forward_orig(ctx, x, timestep, context, pe, clip_fea, vace_context, vace_strength, N, audio_embed, reference_latent); // [N, (t_len [+t_ref]) *h_len*w_len, pt*ph*pw*C]
if (reference_latent != nullptr) {
// Exclude reference tokens from the generated video.
out = ggml_ext_slice(ctx->ggml_ctx, out, 1, 0, t_len * h_len * w_len);
}
out = unpatchify(ctx->ggml_ctx, out, t_len, h_len, w_len); // [N*C, (T+pad_t) + (T2+pad_t2), H + pad_h, W + pad_w]
@@ -839,7 +930,10 @@ namespace WAN {
config.text_len = 512;
}
} else if (config.num_layers == 40) {
if (config.model_type == "t2v") {
if (version == VERSION_WAN2_2_S2V) {
desc = "Wan2.2-S2V-14B";
config.in_dim = 16;
} else if (config.model_type == "t2v") {
if (version == VERSION_WAN2_2_I2V) {
desc = "Wan2.2-I2V-14B";
config.in_dim = 36;
@@ -891,7 +985,9 @@ namespace WAN {
const sd::Tensor<float>& c_concat_tensor = {},
const sd::Tensor<float>& time_dim_concat_tensor = {},
const sd::Tensor<float>& vace_context_tensor = {},
float vace_strength = 1.f) {
float vace_strength = 1.f,
const sd::Tensor<float>& audio_embed_tensor = {},
const sd::Tensor<float>& ref_latent_tensor = {}) {
ggml_cgraph* gf = new_graph_custom(WAN_GRAPH_SIZE);
ggml_tensor* x = make_input(x_tensor);
@@ -901,16 +997,33 @@ namespace WAN {
ggml_tensor* c_concat = make_optional_input(c_concat_tensor);
ggml_tensor* time_dim_concat = make_optional_input(time_dim_concat_tensor);
ggml_tensor* vace_context = make_optional_input(vace_context_tensor);
ggml_tensor* audio_embed = make_optional_input(audio_embed_tensor);
ggml_tensor* ref_latent = make_optional_input(ref_latent_tensor);
pe_vec = Rope::gen_wan_pe(static_cast<int>(x->ne[2]),
static_cast<int>(x->ne[1]),
static_cast<int>(x->ne[0]),
std::get<0>(config.patch_size),
std::get<1>(config.patch_size),
std::get<2>(config.patch_size),
1,
config.theta,
config.axes_dim);
pe_vec = Rope::gen_wan_pe(static_cast<int>(x->ne[2]),
static_cast<int>(x->ne[1]),
static_cast<int>(x->ne[0]),
std::get<0>(config.patch_size),
std::get<1>(config.patch_size),
std::get<2>(config.patch_size),
1,
config.theta,
config.axes_dim);
if (ref_latent != nullptr) {
// Match S2V's reference-frame temporal offset.
int t_start = std::max(30, static_cast<int>(x->ne[2]) + 9);
auto ref_pe = Rope::gen_wan_pe(static_cast<int>(ref_latent->ne[2]),
static_cast<int>(ref_latent->ne[1]),
static_cast<int>(ref_latent->ne[0]),
std::get<0>(config.patch_size),
std::get<1>(config.patch_size),
std::get<2>(config.patch_size),
1,
config.theta,
config.axes_dim,
t_start);
pe_vec.insert(pe_vec.end(), ref_pe.begin(), ref_pe.end());
}
int pos_len = static_cast<int>(pe_vec.size() / config.axes_dim_sum / 2);
// LOG_VERBOSE("pos_len %d", pos_len);
auto pe = ggml_new_tensor_4d(compute_ctx, GGML_TYPE_F32, 2, 2, config.axes_dim_sum / 2, pos_len);
@@ -933,7 +1046,10 @@ namespace WAN {
clip_fea,
time_dim_concat,
vace_context,
vace_strength);
vace_strength,
1,
audio_embed,
ref_latent);
ggml_build_forward_expand(gf, out);
@@ -948,9 +1064,11 @@ namespace WAN {
const sd::Tensor<float>& c_concat = {},
const sd::Tensor<float>& time_dim_concat = {},
const sd::Tensor<float>& vace_context = {},
float vace_strength = 1.f) {
float vace_strength = 1.f,
const sd::Tensor<float>& audio_embed = {},
const sd::Tensor<float>& ref_latent = {}) {
auto get_graph = [&]() -> ggml_cgraph* {
return build_graph(x, timesteps, context, clip_fea, c_concat, time_dim_concat, vace_context, vace_strength);
return build_graph(x, timesteps, context, clip_fea, c_concat, time_dim_concat, vace_context, vace_strength, audio_embed, ref_latent);
};
return restore_trailing_singleton_dims(GGMLRunner::compute(get_graph, n_threads, false), x.dim());
@@ -961,6 +1079,12 @@ namespace WAN {
GGML_ASSERT(diffusion_params.x != nullptr);
GGML_ASSERT(diffusion_params.timesteps != nullptr);
const auto* extra = diffusion_extra_as<WanDiffusionExtra>(diffusion_params);
static const std::vector<sd::Tensor<float>> no_ref_latents;
const auto& ref_latents = config.model_type == "s2v" && diffusion_params.ref_latents != nullptr
? *diffusion_params.ref_latents
: no_ref_latents;
const sd::Tensor<float> empty_tensor;
const sd::Tensor<float>& ref_latent = ref_latents.empty() ? empty_tensor : ref_latents[0];
return compute(n_threads,
*diffusion_params.x,
*diffusion_params.timesteps,
@@ -969,7 +1093,9 @@ namespace WAN {
tensor_or_empty(diffusion_params.c_concat),
sd::Tensor<float>(),
tensor_or_empty(extra->vace_context),
extra->vace_strength);
extra->vace_strength,
tensor_or_empty(extra->audio_embed),
ref_latent);
}
void test() {
+215
View File
@@ -0,0 +1,215 @@
#ifndef __SD_MODEL_DIFFUSION_WAN_AUDIO_HPP__
#define __SD_MODEL_DIFFUSION_WAN_AUDIO_HPP__
#include <cstdint>
#include <memory>
#include <string>
#include <utility>
#include <vector>
#include "model/common/ggml_block.hpp"
namespace WAN {
class WanCausalConv1d : public UnaryBlock {
private:
int kernel_size_;
public:
WanCausalConv1d(int64_t in_dim,
int64_t out_dim,
int kernel_size = 3,
int stride = 1)
: kernel_size_(kernel_size) {
blocks["conv"] = std::make_shared<Conv1d>(in_dim, out_dim, kernel_size, stride, 0, 1, 1, true, true);
}
ggml_tensor* forward(GGMLRunnerContext* ctx, ggml_tensor* x) override {
// Replicate the first sample for causal left padding.
if (kernel_size_ > 1) {
auto first = ggml_ext_slice(ctx->ggml_ctx, x, 0, 0, 1);
for (int i = 0; i < kernel_size_ - 1; i++) {
x = ggml_concat(ctx->ggml_ctx, first, x, 0);
}
}
return std::dynamic_pointer_cast<Conv1d>(blocks["conv"])->forward(ctx, x);
}
};
class WanMotionEncoder : public GGMLBlock {
private:
int64_t hidden_dim_;
int num_token_;
bool need_global_;
void init_params(ggml_context* ctx, const String2TensorStorage& tensor_storage_map = {}, const std::string prefix = "") override {
// The padding token is combined with F32 activations.
params["padding_tokens"] = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, hidden_dim_);
}
ggml_tensor* conv_norm_silu(GGMLRunnerContext* ctx,
ggml_tensor* x,
const std::string& conv_key,
const std::string& norm_key,
bool to_conv_layout) {
x = std::dynamic_pointer_cast<WanCausalConv1d>(blocks[conv_key])->forward(ctx, x);
x = ggml_permute(ctx->ggml_ctx, x, 1, 0, 2, 3);
x = std::dynamic_pointer_cast<LayerNorm>(blocks[norm_key])->forward(ctx, x);
x = ggml_silu(ctx->ggml_ctx, x);
if (to_conv_layout) {
x = ggml_ext_cont(ctx->ggml_ctx, ggml_permute(ctx->ggml_ctx, x, 1, 0, 2, 3));
}
return x;
}
public:
WanMotionEncoder(int64_t in_dim,
int64_t hidden_dim,
int num_token,
bool need_global = true)
: hidden_dim_(hidden_dim), num_token_(num_token), need_global_(need_global) {
blocks["conv1_local"] = std::make_shared<WanCausalConv1d>(in_dim, hidden_dim / 4 * num_token);
if (need_global) {
blocks["conv1_global"] = std::make_shared<WanCausalConv1d>(in_dim, hidden_dim / 4);
}
blocks["norm1"] = std::make_shared<LayerNorm>(hidden_dim / 4, 1e-6f, false);
blocks["conv2"] = std::make_shared<WanCausalConv1d>(hidden_dim / 4, hidden_dim / 2, 3, 2);
blocks["norm2"] = std::make_shared<LayerNorm>(hidden_dim / 2, 1e-6f, false);
blocks["conv3"] = std::make_shared<WanCausalConv1d>(hidden_dim / 2, hidden_dim, 3, 2);
blocks["norm3"] = std::make_shared<LayerNorm>(hidden_dim, 1e-6f, false);
if (need_global) {
blocks["final_linear"] = std::make_shared<Linear>(hidden_dim, hidden_dim);
}
}
std::pair<ggml_tensor*, ggml_tensor*> forward(GGMLRunnerContext* ctx, ggml_tensor* x) {
auto local = std::dynamic_pointer_cast<WanCausalConv1d>(blocks["conv1_local"])->forward(ctx, x);
auto norm1 = std::dynamic_pointer_cast<LayerNorm>(blocks["norm1"]);
std::vector<ggml_tensor*> tokens;
// Each token group is normalized independently over channels.
for (auto& group : ggml_ext_chunk(ctx->ggml_ctx, local, num_token_, 1)) {
ggml_tensor* s = ggml_permute(ctx->ggml_ctx, group, 1, 0, 2, 3);
s = norm1->forward(ctx, s);
s = ggml_silu(ctx->ggml_ctx, s);
s = ggml_ext_cont(ctx->ggml_ctx, ggml_permute(ctx->ggml_ctx, s, 1, 0, 2, 3));
s = conv_norm_silu(ctx, s, "conv2", "norm2", true);
s = conv_norm_silu(ctx, s, "conv3", "norm3", false);
tokens.push_back(ggml_reshape_3d(ctx->ggml_ctx, s, s->ne[0], 1, s->ne[1]));
}
auto padding = ggml_reshape_3d(ctx->ggml_ctx, params["padding_tokens"], hidden_dim_, 1, 1);
padding = ggml_repeat(ctx->ggml_ctx, padding, tokens[0]);
tokens.push_back(padding);
ggml_tensor* local_out = ggml_ext_vec_concat(ctx->ggml_ctx, tokens, 1);
if (!need_global_) {
return {local_out, nullptr};
}
ggml_tensor* g = conv_norm_silu(ctx, x, "conv1_global", "norm1", true);
g = conv_norm_silu(ctx, g, "conv2", "norm2", true);
g = conv_norm_silu(ctx, g, "conv3", "norm3", false);
g = std::dynamic_pointer_cast<Linear>(blocks["final_linear"])->forward(ctx, g);
return {local_out, g};
}
};
class WanCausalAudioEncoder : public GGMLBlock {
private:
int num_layers_;
void init_params(ggml_context* ctx, const String2TensorStorage& tensor_storage_map = {}, const std::string prefix = "") override {
// Preserve the checkpoint shape for loading; layer mixing requires F32.
auto it = tensor_storage_map.find(prefix + "weights");
if (it != tensor_storage_map.end()) {
params["weights"] = ggml_new_tensor(ctx, GGML_TYPE_F32, it->second.n_dims, it->second.ne);
} else {
params["weights"] = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, num_layers_);
}
}
public:
WanCausalAudioEncoder(int64_t audio_dim,
int64_t dim,
int num_token,
int num_layers = 25)
: num_layers_(num_layers) {
blocks["encoder"] = std::make_shared<WanMotionEncoder>(audio_dim, dim, num_token, true);
}
// features: [layers, frames, audio_dim]; outputs: [T, tokens+1, dim] and [T, dim].
std::pair<ggml_tensor*, ggml_tensor*> forward(GGMLRunnerContext* ctx, ggml_tensor* features) {
auto weights = ggml_silu(ctx->ggml_ctx, params["weights"]);
auto x = ggml_mul(ctx->ggml_ctx, features, ggml_reshape_3d(ctx->ggml_ctx, weights, 1, 1, num_layers_));
x = ggml_div(ctx->ggml_ctx, x, ggml_sum(ctx->ggml_ctx, weights));
// Move the layer axis to ggml dimension 0 for reduction.
x = ggml_ext_cont(ctx->ggml_ctx, ggml_ext_torch_permute(ctx->ggml_ctx, x, 2, 0, 1, 3));
x = ggml_sum_rows(ctx->ggml_ctx, x);
x = ggml_reshape_2d(ctx->ggml_ctx, x, x->ne[1], x->ne[2]);
x = ggml_ext_cont(ctx->ggml_ctx, ggml_ext_torch_permute(ctx->ggml_ctx, x, 1, 0, 2, 3));
return std::dynamic_pointer_cast<WanMotionEncoder>(blocks["encoder"])->forward(ctx, x);
}
};
class WanAudioInjector : public GGMLBlock {
private:
int64_t dim_;
public:
WanAudioInjector(int64_t dim,
int64_t num_heads,
int count,
bool qk_norm = true,
float eps = 1e-6f)
: dim_(dim) {
for (int i = 0; i < count; i++) {
blocks["injector." + std::to_string(i)] =
std::make_shared<WanT2VCrossAttention>(dim, num_heads, qk_norm, eps);
blocks["injector_adain_layers." + std::to_string(i) + ".linear"] =
std::make_shared<Linear>(dim, dim * 2);
}
// S2V AdaLayerNorm uses its own epsilon, independent of attention norms.
blocks["adain_norm"] = std::make_shared<LayerNorm>(dim, 1e-5f, false);
}
// Inject into the video prefix; trailing reference tokens pass through unchanged.
ggml_tensor* forward(GGMLRunnerContext* ctx,
ggml_tensor* x,
int64_t seq_len,
int64_t T,
int injector_id,
ggml_tensor* audio_local,
ggml_tensor* audio_global) {
int64_t n_tok = seq_len / T;
int64_t n_token = x->ne[1];
auto adain_linear = std::dynamic_pointer_cast<Linear>(blocks["injector_adain_layers." + std::to_string(injector_id) + ".linear"]);
auto injector = std::dynamic_pointer_cast<WanT2VCrossAttention>(blocks["injector." + std::to_string(injector_id)]);
auto adain_norm = std::dynamic_pointer_cast<LayerNorm>(blocks["adain_norm"]);
auto temb = ggml_silu(ctx->ggml_ctx, audio_global);
temb = adain_linear->forward(ctx, temb);
auto shift = ggml_ext_slice(ctx->ggml_ctx, temb, 0, 0, dim_);
auto scale = ggml_ext_slice(ctx->ggml_ctx, temb, 0, dim_, dim_ * 2);
shift = ggml_reshape_3d(ctx->ggml_ctx, shift, dim_, 1, T);
scale = ggml_reshape_3d(ctx->ggml_ctx, scale, dim_, 1, T);
auto x_vid = ggml_ext_slice(ctx->ggml_ctx, x, 1, 0, seq_len);
auto h = ggml_reshape_3d(ctx->ggml_ctx, x_vid, dim_, n_tok, T);
h = adain_norm->forward(ctx, h);
h = ggml_add(ctx->ggml_ctx, h, ggml_mul(ctx->ggml_ctx, h, scale));
h = ggml_add(ctx->ggml_ctx, h, shift);
auto res = injector->forward(ctx, h, audio_local, 0);
res = ggml_reshape_2d(ctx->ggml_ctx, res, dim_, seq_len);
auto x_head = ggml_add(ctx->ggml_ctx, x_vid, res);
if (seq_len < n_token) {
auto x_tail = ggml_ext_slice(ctx->ggml_ctx, x, 1, seq_len, n_token);
return ggml_concat(ctx->ggml_ctx, x_head, x_tail, 1);
}
return x_head;
}
};
} // namespace WAN
#endif // __SD_MODEL_DIFFUSION_WAN_AUDIO_HPP__