| 367 | } |
| 368 | |
| 369 | MioCodecTransformerWeights bind_transformer( |
| 370 | MioCodecWeights & weights, |
| 371 | const engine::assets::TensorSource & source, |
| 372 | const std::string & prefix, |
| 373 | int64_t dim, |
| 374 | int64_t layers, |
| 375 | int64_t heads, |
| 376 | int64_t window_size, |
| 377 | bool use_adaln, |
| 378 | int64_t condition_dim = 0, |
| 379 | std::optional<int64_t> output_dim = std::nullopt, |
| 380 | engine::assets::TensorStorageType storage_type = engine::assets::TensorStorageType::F32) { |
| 381 | MioCodecTransformerWeights transformer; |
| 382 | transformer.dim = dim; |
| 383 | transformer.heads = heads; |
| 384 | transformer.head_dim = dim / heads; |
| 385 | transformer.window_size = window_size; |
| 386 | transformer.intermediate_dim = swiglu_hidden_dim(dim); |
| 387 | transformer.use_adaln = use_adaln; |
| 388 | transformer.layers.reserve(static_cast<size_t>(layers)); |
| 389 | for (int64_t layer = 0; layer < layers; ++layer) { |
| 390 | const std::string layer_prefix = prefix + ".layers." + std::to_string(layer); |
| 391 | MioCodecTransformerLayerWeights w; |
| 392 | w.qkv_proj = bind_qkv_linear(weights, source, layer_prefix, dim, storage_type); |
| 393 | w.out_proj = bind_linear(weights, layer_prefix + ".attention.wo", dim, dim, false); |
| 394 | w.feed_forward_w1 = bind_linear(weights, layer_prefix + ".feed_forward.w1", dim, transformer.intermediate_dim, false); |
| 395 | w.feed_forward_w2 = bind_linear(weights, layer_prefix + ".feed_forward.w2", transformer.intermediate_dim, dim, false); |
| 396 | w.feed_forward_w3 = bind_linear(weights, layer_prefix + ".feed_forward.w3", dim, transformer.intermediate_dim, false); |
| 397 | if (use_adaln) { |
| 398 | w.attention_adaln = bind_adaln(weights, layer_prefix + ".attention_norm", condition_dim, 3 * dim); |
| 399 | w.feed_forward_adaln = bind_adaln(weights, layer_prefix + ".ffn_norm", condition_dim, 3 * dim); |
| 400 | } else { |
| 401 | w.attention_norm = bind_norm(weights, layer_prefix + ".attention_norm", dim); |
| 402 | w.feed_forward_norm = bind_norm(weights, layer_prefix + ".ffn_norm", dim); |
| 403 | } |
| 404 | transformer.layers.push_back(std::move(w)); |
| 405 | } |
| 406 | if (use_adaln) { |
| 407 | transformer.adaln_norm = bind_adaln(weights, prefix + ".norm", condition_dim, 2 * dim); |
| 408 | } else { |
| 409 | transformer.norm = bind_norm(weights, prefix + ".norm", dim); |
| 410 | } |
| 411 | if (output_dim.has_value()) { |
| 412 | transformer.output_projection = bind_linear(weights, prefix + ".output_proj", dim, *output_dim); |
| 413 | } |
| 414 | return transformer; |
| 415 | } |
| 416 | |
| 417 | MioCodecSnakeBetaWeights bind_snake_beta( |
| 418 | MioCodecWeights & weights, |
no test coverage detected