| 783 | return modulated; |
| 784 | } |
| 785 | |
| 786 | engine::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)); |
no test coverage detected