Skip to main content

rutensor/
tensor.rs

1use crate::{BinaryOp, ComputeType, Error, Mode, OperandDescriptor, OperationDescriptor, Plan, ReductionOp, Result};
2use crate::einsum::{infer_output, output_dtype};
3use ruda_kernel::{dsl::Runtime, tensor::RudaTensor};
4
5/// Sum the product of A and B over modes absent from the output.
6pub fn contract<R: Runtime>(
7    a: &RudaTensor<R>, modes_a: &[Mode], b: &RudaTensor<R>, modes_b: &[Mode], output_modes: &[Mode],
8) -> Result<RudaTensor<R>> {
9    let operands = vec![OperandDescriptor::from_tensor(a, modes_a)?, OperandDescriptor::from_tensor(b, modes_b)?];
10    let output = infer_output(&operands, output_modes, output_dtype(&operands))?;
11    let compute = ComputeType::for_operands(&operands);
12    Plan::new(OperationDescriptor::sum_product(operands, None, output, output_modes, compute)?)?
13        .execute(&[a, b], &[1.0])
14}
15
16/// Reduce modes absent from `output_modes`, in the specified output order.
17pub fn reduce<R: Runtime>(
18    input: &RudaTensor<R>, input_modes: &[Mode], output_modes: &[Mode], reduction: ReductionOp,
19) -> Result<RudaTensor<R>> {
20    let operand = OperandDescriptor::from_tensor(input, input_modes)?;
21    let output = infer_output(std::slice::from_ref(&operand), output_modes, input.dtype)?;
22    let compute = ComputeType::for_operands(std::slice::from_ref(&operand));
23    Plan::new(OperationDescriptor::reduction(operand, None, output, output_modes, reduction, compute)?)?
24        .execute(&[input], &[1.0])
25}
26
27/// Materialize a permutation in a new packed row-major tensor; do not return an aliased view.
28pub fn permute<R: Runtime>(input: &RudaTensor<R>, axes: &[usize]) -> Result<RudaTensor<R>> {
29    let rank = input.meta.rank();
30    if axes.len() != rank || axes.iter().any(|&axis| axis >= rank)
31        || axes.iter().enumerate().any(|(i, axis)| axes[..i].contains(axis)) {
32        return Err(Error::InvalidOperation("axes must be a permutation of 0..rank".into()));
33    }
34    let modes = (0..rank).map(|axis| Mode::try_from(axis).map_err(|_| Error::Overflow)).collect::<Result<Vec<_>>>()?;
35    let output_modes: Vec<_> = axes.iter().map(|&axis| modes[axis]).collect();
36    let operand = OperandDescriptor::from_tensor(input, &modes)?;
37    let output = infer_output(std::slice::from_ref(&operand), &output_modes, input.dtype)?;
38    let compute = ComputeType::for_operands(std::slice::from_ref(&operand));
39    Plan::new(OperationDescriptor::permutation(operand, output, &output_modes, compute)?)?
40        .execute(&[input], &[1.0])
41}
42
43/// Named-axis broadcasting followed by a binary operation.
44pub fn elementwise_binary<R: Runtime>(
45    a: &RudaTensor<R>, modes_a: &[Mode], b: &RudaTensor<R>, modes_b: &[Mode],
46    output_modes: &[Mode], combine: BinaryOp,
47) -> Result<RudaTensor<R>> {
48    let a_desc = OperandDescriptor::from_tensor(a, modes_a)?;
49    let b_desc = OperandDescriptor::from_tensor(b, modes_b)?;
50    let operands = [a_desc, b_desc];
51    let output = infer_output(&operands, output_modes, output_dtype(&operands))?;
52    let compute = ComputeType::for_operands(&operands);
53    let [a_desc, b_desc] = operands;
54    Plan::new(OperationDescriptor::elementwise_binary(a_desc, b_desc, output, output_modes, combine, compute)?)?
55        .execute(&[a, b], &[1.0, 1.0])
56}
57
58/// Combine three broadcast operands as `(A op_ab B) op_abc C`.
59pub fn elementwise_trinary<R: Runtime>(
60    inputs: [(&RudaTensor<R>, &[Mode]); 3], output_modes: &[Mode],
61    combine_ab: BinaryOp, combine_abc: BinaryOp,
62) -> Result<RudaTensor<R>> {
63    let operands = [OperandDescriptor::from_tensor(inputs[0].0, inputs[0].1)?,
64        OperandDescriptor::from_tensor(inputs[1].0, inputs[1].1)?,
65        OperandDescriptor::from_tensor(inputs[2].0, inputs[2].1)?];
66    let output = infer_output(&operands, output_modes, output_dtype(&operands))?;
67    let compute = ComputeType::for_operands(&operands);
68    Plan::new(OperationDescriptor::elementwise_trinary(operands, output, output_modes, combine_ab, combine_abc, compute)?)?
69        .execute(&[inputs[0].0, inputs[1].0, inputs[2].0], &[1.0, 1.0, 1.0])
70}