| 268 | } |
| 269 | |
| 270 | runtime::TaskResult StableAudioSession::run(const runtime::TaskRequest & request) { |
| 271 | require_prepared("Stable Audio run"); |
| 272 | if (!conditioner_inputs_) { |
| 273 | throw std::runtime_error("Stable Audio conditioner was not prepared"); |
| 274 | } |
| 275 | if (!conditioner_runtime_) { |
| 276 | throw std::runtime_error("Stable Audio conditioner runtime was not prepared"); |
| 277 | } |
| 278 | if (!rf_dit_) { |
| 279 | throw std::runtime_error("Stable Audio RF DiT runtime was not prepared"); |
| 280 | } |
| 281 | if (!same_) { |
| 282 | throw std::runtime_error("Stable Audio SAME runtime was not prepared"); |
| 283 | } |
| 284 | const auto wall_start = Clock::now(); |
| 285 | const StableAudioRequest parsed = clamp_request_to_model_limit( |
| 286 | parse_stable_audio_request(request), |
| 287 | assets_->config); |
| 288 | const StableAudioConditioningBatch conditioning = conditioner_inputs_->build(parsed); |
| 289 | const auto rng_policy = engine::sampling::resolve_torch_cuda_sampling_policy( |
| 290 | execution_context().backend_type(), |
| 291 | execution_context().config().device, |
| 292 | "stable_audio.rng", |
| 293 | "Stable Audio"); |
| 294 | const auto sampling_start = Clock::now(); |
| 295 | StableAudioSamplingState sampling = prepare_stable_audio_sampling_state(assets_->config, parsed, rng_policy); |
| 296 | uint64_t rng_offset_blocks = engine::sampling::torch_cuda_tensor_iterator_offset_blocks( |
| 297 | static_cast<uint64_t>(sampling.noise.size()), |
| 298 | rng_policy); |
| 299 | engine::debug::timing_log_scalar( |
| 300 | "stable_audio.sampling_prepare_ms", |
| 301 | engine::debug::elapsed_ms(sampling_start, Clock::now())); |
| 302 | if (parsed.init_audio.has_value()) { |
| 303 | const auto init_latents = same_->encode( |
| 304 | *parsed.init_audio, |
| 305 | sampling.audio_sample_size, |
| 306 | parsed.seed, |
| 307 | rng_offset_blocks); |
| 308 | apply_init_audio(sampling, init_latents, parsed.init_noise_level, assets_->config); |
| 309 | } |
| 310 | if (parsed.inpaint_audio.has_value() || !parsed.inpaint_regions.empty()) { |
| 311 | const auto mask = make_inpaint_mask(parsed, sampling, assets_->config); |
| 312 | std::vector<float> inpaint_latents(static_cast<size_t>(assets_->config.latent_dim * sampling.latent_sample_size), 0.0F); |
| 313 | if (parsed.inpaint_audio.has_value()) { |
| 314 | inpaint_latents = same_->encode( |
| 315 | *parsed.inpaint_audio, |
| 316 | sampling.audio_sample_size, |
| 317 | parsed.seed, |
| 318 | rng_offset_blocks); |
| 319 | } |
| 320 | apply_inpaint_conditioning(sampling, inpaint_latents, mask, assets_->config); |
| 321 | } |
| 322 | const auto conditioner_start = Clock::now(); |
| 323 | const StableAudioConditioningInputs conditioning_inputs = conditioner_runtime_->encode(conditioning); |
| 324 | engine::debug::timing_log_scalar( |
| 325 | "stable_audio.conditioner_ms", |
| 326 | engine::debug::elapsed_ms(conditioner_start, Clock::now())); |
| 327 | if (mem_saver_) { |
nothing calls this directly
no test coverage detected