| 341 | } |
| 342 | |
| 343 | core::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 | |
| 389 | StableAudioRfAttentionWeights load_self_attention( |
| 390 | core::BackendWeightStore & store, |
no test coverage detected