| 226 | } |
| 227 | |
| 228 | std::shared_ptr<const AceStepDiffusionWeights> load_diffusion_weights( |
| 229 | const std::shared_ptr<core::BackendWeightStore> & store, |
| 230 | const AceStepAssets & assets, |
| 231 | assets::TensorStorageType storage_type) { |
| 232 | const auto & config = assets.config.diffusion; |
| 233 | const auto & source = *assets.dit_weights; |
| 234 | auto weights = std::make_shared<AceStepDiffusionWeights>(); |
| 235 | weights->store = store; |
| 236 | weights->one = store->make_f32(core::TensorShape::from_dims({1}), std::vector<float>{1.0F}); |
| 237 | weights->proj_in = { |
| 238 | store->load_tensor(source, "decoder.proj_in.1.weight", storage_type, {config.hidden_size, config.in_channels, config.patch_size}), |
| 239 | store->load_tensor(source, "decoder.proj_in.1.bias", assets::TensorStorageType::F32, {config.hidden_size}), |
| 240 | }; |
| 241 | weights->time_embed = load_time_embedding_weights(*store, source, "decoder.time_embed", storage_type, config.hidden_size); |
| 242 | weights->time_embed_r = load_time_embedding_weights(*store, source, "decoder.time_embed_r", storage_type, config.hidden_size); |
| 243 | weights->condition_embedder = { |
| 244 | store->load_tensor(source, "decoder.condition_embedder.weight", storage_type, {config.hidden_size, config.hidden_size}), |
| 245 | store->load_tensor(source, "decoder.condition_embedder.bias", assets::TensorStorageType::F32, {config.hidden_size}), |
| 246 | }; |
| 247 | if (!config.is_turbo) { |
| 248 | weights->null_condition_emb_host = source.require_f32("null_condition_emb", {1, 1, config.hidden_size}); |
| 249 | } |
| 250 | weights->layers.reserve(static_cast<size_t>(config.num_hidden_layers)); |
| 251 | for (int64_t i = 0; i < config.num_hidden_layers; ++i) { |
| 252 | const std::string prefix = "decoder.layers." + std::to_string(i); |
| 253 | AceStepDiTLayerWeights layer; |
| 254 | layer.self_attn_norm = store->load_f32_tensor(source, prefix + ".self_attn_norm.weight", {config.hidden_size}); |
| 255 | layer.self_attn = load_dit_attention_weights(*store, source, prefix + ".self_attn", storage_type, config); |
| 256 | layer.cross_attn_norm = store->load_f32_tensor(source, prefix + ".cross_attn_norm.weight", {config.hidden_size}); |
| 257 | layer.cross_attn = load_dit_attention_weights(*store, source, prefix + ".cross_attn", storage_type, config); |
| 258 | layer.mlp_norm = store->load_f32_tensor(source, prefix + ".mlp_norm.weight", {config.hidden_size}); |
| 259 | layer.mlp_gate = { |
| 260 | store->load_tensor(source, prefix + ".mlp.gate_proj.weight", storage_type, {config.intermediate_size, config.hidden_size}), |
| 261 | std::nullopt, |
| 262 | }; |
| 263 | layer.mlp_up = { |
| 264 | store->load_tensor(source, prefix + ".mlp.up_proj.weight", storage_type, {config.intermediate_size, config.hidden_size}), |
| 265 | std::nullopt, |
| 266 | }; |
| 267 | layer.mlp_down = { |
| 268 | store->load_tensor(source, prefix + ".mlp.down_proj.weight", storage_type, {config.hidden_size, config.intermediate_size}), |
| 269 | std::nullopt, |
| 270 | }; |
| 271 | layer.scale_shift_table = store->load_tensor( |
| 272 | source, |
| 273 | prefix + ".scale_shift_table", |
| 274 | storage_type, |
| 275 | {1, 6, config.hidden_size}); |
| 276 | weights->layers.push_back(std::move(layer)); |
| 277 | } |
| 278 | weights->norm_out = store->load_f32_tensor(source, "decoder.norm_out.weight", {config.hidden_size}); |
| 279 | weights->proj_out = { |
| 280 | store->load_tensor(source, "decoder.proj_out.1.weight", storage_type, {config.hidden_size, config.latent_channels, config.patch_size}), |
| 281 | store->load_tensor(source, "decoder.proj_out.1.bias", assets::TensorStorageType::F32, {config.latent_channels}), |
| 282 | }; |
| 283 | weights->final_scale_shift_table = store->load_tensor( |
| 284 | source, |
| 285 | "decoder.scale_shift_table", |
no test coverage detected