| 194 | std::string(family) + " native runtime does not yet support skip_connection"); |
| 195 | } |
| 196 | config.stft_freq_bins = config.n_fft / 2 + 1; |
| 197 | config.chunk_frames = 1 + config.chunk_size / config.hop_length; |
| 198 | config.instruments = {"vocals"}; |
| 199 | config.target_instrument = std::string("vocals"); |
| 200 | if (family == kBsRoformerFamily) { |
| 201 | config.fused_qkv = true; |
| 202 | config.transformer_output_norm = false; |
| 203 | config.has_final_norm = true; |
| 204 | fill_bs_band_layout(config, json::require_i64_array(parsed, "freqs_per_bands")); |
| 205 | } else { |
| 206 | config.num_bands = json::require_i32(parsed, "num_bands"); |
| 207 | config.fused_qkv = false; |
| 208 | config.transformer_output_norm = true; |
| 209 | config.has_final_norm = json::optional_bool(parsed, "has_final_norm", false); |
| 210 | fill_mel_band_layout(config, config.num_bands); |
| 211 | } |
| 212 | return config; |
| 213 | } |
| 214 | |
| 215 | } // namespace |
| 216 | |
| 217 | std::shared_ptr<const RoformerAssets> load_mel_band_roformer_assets( |
| 218 | const runtime::ModelLoadRequest & request) { |
| 219 | return load_roformer_assets(request, kMelBandRoformerFamily); |
| 220 | } |
| 221 | |
| 222 | std::shared_ptr<const RoformerAssets> load_bs_roformer_assets( |
| 223 | const runtime::ModelLoadRequest & request) { |
| 224 | return load_roformer_assets(request, kBsRoformerFamily); |
| 225 | } |
| 226 | |
| 227 | std::shared_ptr<const RoformerAssets> load_roformer_assets( |
| 228 | const runtime::ModelLoadRequest & request, |
| 229 | std::string_view family) { |
| 230 | auto assets = std::make_shared<RoformerAssets>(); |
| 231 | assets->resources = load_resources(request, family); |
| 232 | assets->tensor_source = assets->resources.open_tensor_source("weights"); |
| 233 | const auto parsed = assets->resources.parse_json("config"); |
| 234 | assets->config = parse_config(parsed, family); |
no test coverage detected