1use crate::{OperandDescriptor, TensorDescriptor, Mode, Error, Result};
2use ruda_core::tensor::DType;
3
4#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
5pub enum UnaryOp {
6 Identity, Negate, Abs, Sqrt, Exp, Log, Sin, Cos, Tanh, Relu, Reciprocal,
7 Conjugate,
9}
10
11#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
12pub enum BinaryOp { Add, Mul, Min, Max }
13
14#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
15pub enum ReductionOp { Sum, Product, Min, Max }
16
17#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
19pub enum ComputeType { F32, F64 }
20
21impl ComputeType {
22 pub fn dtype(self) -> DType {
23 match self { Self::F32 => DType::F32, Self::F64 => DType::F64 }
24 }
25
26 pub(crate) fn for_operands(inputs: &[OperandDescriptor]) -> Self {
27 if inputs.iter().any(|x| x.tensor.dtype == DType::F64) { Self::F64 } else { Self::F32 }
28 }
29}
30
31#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
32pub(crate) enum Kind {
33 Contraction,
34 Reduction(ReductionOp),
35 Permutation,
36 Elementwise(BinaryOp, BinaryOp),
37}
38
39#[derive(Clone, Debug, PartialEq, Eq)]
41pub struct OperationDescriptor {
42 pub(crate) inputs: Vec<OperandDescriptor>,
43 pub(crate) output: TensorDescriptor,
44 pub(crate) output_modes: Vec<Mode>,
45 pub(crate) compute: ComputeType,
46 pub(crate) kind: Kind,
47 pub(crate) terms: usize,
48 pub(crate) addend: bool,
49}
50
51impl OperationDescriptor {
52 pub fn contraction(
55 a: OperandDescriptor, b: OperandDescriptor, c: Option<OperandDescriptor>,
56 output: TensorDescriptor, modes: &[Mode], compute: ComputeType,
57 ) -> Result<Self> {
58 Self::sum_product(vec![a, b], c, output, modes, compute)
59 }
60
61 pub fn sum_product(
63 mut inputs: Vec<OperandDescriptor>, c: Option<OperandDescriptor>,
64 output: TensorDescriptor, modes: &[Mode], compute: ComputeType,
65 ) -> Result<Self> {
66 if inputs.is_empty() {
67 return Err(Error::InvalidOperation("a sum of products requires at least one input".into()));
68 }
69 let terms = inputs.len();
70 let addend = c.is_some();
71 inputs.extend(c);
72 Self::new(inputs, output, modes, compute, Kind::Contraction, terms, addend)
73 }
74
75 pub fn reduction(
77 a: OperandDescriptor, c: Option<OperandDescriptor>, output: TensorDescriptor,
78 modes: &[Mode], reduction: ReductionOp, compute: ComputeType,
79 ) -> Result<Self> {
80 let addend = c.is_some();
81 let mut inputs = vec![a];
82 inputs.extend(c);
83 Self::new(inputs, output, modes, compute, Kind::Reduction(reduction), 1, addend)
84 }
85
86 pub fn permutation(
88 a: OperandDescriptor, output: TensorDescriptor, modes: &[Mode], compute: ComputeType,
89 ) -> Result<Self> {
90 Self::new(vec![a], output, modes, compute, Kind::Permutation, 1, false)
91 }
92
93 pub fn elementwise_binary(
95 a: OperandDescriptor, b: OperandDescriptor, output: TensorDescriptor,
96 modes: &[Mode], combine: BinaryOp, compute: ComputeType,
97 ) -> Result<Self> {
98 Self::new(vec![a, b], output, modes, compute,
99 Kind::Elementwise(combine, BinaryOp::Add), 2, false)
100 }
101
102 pub fn elementwise_trinary(
104 inputs: [OperandDescriptor; 3], output: TensorDescriptor, modes: &[Mode],
105 combine_ab: BinaryOp, combine_abc: BinaryOp, compute: ComputeType,
106 ) -> Result<Self> {
107 Self::new(inputs.into(), output, modes, compute,
108 Kind::Elementwise(combine_ab, combine_abc), 3, false)
109 }
110
111 fn new(
112 inputs: Vec<OperandDescriptor>, output: TensorDescriptor, modes: &[Mode],
113 compute: ComputeType, kind: Kind, terms: usize, addend: bool,
114 ) -> Result<Self> {
115 if output.rank() != modes.len() {
116 return Err(Error::InvalidDescriptor("output modes and axes differ".into()));
117 }
118 for (i, mode) in modes.iter().enumerate() {
119 if modes[..i].contains(mode) {
120 return Err(Error::InvalidDescriptor("output modes must be unique".into()));
121 }
122 }
123 if !output.is_nonoverlapping() {
124 return Err(Error::InvalidDescriptor("output strides overlap".into()));
125 }
126 if compute == ComputeType::F32 && inputs.iter().any(|x| x.tensor.dtype == DType::F64) {
127 return Err(Error::InvalidOperation("F64 inputs require F64 computation".into()));
128 }
129 Ok(Self { inputs, output, output_modes: modes.to_vec(), compute, kind, terms, addend })
130 }
131
132 pub fn inputs(&self) -> &[OperandDescriptor] { &self.inputs }
133 pub fn output(&self) -> &TensorDescriptor { &self.output }
134 pub fn output_modes(&self) -> &[Mode] { &self.output_modes }
135 pub fn compute_type(&self) -> ComputeType { self.compute }
136 pub fn scalar_count(&self) -> usize {
137 match self.kind {
138 Kind::Elementwise(_, _) => self.terms,
139 _ => 1 + usize::from(self.addend),
140 }
141 }
142}