| 326 | encoder_hidden_states); |
| 327 | x = modules::AddModule{}.build(ctx, x, cross_out); |
| 328 | |
| 329 | auto mlp_norm = apply_modulated_rms_norm( |
| 330 | ctx, |
| 331 | x, |
| 332 | weights.mlp_norm, |
| 333 | one, |
| 334 | c_shift_msa, |
| 335 | c_scale_msa, |
| 336 | config.rms_norm_eps); |
| 337 | auto ff = build_mlp(ctx, mlp_norm, weights.mlp_gate, weights.mlp_up, weights.mlp_down, config); |
| 338 | ff = modules::MulModule{}.build(ctx, ff, c_gate_msa); |
| 339 | return modules::AddModule{}.build(ctx, x, ff); |
| 340 | } |
| 341 | |
| 342 | std::vector<float> build_cross_attention_mask_values( |
| 343 | int64_t seq_len, |
| 344 | const std::vector<int32_t> & encoder_attention_mask, |
| 345 | int64_t valid_seq_len) { |
| 346 | const int64_t encoder_tokens = static_cast<int64_t>(encoder_attention_mask.size()); |
| 347 | if (seq_len <= 0 || encoder_tokens <= 0 || valid_seq_len <= 0 || valid_seq_len > seq_len) { |
| 348 | throw std::runtime_error("ACE-Step diffusion cross attention mask shape is invalid"); |
| 349 | } |
| 350 | std::vector<float> values(static_cast<size_t>(seq_len * encoder_tokens), 0.0F); |
| 351 | const float masked = std::numeric_limits<float>::lowest(); |
| 352 | for (int64_t q = 0; q < seq_len; ++q) { |
| 353 | for (int64_t k = 0; k < encoder_tokens; ++k) { |
| 354 | if (q >= valid_seq_len) { |
| 355 | values[static_cast<size_t>(q * encoder_tokens + k)] = masked; |
| 356 | } |
| 357 | } |
| 358 | } |
| 359 | return values; |
| 360 | } |
| 361 | |
| 362 | std::vector<float> build_sliding_mask_values(int64_t tokens, int64_t sliding_window) { |
| 363 | std::vector<float> values(static_cast<size_t>(tokens * tokens), 0.0F); |
| 364 | const float masked = std::numeric_limits<float>::lowest(); |
| 365 | for (int64_t q = 0; q < tokens; ++q) { |
| 366 | for (int64_t k = 0; k < tokens; ++k) { |
| 367 | if (std::llabs(q - k) > sliding_window) { |
| 368 | values[static_cast<size_t>(q * tokens + k)] = masked; |
| 369 | } |
| 370 | } |
| 371 | } |
| 372 | return values; |
| 373 | } |
| 374 | |
| 375 | std::vector<float> build_self_attention_padding_mask_values(int64_t seq_len, int64_t valid_seq_len) { |
| 376 | if (seq_len <= 0 || valid_seq_len <= 0 || valid_seq_len > seq_len) { |
| 377 | throw std::runtime_error("ACE-Step diffusion self attention padding mask is invalid"); |
| 378 | } |
| 379 | std::vector<float> values(static_cast<size_t>(seq_len * seq_len), 0.0F); |
| 380 | const float masked = std::numeric_limits<float>::lowest(); |
| 381 | for (int64_t q = 0; q < valid_seq_len; ++q) { |
| 382 | for (int64_t k = valid_seq_len; k < seq_len; ++k) { |
| 383 | values[static_cast<size_t>(q * seq_len + k)] = masked; |
| 384 | } |
| 385 | } |
no test coverage detected