use crate::ElementConversion;
use core::num::NonZeroUsize;
pub(crate) 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, 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: padding.map(|padding| (padding, padding)),
dilation: dilation.map(|d| check_nonzero(d, "dilation must be non-zero")),
groups: check_nonzero(groups, "groups must be non-zero"),
}
}
pub fn new_with_padding(
stride: [usize; N],
padding: [(usize, 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"),
}
}
pub fn is_asymmetric(&self) -> bool {
self.padding.iter().any(|(begin, end)| begin != end)
}
pub fn padding_begin(&self) -> [usize; N] {
self.padding.map(|(begin, _)| begin)
}
pub fn padding_end(&self) -> [usize; N] {
self.padding.map(|(_, after)| after)
}
pub fn symmetric_padding(&self) -> [usize; N] {
assert!(
!self.is_asymmetric(),
"expected symmetric convolution padding"
);
self.padding_begin()
}
}
#[deprecated(since = "0.22.0", note = "Use `ConvOptions::new_with_padding` instead")]
#[derive(Debug, Clone)]
pub struct PaddedConvOptions<const N: usize> {
pub options: ConvOptions<N>,
pub padding_end: Option<[usize; N]>,
}
#[allow(deprecated)]
impl<const N: usize> PaddedConvOptions<N> {
pub fn asymmetric(
stride: [usize; N],
padding_begin: [usize; N],
padding_end: [usize; N],
dilation: [usize; N],
groups: usize,
) -> Self {
let options = ConvOptions::new(stride, padding_begin, dilation, groups);
let padding_end = (padding_begin != padding_end).then_some(padding_end);
Self {
options,
padding_end,
}
}
pub fn is_asymmetric(&self) -> bool {
self.padding_end.is_some()
}
}
#[allow(deprecated)]
impl<const N: usize> From<PaddedConvOptions<N>> for ConvOptions<N> {
fn from(value: PaddedConvOptions<N>) -> Self {
let Some(padding_end) = value.padding_end else {
return value.options;
};
let padding_begin = value.options.padding_begin();
let padding = core::array::from_fn(|i| (padding_begin[i], padding_end[i]));
ConvOptions::new_with_padding(
value.options.stride,
padding,
value.options.dilation,
value.options.groups,
)
}
}
#[allow(deprecated)]
impl<const N: usize> From<ConvOptions<N>> for PaddedConvOptions<N> {
fn from(options: ConvOptions<N>) -> Self {
if options.is_asymmetric() {
let padding_begin = options.padding_begin();
let padding_end = options.padding_end();
Self {
options: ConvOptions::new(
options.stride,
padding_begin,
options.dilation,
options.groups,
),
padding_end: Some(padding_end),
}
} else {
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,
}
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"),
}
}
}
#[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")),
}
}
}
#[derive(new, Debug, Clone, serde::Deserialize, serde::Serialize)]
pub enum InterpolateMode {
Nearest,
NearestExact,
Bilinear,
Bicubic,
Lanczos3,
}
#[derive(Debug, Clone)]
pub struct InterpolateOptions {
pub mode: InterpolateMode,
pub align_corners: bool,
}
impl InterpolateOptions {
pub fn new(mode: InterpolateMode) -> Self {
Self {
mode,
align_corners: true,
}
}
pub fn with_align_corners(mut self, align_corners: bool) -> Self {
self.align_corners = align_corners;
self
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Deserialize, serde::Serialize)]
pub enum GridSamplePaddingMode {
#[default]
Zeros,
Border,
Reflection,
}
#[derive(Debug, Clone)]
pub struct GridSampleOptions {
pub mode: InterpolateMode,
pub padding_mode: GridSamplePaddingMode,
pub align_corners: bool,
}
impl Default for GridSampleOptions {
fn default() -> Self {
Self {
mode: InterpolateMode::Bilinear,
padding_mode: GridSamplePaddingMode::Zeros,
align_corners: false,
}
}
}
impl From<InterpolateMode> for GridSampleOptions {
fn from(value: InterpolateMode) -> Self {
GridSampleOptions::new(value)
}
}
impl GridSampleOptions {
pub fn new(mode: InterpolateMode) -> Self {
Self {
mode,
..Default::default()
}
}
pub fn with_padding_mode(mut self, padding_mode: GridSamplePaddingMode) -> Self {
self.padding_mode = padding_mode;
self
}
pub fn with_align_corners(mut self, align_corners: bool) -> Self {
self.align_corners = align_corners;
self
}
}
#[derive(Debug, Clone, Copy, PartialEq, serde::Deserialize, serde::Serialize)]
pub enum PadMode {
Constant(f32),
Reflect,
Edge,
}
impl Default for PadMode {
fn default() -> Self {
PadMode::Constant(0.0)
}
}
impl<E: ElementConversion> From<E> for PadMode {
fn from(value: E) -> Self {
PadMode::Constant(value.elem())
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, serde::Deserialize, serde::Serialize)]
pub struct AttentionModuleOptions {
pub scale: Option<f64>,
pub softcap: Option<f64>,
pub is_causal: bool,
}
#[derive(Debug, Clone, Copy, Hash, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub enum IndexingUpdateOp {
Assign,
Add,
Mul,
Min,
Max,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn conv_options_symmetric_constructor() {
let options = ConvOptions::new([1, 2], [3, 4], [5, 6], 7);
assert_eq!(options.padding, [(3, 3), (4, 4)]);
assert_eq!(options.padding_begin(), [3, 4]);
assert_eq!(options.padding_end(), [3, 4]);
assert_eq!(options.symmetric_padding(), [3, 4]);
assert!(!options.is_asymmetric());
}
#[test]
fn conv_options_explicit_padding_constructor() {
let options = ConvOptions::new_with_padding([1, 2], [(3, 4), (5, 6)], [7, 8], 9);
assert_eq!(options.padding, [(3, 4), (5, 6)]);
assert_eq!(options.padding_begin(), [3, 5]);
assert_eq!(options.padding_end(), [4, 6]);
assert!(options.is_asymmetric());
}
#[test]
#[allow(deprecated)]
fn padded_conv_options_convert_to_conv_options() {
let options: ConvOptions<2> =
PaddedConvOptions::asymmetric([1, 2], [3, 5], [4, 6], [7, 8], 9).into();
assert_eq!(options.padding, [(3, 4), (5, 6)]);
assert_eq!(options.stride, [1, 2]);
assert_eq!(options.dilation, [7, 8]);
assert_eq!(options.groups, 9);
}
#[test]
#[allow(deprecated)]
fn asymmetric_conv_options_roundtrip_through_padded_options() {
let expected = ConvOptions::new_with_padding([1, 2], [(3, 4), (5, 6)], [7, 8], 9);
let padded: PaddedConvOptions<2> = expected.clone().into();
let actual: ConvOptions<2> = padded.into();
assert_eq!(actual, expected);
}
#[test]
#[should_panic = "expected symmetric convolution padding"]
fn conv_options_symmetric_padding_with_asymmetric_options() {
let options = ConvOptions::new_with_padding([1], [(1, 2)], [1], 1);
let _ = options.symmetric_padding();
}
#[test]
#[should_panic = "stride must be non-zero"]
fn conv_options_stride_zero() {
let _opt = ConvOptions::new([0, 1], [0, 0], [1, 1], 1);
}
#[test]
#[should_panic = "dilation must be non-zero"]
fn conv_options_dilation_zero() {
let _opt = ConvOptions::new([1, 1], [0, 0], [0, 0], 1);
}
#[test]
#[should_panic = "groups must be non-zero"]
fn conv_options_groups_zero() {
let _opt = ConvOptions::new([1, 1], [0, 0], [1, 1], 0);
}
#[test]
#[should_panic = "stride must be non-zero"]
fn conv_transpose_options_stride_zero() {
let _opt = ConvTransposeOptions::new([0, 1], [0, 0], [0, 0], [1, 1], 1);
}
#[test]
#[should_panic = "dilation must be non-zero"]
fn conv_transpose_options_dilation_zero() {
let _opt = ConvTransposeOptions::new([1, 1], [0, 0], [0, 0], [0, 0], 1);
}
#[test]
#[should_panic = "groups must be non-zero"]
fn conv_transpose_options_groups_zero() {
let _opt = ConvTransposeOptions::new([1, 1], [0, 0], [0, 0], [1, 1], 0);
}
#[test]
#[should_panic = "stride must be non-zero"]
fn deform_conv_options_stride_zero() {
let _opt = DeformConvOptions::new([0, 1], [0, 0], [1, 1], 1, 1);
}
#[test]
#[should_panic = "dilation must be non-zero"]
fn deform_conv_options_dilation_zero() {
let _opt = DeformConvOptions::new([1, 1], [0, 0], [0, 0], 1, 1);
}
#[test]
#[should_panic = "weight groups must be non-zero"]
fn deform_conv_options_weights_groups_zero() {
let _opt = DeformConvOptions::new([1, 1], [0, 0], [1, 1], 0, 1);
}
#[test]
#[should_panic = "offset groups must be non-zero"]
fn deform_conv_options_offset_groups_zero() {
let _opt = DeformConvOptions::new([1, 1], [0, 0], [1, 1], 1, 0);
}
#[test]
#[should_panic = "stride must be non-zero"]
fn unfold_options_stride_zero() {
let _opt = UnfoldOptions::new([0, 1], [0, 0], [1, 1]);
}
#[test]
#[should_panic = "dilation must be non-zero"]
fn unfold_options_dilation_zero() {
let _opt = UnfoldOptions::new([1, 1], [0, 0], [0, 0]);
}
}