MCPcopy Create free account
hub / github.com/0xShug0/audio.cpp / build_rf_layer

Function build_rf_layer

src/models/stable_audio/rf_dit.cpp:343–387  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

341}
342
343core::TensorValue build_rf_layer(
344 core::ModuleBuildContext & ctx,
345 const core::TensorValue & input,
346 const core::TensorValue & cross_context,
347 const core::TensorValue & global_six,
348 const core::TensorValue & local_add,
349 const core::TensorValue & positions,
350 const core::TensorValue & padding_mask,
351 const StableAudioRfLayerWeights & weights,
352 const StableAudioConfig & config) {
353 auto gate_bias = core::reshape_tensor(ctx, weights.scale_shift_gate, core::TensorShape::from_dims({1, 6 * config.embed_dim}));
354 gate_bias = modules::RepeatModule({core::TensorShape::from_dims({input.shape.dims[0], 6 * config.embed_dim})}).build(ctx, gate_bias);
355 auto adaln = modules::AddModule{}.build(ctx, global_six, gate_bias);
356 adaln = core::reshape_tensor(ctx, adaln, core::TensorShape::from_dims({input.shape.dims[0], 1, 6 * config.embed_dim}));
357 auto scale_self = modules::SliceModule({2, 0, config.embed_dim}).build(ctx, adaln);
358 auto shift_self = modules::SliceModule({2, config.embed_dim, config.embed_dim}).build(ctx, adaln);
359 auto gate_self = modules::SliceModule({2, 2 * config.embed_dim, config.embed_dim}).build(ctx, adaln);
360 auto scale_ff = modules::SliceModule({2, 3 * config.embed_dim, config.embed_dim}).build(ctx, adaln);
361 auto shift_ff = modules::SliceModule({2, 4 * config.embed_dim, config.embed_dim}).build(ctx, adaln);
362 auto gate_ff = modules::SliceModule({2, 5 * config.embed_dim, config.embed_dim}).build(ctx, adaln);
363 scale_self = modules::RepeatModule({input.shape}).build(ctx, scale_self);
364 shift_self = modules::RepeatModule({input.shape}).build(ctx, shift_self);
365 gate_self = modules::RepeatModule({input.shape}).build(ctx, gate_self);
366 scale_ff = modules::RepeatModule({input.shape}).build(ctx, scale_ff);
367 shift_ff = modules::RepeatModule({input.shape}).build(ctx, shift_ff);
368 gate_ff = modules::RepeatModule({input.shape}).build(ctx, gate_ff);
369
370 auto hidden = rms_norm(ctx, input, weights.pre_norm_gamma, kTransformerNormEps);
371 hidden = modules::AddModule{}.build(ctx, modules::MulModule{}.build(ctx, hidden, one_plus(ctx, scale_self)), shift_self);
372 auto attn = self_attention(ctx, hidden, positions, padding_mask, weights.self_attn, config);
373 attn = modules::MulModule{}.build(ctx, attn, sigmoid_one_minus(ctx, gate_self));
374 hidden = modules::AddModule{}.build(ctx, input, attn);
375
376 auto cross_norm = rms_norm(ctx, hidden, weights.cross_attend_norm_gamma, kTransformerNormEps);
377 auto cross = cross_attention(ctx, cross_norm, cross_context, weights.cross_attn, config);
378 hidden = modules::AddModule{}.build(ctx, hidden, cross);
379 auto local = build_local_embedding(ctx, hidden, local_add, weights, config);
380 hidden = modules::AddModule{}.build(ctx, hidden, local);
381
382 auto ff_in = rms_norm(ctx, hidden, weights.ff_norm_gamma, kTransformerNormEps);
383 ff_in = modules::AddModule{}.build(ctx, modules::MulModule{}.build(ctx, ff_in, one_plus(ctx, scale_ff)), shift_ff);
384 auto ff = swiglu_ff(ctx, ff_in, weights.ff, config);
385 ff = modules::MulModule{}.build(ctx, ff, sigmoid_one_minus(ctx, gate_ff));
386 return modules::AddModule{}.build(ctx, hidden, ff);
387}
388
389StableAudioRfAttentionWeights load_self_attention(
390 core::BackendWeightStore & store,

Callers 1

build_graph_outputMethod · 0.70

Calls 11

reshape_tensorFunction · 0.85
RepeatModuleClass · 0.85
SliceModuleClass · 0.85
rms_normFunction · 0.85
one_plusFunction · 0.85
sigmoid_one_minusFunction · 0.85
build_local_embeddingFunction · 0.85
self_attentionFunction · 0.70
cross_attentionFunction · 0.70
swiglu_ffFunction · 0.70
buildMethod · 0.45

Tested by

no test coverage detected