use crate::SignalError;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum LinearOutput {
Full,
Same,
Valid,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ConvolutionMode {
Linear(LinearOutput),
Circular {
period: usize,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum BoundaryPolicy {
ZeroPad,
Periodic,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ConvolutionNormalization {
None,
KernelSum,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ConvolutionAlgorithm {
Auto,
Direct,
Fft,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ConvolutionCostPlan {
pub requested: ConvolutionAlgorithm,
pub selected: ConvolutionAlgorithm,
pub direct_cost_units: usize,
pub fft_cost_units: usize,
pub fft_len: usize,
pub fft_scratch_bytes: usize,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ConvolutionPlan {
pub mode: ConvolutionMode,
pub algorithm: ConvolutionAlgorithm,
pub boundary: BoundaryPolicy,
pub normalization: ConvolutionNormalization,
}
impl ConvolutionPlan {
pub const fn linear_full() -> Self {
Self {
mode: ConvolutionMode::Linear(LinearOutput::Full),
algorithm: ConvolutionAlgorithm::Auto,
boundary: BoundaryPolicy::ZeroPad,
normalization: ConvolutionNormalization::None,
}
}
pub const fn circular(period: usize) -> Self {
Self {
mode: ConvolutionMode::Circular { period },
algorithm: ConvolutionAlgorithm::Auto,
boundary: BoundaryPolicy::Periodic,
normalization: ConvolutionNormalization::None,
}
}
pub fn inspect(
&self,
signal_len: usize,
kernel_len: usize,
) -> Result<ConvolutionCostPlan, SignalError> {
self.validate(signal_len, kernel_len)?;
let full_len = linear_full_len(signal_len, kernel_len)?;
let fft_len = match self.mode {
ConvolutionMode::Linear(_) => {
full_len
.checked_next_power_of_two()
.ok_or(SignalError::InvalidLength {
len: full_len,
reason: "convolution FFT length overflowed",
})?
}
ConvolutionMode::Circular { period } => period,
};
let direct_cost_units = match self.mode {
ConvolutionMode::Linear(_) => signal_len.checked_mul(kernel_len),
ConvolutionMode::Circular { period } => period.checked_mul(period),
}
.ok_or(SignalError::InvalidLength {
len: signal_len,
reason: "direct convolution cost overflowed",
})?;
let stages = usize::try_from(usize::BITS - fft_len.leading_zeros()).unwrap_or(usize::MAX);
let fft_cost_units = fft_len
.checked_mul(stages)
.and_then(|cost| cost.checked_mul(3))
.and_then(|cost| cost.checked_add(fft_len))
.ok_or(SignalError::InvalidLength {
len: fft_len,
reason: "FFT convolution cost overflowed",
})?;
let fft_scratch_bytes = fft_len
.checked_mul(3)
.and_then(|cells| cells.checked_mul(2 * size_of::<f64>()))
.ok_or(SignalError::InvalidLength {
len: fft_len,
reason: "FFT convolution scratch size overflowed",
})?;
let selected = match self.algorithm {
ConvolutionAlgorithm::Auto if direct_cost_units <= fft_cost_units => {
ConvolutionAlgorithm::Direct
}
ConvolutionAlgorithm::Auto => ConvolutionAlgorithm::Fft,
selected => selected,
};
Ok(ConvolutionCostPlan {
requested: self.algorithm,
selected,
direct_cost_units,
fft_cost_units,
fft_len,
fft_scratch_bytes,
})
}
fn validate(&self, signal_len: usize, kernel_len: usize) -> Result<(), SignalError> {
if signal_len == 0 || kernel_len == 0 {
return Err(SignalError::InvalidLength {
len: signal_len.min(kernel_len),
reason: "convolution inputs must both be non-empty",
});
}
match (self.mode, self.boundary) {
(ConvolutionMode::Linear(_), BoundaryPolicy::ZeroPad)
| (ConvolutionMode::Circular { .. }, BoundaryPolicy::Periodic) => {}
(ConvolutionMode::Linear(_), BoundaryPolicy::Periodic) => {
return Err(SignalError::InvalidPolicy {
policy: "boundary",
reason: "linear convolution requires zero padding",
});
}
(ConvolutionMode::Circular { .. }, BoundaryPolicy::ZeroPad) => {
return Err(SignalError::InvalidPolicy {
policy: "boundary",
reason: "circular convolution requires periodic extension",
});
}
}
if let ConvolutionMode::Circular { period: 0 } = self.mode {
return Err(SignalError::InvalidLength {
len: 0,
reason: "circular convolution period must be nonzero",
});
}
if self.mode == ConvolutionMode::Linear(LinearOutput::Valid) && signal_len < kernel_len {
return Err(SignalError::InvalidLength {
len: signal_len,
reason: "valid convolution requires signal length at least kernel length",
});
}
Ok(())
}
}
pub(crate) fn linear_full_len(signal_len: usize, kernel_len: usize) -> Result<usize, SignalError> {
signal_len
.checked_add(kernel_len)
.and_then(|len| len.checked_sub(1))
.ok_or(SignalError::InvalidLength {
len: signal_len,
reason: "linear convolution output length overflowed",
})
}
pub(crate) fn retained_span(
mode: ConvolutionMode,
signal_len: usize,
kernel_len: usize,
) -> Result<(usize, usize), SignalError> {
Ok(match mode {
ConvolutionMode::Linear(LinearOutput::Full) => {
(0, linear_full_len(signal_len, kernel_len)?)
}
ConvolutionMode::Linear(LinearOutput::Same) => ((kernel_len - 1) / 2, signal_len),
ConvolutionMode::Linear(LinearOutput::Valid) => {
(kernel_len - 1, signal_len - kernel_len + 1)
}
ConvolutionMode::Circular { period } => (0, period),
})
}