Skip to main content

rutensor/
operation.rs

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    /// Conjugation is the identity for the real storage types accepted by ruTENSOR.
8    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/// Arithmetic precision, including operand transforms and scalar coefficients.
18#[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/// A buffer-independent operation. Input order is the order passed to its constructor.
40#[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    /// D = alpha * sum(op(A) * op(B)) + beta * op(C).
53    /// C, if supplied, follows A and B in the execution input list.
54    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    /// General multi-operand sum of products. Modes absent from D are summed.
62    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    /// D = alpha * reduce(op(A)) + beta * op(C).
76    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    /// D = alpha * op(A), with a physical copy into the specified output layout.
87    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    /// D = combine(alpha * op(A), beta * op(B)).
94    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    /// D = combine_abc(combine_ab(alpha * op(A), beta * op(B)), gamma * op(C)).
103    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}