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

Function compute_roformer_istft

src/models/roformer/runtime.cpp:785–890  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

783 return modulated;
784}
785
786engine::audio::AudioTensor compute_roformer_stft(
787 const std::vector<float> & waveform,
788 const std::vector<float> & window,
789 int64_t channels,
790 int64_t samples,
791 const engine::audio::STFTConfig & config,
792 size_t threads) {
793 if (static_cast<int64_t>(waveform.size()) != channels * samples ||
794 static_cast<int64_t>(window.size()) != config.win_length) {
795 throw std::runtime_error("RoFormer local STFT input size mismatch");
796 }
797 const int64_t pad = config.center ? config.n_fft / 2 : 0;
798 const int64_t frames = 1 + (samples + 2 * pad - config.n_fft) / config.hop_length;
799 const int64_t freq_bins = (config.n_fft / 2) + 1;
800 const int64_t window_offset = (config.n_fft - config.win_length) / 2;
801
802 std::vector<float> framed(static_cast<size_t>(channels * config.n_fft * frames), 0.0f);
803#ifdef _OPENMP
804 #pragma omp parallel for if(channels * frames >= 8)
805#endif
806 for (int64_t ch = 0; ch < channels; ++ch) {
807 const float * signal = waveform.data() + static_cast<size_t>(ch * samples);
808 for (int64_t frame_index = 0; frame_index < frames; ++frame_index) {
809 const int64_t start = frame_index * config.hop_length - pad;
810 for (int64_t i = 0; i < config.win_length; ++i) {
811 const int64_t sample_index = start + window_offset + i;
812 float sample = 0.0f;
813 if (sample_index >= 0 && sample_index < samples) {
814 sample = signal[sample_index];
815 } else if (config.pad_mode == engine::audio::STFTPadMode::Reflect) {
816 sample = signal[reflect_index(sample_index, samples)];
817 }
818 framed[static_cast<size_t>(((ch * config.n_fft) + (window_offset + i)) * frames + frame_index)] =
819 sample * window[static_cast<size_t>(i)];
820 }
821 }
822 }
823
824 engine::audio::TensorShape input_shape{
825 static_cast<size_t>(channels),
826 static_cast<size_t>(config.n_fft),
827 static_cast<size_t>(frames),
828 };
829 engine::audio::TensorStrideBytes input_strides{
830 static_cast<std::ptrdiff_t>(config.n_fft * frames * static_cast<int64_t>(sizeof(float))),
831 static_cast<std::ptrdiff_t>(frames * static_cast<int64_t>(sizeof(float))),
832 static_cast<std::ptrdiff_t>(sizeof(float)),
833 };
834 engine::audio::TensorStrideBytes output_strides{
835 static_cast<std::ptrdiff_t>(freq_bins * frames * static_cast<int64_t>(sizeof(std::complex<float>))),
836 static_cast<std::ptrdiff_t>(frames * static_cast<int64_t>(sizeof(std::complex<float>))),
837 static_cast<std::ptrdiff_t>(sizeof(std::complex<float>)),
838 };
839
840 std::vector<std::complex<float>> spectrum(
841 static_cast<size_t>(channels * freq_bins * frames),
842 std::complex<float>(0.0f, 0.0f));

Callers 1

separate_runtime_chunkFunction · 0.85

Calls 4

real_fft_inverseFunction · 0.85
assignMethod · 0.80
sizeMethod · 0.45
dataMethod · 0.45

Tested by

no test coverage detected