use core::num::NonZeroUsize;
fn check_nonzero(value: usize, msg: &str) -> usize {
NonZeroUsize::new(value).expect(msg);
value
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
pub struct ConvOptions<const N: usize> {
pub stride: [usize; N],
pub padding: [usize; N],
pub dilation: [usize; N],
pub groups: usize,
}
impl<const N: usize> ConvOptions<N> {
pub fn new(
stride: [usize; N],
padding: [usize; N],
dilation: [usize; N],
groups: usize,
) -> Self {
Self {
stride: stride.map(|s| check_nonzero(s, "stride must be non-zero")),
padding,
dilation: dilation.map(|d| check_nonzero(d, "dilation must be non-zero")),
groups: check_nonzero(groups, "groups must be non-zero"),
}
}
}
#[derive(Debug, Clone)]
pub struct PaddedConvOptions<const N: usize> {
pub options: ConvOptions<N>,
pub padding_end: Option<[usize; N]>,
}
impl<const N: usize> PaddedConvOptions<N> {
pub fn asymmetric(
stride: [usize; N],
padding_start: [usize; N],
padding_end: [usize; N],
dilation: [usize; N],
groups: usize,
) -> Self {
let options = ConvOptions::new(stride, padding_start, dilation, groups);
if padding_start == padding_end {
Self {
options,
padding_end: None,
}
} else {
Self {
options,
padding_end: Some(padding_end),
}
}
}
pub fn is_asymmetric(&self) -> bool {
self.padding_end.is_some()
}
}
impl<const N: usize> From<ConvOptions<N>> for PaddedConvOptions<N> {
fn from(options: ConvOptions<N>) -> Self {
Self {
options,
padding_end: None,
}
}
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
pub struct DeformConvOptions<const N: usize> {
pub stride: [usize; N],
pub padding: [usize; N],
pub dilation: [usize; N],
pub weight_groups: usize,
pub offset_groups: usize,
pub padding_end: Option<[usize; N]>,
}
impl<const N: usize> DeformConvOptions<N> {
pub fn new(
stride: [usize; N],
padding: [usize; N],
dilation: [usize; N],
weight_groups: usize,
offset_groups: usize,
) -> Self {
Self {
stride: stride.map(|s| check_nonzero(s, "stride must be non-zero")),
padding,
dilation: dilation.map(|d| check_nonzero(d, "dilation must be non-zero")),
weight_groups: check_nonzero(weight_groups, "weight groups must be non-zero"),
offset_groups: check_nonzero(offset_groups, "offset groups must be non-zero"),
padding_end: None,
}
}
pub fn with_padding_end(mut self, padding_end: [usize; N]) -> Self {
self.padding_end = (padding_end != self.padding).then_some(padding_end);
self
}
pub fn output_size(&self, input: [usize; N], kernel: [usize; N]) -> [usize; N] {
let padding_end = self.padding_end.unwrap_or(self.padding);
core::array::from_fn(|axis| {
let extent = kernel[axis].checked_sub(1)
.and_then(|size| size.checked_mul(self.dilation[axis]))
.and_then(|size| size.checked_add(1))
.expect("deform_conv kernel extent must be non-zero and fit usize");
let padded = input[axis].checked_add(self.padding[axis])
.and_then(|size| size.checked_add(padding_end[axis]))
.expect("deform_conv padded input must fit usize");
padded.checked_sub(extent)
.expect("deform_conv kernel exceeds padded input") / self.stride[axis] + 1
})
}
}
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
pub struct ConvTransposeOptions<const N: usize> {
pub stride: [usize; N],
pub padding: [usize; N],
pub padding_out: [usize; N],
pub dilation: [usize; N],
pub groups: usize,
}
impl<const N: usize> ConvTransposeOptions<N> {
pub fn new(
stride: [usize; N],
padding: [usize; N],
padding_out: [usize; N],
dilation: [usize; N],
groups: usize,
) -> Self {
Self {
stride: stride.map(|s| check_nonzero(s, "stride must be non-zero")),
padding,
padding_out,
dilation: dilation.map(|d| check_nonzero(d, "dilation must be non-zero")),
groups: check_nonzero(groups, "groups must be non-zero"),
}
}
}
#[derive(Debug, Clone)]
pub struct UnfoldOptions {
pub stride: [usize; 2],
pub padding: [usize; 2],
pub dilation: [usize; 2],
}
impl UnfoldOptions {
pub fn new(stride: [usize; 2], padding: [usize; 2], dilation: [usize; 2]) -> Self {
Self {
stride: stride.map(|s| check_nonzero(s, "stride must be non-zero")),
padding,
dilation: dilation.map(|d| check_nonzero(d, "dilation must be non-zero")),
}
}
}