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

Method compute

src/framework/audio/dsp.cpp:468–564  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

466}
467
468AudioTensor ISTFT::compute(
469 const std::vector<float> & complex_spec,
470 const std::vector<float> & window,
471 int64_t batch,
472 int64_t freq_bins,
473 int64_t frames,
474 int64_t samples,
475 const STFTConfig & config,
476 size_t threads) const {
477 if (static_cast<int64_t>(complex_spec.size()) != checked_product({batch, freq_bins, frames, 2}) ||
478 static_cast<int64_t>(window.size()) != config.win_length) {
479 throw std::runtime_error("ISTFT input size mismatch");
480 }
481
482 const int64_t pad = config.n_fft / 2;
483 const int64_t padded_samples = samples + 2 * pad;
484 const int64_t usable_window = std::min<int64_t>(config.win_length, config.n_fft);
485
486 std::vector<std::complex<float>> spectrum(
487 static_cast<size_t>(checked_product({batch, frames, freq_bins})),
488 std::complex<float>(0.0f, 0.0f));
489#ifdef _OPENMP
490 #pragma omp parallel for collapse(3) if(batch * frames * freq_bins >= 4096)
491#endif
492 for (int64_t b = 0; b < batch; ++b) {
493 for (int64_t frame_index = 0; frame_index < frames; ++frame_index) {
494 for (int64_t k = 0; k < freq_bins; ++k) {
495 const size_t src = static_cast<size_t>((((b * freq_bins) + k) * frames + frame_index) * 2);
496 spectrum[static_cast<size_t>((b * frames + frame_index) * freq_bins + k)] = {
497 complex_spec[src],
498 complex_spec[src + 1],
499 };
500 }
501 }
502 }
503
504 TensorShape output_shape{
505 static_cast<size_t>(batch),
506 static_cast<size_t>(frames),
507 static_cast<size_t>(config.n_fft),
508 };
509 TensorStrideBytes input_strides{
510 static_cast<std::ptrdiff_t>(frames * freq_bins * static_cast<int64_t>(sizeof(std::complex<float>))),
511 static_cast<std::ptrdiff_t>(freq_bins * static_cast<int64_t>(sizeof(std::complex<float>))),
512 static_cast<std::ptrdiff_t>(sizeof(std::complex<float>)),
513 };
514 TensorStrideBytes output_strides{
515 static_cast<std::ptrdiff_t>(frames * config.n_fft * static_cast<int64_t>(sizeof(float))),
516 static_cast<std::ptrdiff_t>(config.n_fft * static_cast<int64_t>(sizeof(float))),
517 static_cast<std::ptrdiff_t>(sizeof(float)),
518 };
519
520 std::vector<float> framed(static_cast<size_t>(checked_product({batch, frames, config.n_fft})), 0.0f);
521 real_fft_inverse(
522 output_shape,
523 input_strides,
524 output_strides,
525 2,

Calls 15

real_fft_inverseFunction · 0.85
STFTClass · 0.85
timing_log_enabledFunction · 0.85
MelFilterbankClass · 0.85
workerClass · 0.85
normalize_batch_implFunction · 0.85
assignMethod · 0.80
compute_magnitudeMethod · 0.80
compute_customMethod · 0.80
checked_productFunction · 0.70
measure_msFunction · 0.70
timing_log_scalarFunction · 0.50