use ruda_kernel::dsl as kernel_dsl;
use core::f32::consts::PI;
use ruda_kernel::dsl::prelude::*;
use ruda_kernel::library::tensor::AsView as _;
use ruda_kernel::library::tensor::AsViewExpand;
use ruda_kernel::library::tensor::AsViewMut as _;
use ruda_kernel::library::tensor::AsViewMutExpand;
use ruda_kernel::library::tensor::TensorHandle;
use crate::{
fft::{
FftMode,
cfft::{CfftBindings, MAX_SHARED_N_FFT, cfft_launch_any_size},
fft_parallel::{bit_reverse, fft_butterfly_parallel},
},
layout::BatchSignalLayout,
};
pub(crate) fn rfft_large_launch<R: Runtime>(
client: &ComputeClient<R>,
signal: TensorBinding<R>,
spectrum_re: TensorBinding<R>,
spectrum_im: TensorBinding<R>,
dim: usize,
signal_len: usize,
dtype: StorageType,
) -> Result<(), LaunchError> {
let n_fft = (spectrum_re.shape[dim] - 1) * 2;
let m = n_fft / 2;
let count: usize = signal
.shape
.iter()
.enumerate()
.filter(|(i, _)| *i != dim)
.map(|(_, e)| *e)
.product();
if m <= MAX_SHARED_N_FFT {
let threads = (m / 2).clamp(1, 256);
let grid =
ruda_kernel::dsl::calculate_ruda_count_elemwise(client, count, RudaDim::new_single());
rfft_fused_kernel::launch::<f32, R>(
client,
grid,
RudaDim::new_1d(threads as u32),
signal.into_tensor_arg(),
spectrum_re.into_tensor_arg(),
spectrum_im.into_tensor_arg(),
count as u32,
signal_len as u32,
n_fft,
m,
m.trailing_zeros() as usize,
threads,
dim,
);
return Ok(());
}
let packed_shape: Vec<usize> = signal
.shape
.iter()
.enumerate()
.map(|(i, &s)| if i == dim { m } else { s })
.collect();
let packed_elems: usize = packed_shape.iter().product();
let packed_re = TensorHandle::<R>::new_contiguous(
packed_shape.clone(),
client.empty(packed_elems * dtype.size()),
dtype,
);
let packed_im = TensorHandle::<R>::new_contiguous(
packed_shape.clone(),
client.empty(packed_elems * dtype.size()),
dtype,
);
{
let ruda_dim = RudaDim::new_1d(256);
let ruda_count = ruda_kernel::dsl::calculate_ruda_count_elemwise(client, count * m, ruda_dim);
rfft_pack_kernel::launch::<f32, R>(
client,
ruda_count,
ruda_dim,
signal.into_tensor_arg(),
packed_re.clone().binding().into_tensor_arg(),
packed_im.clone().binding().into_tensor_arg(),
(count * m) as u32,
signal_len as u32,
m,
dim,
);
}
cfft_launch_any_size::<R>(
client,
CfftBindings {
input_re: packed_re.clone().binding(),
input_im: packed_im.clone().binding(),
output_re: packed_re.clone().binding(),
output_im: packed_im.clone().binding(),
},
dim,
dtype,
FftMode::Forward,
)?;
{
let n_freq = m + 1;
let ruda_dim = RudaDim::new_1d(256);
let ruda_count = ruda_kernel::dsl::calculate_ruda_count_elemwise(client, count * n_freq, ruda_dim);
rfft_post_kernel::launch::<f32, R>(
client,
ruda_count,
ruda_dim,
packed_re.binding().into_tensor_arg(),
packed_im.binding().into_tensor_arg(),
spectrum_re.into_tensor_arg(),
spectrum_im.into_tensor_arg(),
(count * n_freq) as u32,
n_fft,
m,
dim,
);
}
Ok(())
}
pub(crate) fn irfft_large_launch<R: Runtime>(
client: &ComputeClient<R>,
spectrum_re: TensorBinding<R>,
spectrum_im: TensorBinding<R>,
signal: TensorBinding<R>,
dim: usize,
spec_bins: usize,
dtype: StorageType,
) -> Result<(), LaunchError> {
let n_fft = signal.shape[dim];
let m = n_fft / 2;
let count: usize = signal
.shape
.iter()
.enumerate()
.filter(|(i, _)| *i != dim)
.map(|(_, e)| *e)
.product();
if m <= MAX_SHARED_N_FFT {
let threads = (m / 2).clamp(1, 256);
let grid =
ruda_kernel::dsl::calculate_ruda_count_elemwise(client, count, RudaDim::new_single());
irfft_fused_kernel::launch::<f32, R>(
client,
grid,
RudaDim::new_1d(threads as u32),
spectrum_re.into_tensor_arg(),
spectrum_im.into_tensor_arg(),
signal.into_tensor_arg(),
count as u32,
spec_bins as u32,
n_fft,
m,
m.trailing_zeros() as usize,
threads,
dim,
);
return Ok(());
}
let packed_shape: Vec<usize> = signal
.shape
.iter()
.enumerate()
.map(|(i, &s)| if i == dim { m } else { s })
.collect();
let packed_elems: usize = packed_shape.iter().product();
let packed_in_re = TensorHandle::<R>::new_contiguous(
packed_shape.clone(),
client.empty(packed_elems * dtype.size()),
dtype,
);
let packed_in_im = TensorHandle::<R>::new_contiguous(
packed_shape.clone(),
client.empty(packed_elems * dtype.size()),
dtype,
);
let packed_out_re = TensorHandle::<R>::new_contiguous(
packed_shape.clone(),
client.empty(packed_elems * dtype.size()),
dtype,
);
let packed_out_im = TensorHandle::<R>::new_contiguous(
packed_shape.clone(),
client.empty(packed_elems * dtype.size()),
dtype,
);
{
let ruda_dim = RudaDim::new_1d(256);
let ruda_count = ruda_kernel::dsl::calculate_ruda_count_elemwise(client, count * m, ruda_dim);
irfft_pre_kernel::launch::<f32, R>(
client,
ruda_count,
ruda_dim,
spectrum_re.into_tensor_arg(),
spectrum_im.into_tensor_arg(),
packed_in_re.clone().binding().into_tensor_arg(),
packed_in_im.clone().binding().into_tensor_arg(),
(count * m) as u32,
spec_bins as u32,
n_fft,
m,
dim,
);
}
cfft_launch_any_size::<R>(
client,
CfftBindings {
input_re: packed_in_re.binding(),
input_im: packed_in_im.binding(),
output_re: packed_out_re.clone().binding(),
output_im: packed_out_im.clone().binding(),
},
dim,
dtype,
FftMode::Inverse,
)?;
{
let ruda_dim = RudaDim::new_1d(256);
let ruda_count = ruda_kernel::dsl::calculate_ruda_count_elemwise(client, count * m, ruda_dim);
irfft_unpack_kernel::launch::<f32, R>(
client,
ruda_count,
ruda_dim,
packed_out_re.binding().into_tensor_arg(),
packed_out_im.binding().into_tensor_arg(),
signal.into_tensor_arg(),
(count * m) as u32,
m,
dim,
);
}
Ok(())
}
mod kernels;
use kernels::*;