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; 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,
}
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]
#[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]);
}
}