| 466 | } |
| 467 | |
| 468 | AudioTensor 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, |