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

Function load_diffusion_weights

src/models/ace_step/dit_weights_runtime.cpp:228–289  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

226}
227
228std::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",

Callers 1

Calls 7

load_tensorMethod · 0.80
load_f32_tensorMethod · 0.80
to_stringFunction · 0.50
make_f32Method · 0.45
require_f32Method · 0.45

Tested by

no test coverage detected