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
5pub 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
16pub 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
27pub 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
43pub 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
58pub 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}