1use std::collections::BTreeMap;
2use crate::{Error, Mode, OperationDescriptor, Result};
3use crate::descriptor::product;
4use crate::kernel::{Config, Launch, convert_operand};
5use crate::operation::Kind;
6use ruda_core::{device::Device, tensor::{Metadata, Shape}};
7use ruda_kernel::{dsl::Runtime, tensor::RudaTensor};
8
9#[derive(Clone, Debug)]
11pub struct Plan {
12 operation: OperationDescriptor,
13 modes: Vec<Mode>,
14 reduction_extents: Vec<usize>,
15 reduction_count: usize,
16 axis_modes: Vec<Vec<usize>>,
18}
19
20pub(crate) fn merge_extent(extents: &mut BTreeMap<Mode, usize>, mode: Mode, value: usize) -> Result<()> {
21 if let Some(&previous) = extents.get(&mode) {
22 if previous != value && previous != 1 && value != 1 {
23 return Err(Error::IncompatibleExtent { mode, left: previous, right: value });
24 }
25 extents.insert(mode, if previous == 1 { value } else { previous });
27 } else {
28 extents.insert(mode, value);
29 }
30 Ok(())
31}
32
33impl Plan {
34 pub fn new(operation: OperationDescriptor) -> Result<Self> {
36 let mut extents = BTreeMap::new();
37 for input in &operation.inputs {
38 for (&mode, &extent) in input.modes.iter().zip(&input.tensor.extents) {
39 merge_extent(&mut extents, mode, extent)?;
40 }
41 }
42 for (&mode, &extent) in operation.output_modes.iter().zip(&operation.output.extents) {
43 merge_extent(&mut extents, mode, extent)?;
45 if extents[&mode] != extent {
46 return Err(Error::InvalidOperation(format!("output extent for mode {mode} would discard values")));
47 }
48 }
49 let mut reduction_modes = Vec::new();
50 for input in &operation.inputs[..operation.terms] {
51 for &mode in &input.modes {
52 if !operation.output_modes.contains(&mode) && !reduction_modes.contains(&mode) {
53 reduction_modes.push(mode);
54 }
55 }
56 }
57 if operation.addend && operation.inputs[operation.terms].modes.iter()
58 .any(|mode| !operation.output_modes.contains(mode)) {
59 return Err(Error::InvalidOperation("C may only contain output modes".into()));
60 }
61 if !reduction_modes.is_empty() && matches!(operation.kind, Kind::Permutation | Kind::Elementwise(_, _)) {
62 return Err(Error::InvalidOperation("elementwise and permutation operations cannot discard modes".into()));
63 }
64 if operation.kind == Kind::Permutation {
65 let input = &operation.inputs[0];
66 if input.modes.len() != operation.output_modes.len()
67 || input.modes.iter().any(|mode| !operation.output_modes.contains(mode))
68 || input.modes.iter().enumerate().any(|(i, mode)| input.modes[..i].contains(mode)) {
69 return Err(Error::InvalidOperation("a permutation must reorder each input mode exactly once".into()));
70 }
71 for (i, mode) in input.modes.iter().enumerate() {
72 if input.tensor.extents[i] != extents[mode] {
73 return Err(Error::InvalidOperation("a permutation cannot expand an input axis".into()));
74 }
75 }
76 }
77 let reduction_extents: Vec<_> = reduction_modes.iter().map(|m| extents[m]).collect();
78 let reduction_count = product(&reduction_extents)?;
79 let mut modes = operation.output_modes.clone();
80 modes.extend(reduction_modes);
81 let axis_modes = operation.inputs.iter().map(|input| input.modes.iter()
82 .map(|m| modes.iter().position(|x| x == m).expect("validated mode"))
83 .collect()).collect();
84 Ok(Self { operation, modes, reduction_extents, reduction_count, axis_modes })
85 }
86
87 pub fn operation(&self) -> &OperationDescriptor { &self.operation }
88 pub fn reduction_extents(&self) -> &[usize] { &self.reduction_extents }
89 pub fn reduction_elements(&self) -> usize { self.reduction_count }
90
91 pub fn execute<R: Runtime>(&self, inputs: &[&RudaTensor<R>], scalars: &[f64]) -> Result<RudaTensor<R>> {
93 self.validate_inputs(inputs, scalars)?;
94 let reference = inputs[0];
95 let descriptor = &self.operation.output;
96 let handle = reference.client.empty(descriptor.storage_bytes().max(descriptor.dtype.size()));
97 let output = RudaTensor::new(reference.client.clone(), handle,
98 Metadata::new(Shape::from(descriptor.extents.clone()), descriptor.strides.clone()),
99 reference.device.clone(), descriptor.dtype);
100 self.submit(inputs, output, scalars)
101 }
102
103 pub fn execute_into<R: Runtime>(
105 &self, inputs: &[&RudaTensor<R>], output: RudaTensor<R>, scalars: &[f64],
106 ) -> Result<RudaTensor<R>> {
107 self.validate_inputs(inputs, scalars)?;
108 if !self.operation.output.matches(&output) { return Err(Error::OutputMismatch); }
109 self.operation.output.check_buffer(&output)?;
110 if inputs[0].device.to_id() != output.device.to_id() { return Err(Error::DeviceMismatch); }
111 if !output.can_mut() { return Err(Error::SharedOutput); }
112 self.submit(inputs, output, scalars)
113 }
114
115 fn validate_inputs<R: Runtime>(&self, inputs: &[&RudaTensor<R>], scalars: &[f64]) -> Result<()> {
116 if inputs.len() != self.operation.inputs.len() {
117 return Err(Error::InputCount { expected: self.operation.inputs.len(), actual: inputs.len() });
118 }
119 if scalars.len() != self.operation.scalar_count() {
120 return Err(Error::ScalarCount { expected: self.operation.scalar_count(), actual: scalars.len() });
121 }
122 for (index, (input, descriptor)) in inputs.iter().zip(&self.operation.inputs).enumerate() {
123 if !descriptor.tensor.matches(input) { return Err(Error::TensorMismatch { input: index }); }
124 descriptor.tensor.check_buffer(input)?;
125 if inputs[0].device.to_id() != input.device.to_id() { return Err(Error::DeviceMismatch); }
126 }
127 Ok(())
128 }
129
130 fn submit<R: Runtime>(
131 &self, inputs: &[&RudaTensor<R>], output: RudaTensor<R>, scalars: &[f64],
132 ) -> Result<RudaTensor<R>> {
133 let count = self.operation.output.num_elements();
134 if count == 0 { return Ok(output); }
135 let hardware = &output.client.properties().hardware;
136 let mut maximum = (hardware.max_ruda_dim.0 as usize)
137 .min(hardware.max_units_per_ruda as usize).min(128);
138 if !self.reduction_extents.is_empty() {
139 maximum = maximum.min(hardware.max_shared_memory_size / self.operation.compute.dtype().size());
140 }
141 if maximum == 0 {
142 return Err(Error::UnsupportedDevice("no workgroup threads available".into()));
143 }
144 let desired = maximum.min(self.reduction_count.max(1));
145 let threads = 1usize << desired.ilog2();
146 count.checked_mul(threads).ok_or(Error::Overflow)?;
147 let compute = self.operation.compute;
148 let converted = inputs.iter().map(|input| convert_operand(input, compute.dtype()))
149 .collect::<Result<Vec<_>>>()?;
150 let stride_count = self.modes.len().checked_mul(inputs.len()).ok_or(Error::Overflow)?;
151 let mut strides = vec![0usize; stride_count];
152 for (operand, tensor) in converted.iter().enumerate() {
153 for (axis, &mode) in self.axis_modes[operand].iter().enumerate() {
154 if tensor.meta.shape()[axis] != 1 {
155 let position = operand * self.modes.len() + mode;
156 strides[position] = strides[position].checked_add(tensor.meta.strides()[axis])
157 .ok_or(Error::Overflow)?;
158 }
159 }
160 }
161 Launch {
162 inputs: &converted, output: &output, strides,
163 reduction_extents: &self.reduction_extents, scalars, count,
164 reduction_count: self.reduction_count, threads, compute,
165 config: Config {
166 kind: self.operation.kind,
167 unary: self.operation.inputs.iter().map(|x| x.unary).collect(),
168 terms: self.operation.terms, addend: self.operation.addend,
169 },
170 }.submit()?;
171 Ok(output)
172 }
173}