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

Method compute_complex

src/framework/audio/dsp.cpp:381–466  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

379}
380
381AudioTensor 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));

Callers 8

source_stft_bctFunction · 0.80
denoise_mono_16k_wholeFunction · 0.80
extract_timbre_melFunction · 0.80
mainFunction · 0.80
mainFunction · 0.80

Calls 6

real_fft_forwardFunction · 0.85
assignMethod · 0.80
checked_productFunction · 0.70
reflect_indexFunction · 0.70
sizeMethod · 0.45
dataMethod · 0.45