Skip to main content

rutensor/
plan.rs

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/// A reusable index map and execution description, independent of buffer addresses.
10#[derive(Clone, Debug)]
11pub struct Plan {
12    operation: OperationDescriptor,
13    modes: Vec<Mode>,
14    reduction_extents: Vec<usize>,
15    reduction_count: usize,
16    /// Logical mode index of every physical axis of each operand.
17    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        // Broadcasting 0 with 1 produces an empty dimension, not a size-one dimension.
26        extents.insert(mode, if previous == 1 { value } else { previous });
27    } else {
28        extents.insert(mode, value);
29    }
30    Ok(())
31}
32
33impl Plan {
34    /// Resolve free, broadcast, diagonal and reduction axes without executing a kernel.
35    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            // Explicit output descriptors may introduce broadcast-only axes.
44            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    /// Allocate D and enqueue the operation. All inputs remain unchanged.
92    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    /// Write into an exclusively owned D buffer and return it. D cannot alias any input.
104    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}