| 61 | config.same_decoder_sinusoidal_blocks = json::optional_i64_array(decoder_config, "sinusoidal_blocks", std::vector<int64_t>(config.same_decoder_transformer_depths.size(), 0)); |
| 62 | config.same_sliding_window = json::optional_i64_array(encoder_config, "sliding_window"); |
| 63 | config.same_differential = json::optional_bool(encoder_config, "differential", true); |
| 64 | if (config.same_differential != json::optional_bool(decoder_config, "differential", config.same_differential)) { |
| 65 | throw std::runtime_error("Stable Audio SAME encoder/decoder differential settings must match"); |
| 66 | } |
| 67 | config.same_variable_stride = json::optional_bool(encoder_config, "variable_stride", false); |
| 68 | if (config.same_variable_stride != json::optional_bool(decoder_config, "variable_stride", config.same_variable_stride)) { |
| 69 | throw std::runtime_error("Stable Audio SAME encoder/decoder variable_stride settings must match"); |
| 70 | } |
| 71 | config.same_encoder_conv_mapping = json::optional_bool(encoder_config, "conv_mapping", false); |
| 72 | config.same_decoder_conv_mapping = json::optional_bool(decoder_config, "conv_mapping", false); |
| 73 | config.same_chunk_midpoint_shift = |
| 74 | json::optional_bool(encoder_config, "chunk_midpoint_shift", false) || |
| 75 | json::optional_bool(decoder_config, "chunk_midpoint_shift", false); |
| 76 | config.same_encoder_mask_noise = json::optional_f32(encoder_config, "mask_noise", 0.0F); |
| 77 | config.same_decoder_mask_noise = json::optional_f32(decoder_config, "mask_noise", 0.0F); |
| 78 | const auto & bottleneck = pretransform_config.require("bottleneck"); |
| 79 | if (const auto * bottleneck_config = bottleneck.find("config"); bottleneck_config != nullptr && bottleneck_config->is_object()) { |
| 80 | config.same_bottleneck_noise_regularize = json::optional_bool(*bottleneck_config, "noise_regularize", false); |
| 81 | config.same_bottleneck_auto_scale = json::optional_bool(*bottleneck_config, "auto_scale", false); |
| 82 | config.same_bottleneck_noise_augment_dim = json::optional_i64(*bottleneck_config, "noise_augment_dim", 0); |
| 83 | } |
| 84 | |
| 85 | const auto & conditioning = model.require("conditioning"); |
| 86 | config.cond_dim = json::require_i64(conditioning, "cond_dim"); |
| 87 | for (const auto & item : conditioning.require("configs").as_array()) { |
| 88 | const std::string id = json::require_string(item, "id"); |
| 89 | const std::string type = json::require_string(item, "type"); |
| 90 | const auto & item_config = item.require("config"); |
| 91 | if (id == "prompt") { |
| 92 | config.prompt_conditioner_type = type; |
| 93 | config.prompt_max_length = json::require_i64(item_config, "max_length"); |
| 94 | } else if (id == "seconds_total" && type == "number") { |
| 95 | config.seconds_min = json::require_f32(item_config, "min_val"); |
| 96 | config.seconds_max = json::require_f32(item_config, "max_val"); |
| 97 | } |
| 98 | } |
| 99 | if (config.prompt_conditioner_type.empty()) { |
| 100 | throw std::runtime_error("Stable Audio config is missing prompt conditioner"); |
| 101 | } |
| 102 | if (config.prompt_conditioner_type != "t5gemma" && config.prompt_conditioner_type != "t5") { |
| 103 | throw std::runtime_error("Stable Audio unsupported prompt conditioner type: " + config.prompt_conditioner_type); |
| 104 | } |
| 105 | |
| 106 | const auto & diffusion = model.require("diffusion"); |
| 107 | config.diffusion_type = json::require_string(diffusion, "type"); |
| 108 | config.diffusion_objective = json::optional_string(diffusion, "diffusion_objective", "v"); |
| 109 | if (const auto * shift = diffusion.find("sampling_distribution_shift_options"); shift != nullptr && !shift->is_null()) { |
| 110 | config.distribution_shift_type = json::optional_string(*shift, "type", "full"); |
| 111 | config.distribution_shift_base_shift = json::optional_f32(*shift, "base_shift", 0.5F); |
| 112 | config.distribution_shift_max_shift = json::optional_f32(*shift, "max_shift", 1.15F); |
| 113 | config.distribution_shift_min_length = json::optional_i64(*shift, "min_length", 256); |
| 114 | config.distribution_shift_max_length = json::optional_i64(*shift, "max_length", 4096); |
| 115 | config.distribution_shift_use_sine = json::optional_bool(*shift, "use_sine", false); |
| 116 | } |
| 117 | const auto & diffusion_config = diffusion.require("config"); |
| 118 | config.diffusion_io_channels = json::require_i64(diffusion_config, "io_channels"); |
| 119 | config.embed_dim = json::require_i64(diffusion_config, "embed_dim"); |
| 120 | config.depth = json::require_i64(diffusion_config, "depth"); |
no test coverage detected