| 379 | } |
| 380 | |
| 381 | AudioTensor STFT::compute_complex( |
| 382 | const std::vector<float> & waveform, |
| 383 | const std::vector<float> & window, |
| 384 | int64_t batch, |
| 385 | int64_t samples, |
| 386 | const STFTConfig & config, |
| 387 | size_t threads) const { |
| 388 | if (static_cast<int64_t>(waveform.size()) != checked_product({batch, samples}) || |
| 389 | static_cast<int64_t>(window.size()) != config.win_length) { |
| 390 | throw std::runtime_error("STFT input size mismatch"); |
| 391 | } |
| 392 | |
| 393 | const int64_t pad = config.center ? config.n_fft / 2 : 0; |
| 394 | const int64_t frames = 1 + (samples + 2 * pad - config.n_fft) / config.hop_length; |
| 395 | const int64_t freq_bins = (config.n_fft / 2) + 1; |
| 396 | const int64_t window_offset = (config.n_fft - config.win_length) / 2; |
| 397 | |
| 398 | std::vector<float> framed(static_cast<size_t>(checked_product({batch, frames, config.n_fft})), 0.0f); |
| 399 | #ifdef _OPENMP |
| 400 | #pragma omp parallel for collapse(2) if(batch * frames >= 8) |
| 401 | #endif |
| 402 | for (int64_t b = 0; b < batch; ++b) { |
| 403 | for (int64_t frame_index = 0; frame_index < frames; ++frame_index) { |
| 404 | const float * signal = waveform.data() + static_cast<size_t>(b * samples); |
| 405 | const int64_t start = frame_index * config.hop_length - pad; |
| 406 | float * frame = framed.data() + static_cast<size_t>((b * frames + frame_index) * config.n_fft); |
| 407 | for (int64_t i = 0; i < config.win_length; ++i) { |
| 408 | const int64_t sample_index = start + window_offset + i; |
| 409 | float sample = 0.0f; |
| 410 | if (sample_index >= 0 && sample_index < samples) { |
| 411 | sample = signal[sample_index]; |
| 412 | } else if (config.pad_mode == STFTPadMode::Reflect) { |
| 413 | sample = signal[reflect_index(sample_index, samples)]; |
| 414 | } |
| 415 | frame[window_offset + i] = sample * window[static_cast<size_t>(i)]; |
| 416 | } |
| 417 | } |
| 418 | } |
| 419 | |
| 420 | TensorShape shape_in{ |
| 421 | static_cast<size_t>(batch), |
| 422 | static_cast<size_t>(frames), |
| 423 | static_cast<size_t>(config.n_fft), |
| 424 | }; |
| 425 | TensorStrideBytes stride_in{ |
| 426 | static_cast<std::ptrdiff_t>(frames * config.n_fft * static_cast<int64_t>(sizeof(float))), |
| 427 | static_cast<std::ptrdiff_t>(config.n_fft * static_cast<int64_t>(sizeof(float))), |
| 428 | static_cast<std::ptrdiff_t>(sizeof(float)), |
| 429 | }; |
| 430 | TensorStrideBytes stride_out{ |
| 431 | static_cast<std::ptrdiff_t>(frames * freq_bins * static_cast<int64_t>(sizeof(std::complex<float>))), |
| 432 | static_cast<std::ptrdiff_t>(freq_bins * static_cast<int64_t>(sizeof(std::complex<float>))), |
| 433 | static_cast<std::ptrdiff_t>(sizeof(std::complex<float>)), |
| 434 | }; |
| 435 | |
| 436 | std::vector<std::complex<float>> spectrum( |
| 437 | static_cast<size_t>(checked_product({batch, frames, freq_bins})), |
| 438 | std::complex<float>(0.0f, 0.0f)); |