Skip to main content

laddu_kernel/ir/
wrappers.rs

1use super::*;
2
3fn validate_root_bounds(values: &[KernelValue], root: KernelValueId) -> Result<(), KernelIrError> {
4    if root.index() >= values.len() {
5        return Err(KernelIrError::RootOutOfBounds {
6            root: root.index(),
7            len: values.len(),
8        });
9    }
10    Ok(())
11}
12
13fn validate_scalar_root(
14    values: &[KernelValue],
15    root: KernelValueId,
16    operation: &'static str,
17    message: &'static str,
18) -> Result<(), KernelIrError> {
19    if !values[root.index()].kind.is_scalar() {
20        return Err(KernelInstruction::shape_error(
21            root.index(),
22            operation,
23            message,
24        ));
25    }
26    Ok(())
27}
28
29fn validate_cache_outputs(
30    values: &[KernelValue],
31    outputs: &[KernelValueId],
32) -> Result<(), KernelIrError> {
33    for output in outputs {
34        if output.index() >= values.len() {
35            return Err(KernelIrError::CacheOutputOutOfBounds {
36                output: output.index(),
37                len: values.len(),
38            });
39        }
40    }
41    Ok(())
42}
43
44fn validate_gradient_outputs(
45    values: &[KernelValue],
46    outputs: &[KernelValueId],
47) -> Result<(), KernelIrError> {
48    for output in outputs {
49        let Some(value) = values.get(output.index()) else {
50            return Err(KernelIrError::GradientOutOfBounds {
51                output: output.index(),
52                len: values.len(),
53            });
54        };
55        if value.kind != KernelValueKind::Real {
56            return Err(KernelIrError::GradientKindMismatch {
57                output: output.index(),
58                actual: value.kind,
59            });
60        }
61    }
62    Ok(())
63}
64
65fn required_values(values: &[KernelValue], outputs: &[KernelValueId]) -> Vec<bool> {
66    let mut required = vec![false; values.len()];
67    let mut pending = outputs.to_vec();
68    while let Some(id) = pending.pop() {
69        if required[id.index()] {
70            continue;
71        }
72        required[id.index()] = true;
73        values[id.index()]
74            .instruction
75            .for_each_operand(|operand| pending.push(operand));
76    }
77    required
78}
79
80impl ScalarKernelIr {
81    /// Validates values and constructs a scalar kernel rooted at `root`.
82    ///
83    /// # Errors
84    ///
85    /// Returns [`KernelIrError`] when the value list is empty, `root` or an
86    /// operand is out of bounds, values are not topologically ordered, a
87    /// value's kind or class is inconsistent with its instruction, or the
88    /// root is not scalar.
89    pub fn new(values: Vec<KernelValue>, root: KernelValueId) -> Result<Self, KernelIrError> {
90        let ir = Self { values, root };
91        ir.validate()?;
92        Ok(ir)
93    }
94
95    /// Revalidates ordering, types, classes, and the scalar root.
96    ///
97    /// # Errors
98    ///
99    /// Returns [`KernelIrError`] when the IR is empty, its root or an operand
100    /// is out of bounds, its values are not topologically ordered, or a
101    /// value's kind, class, or shape is inconsistent with its instruction.
102    pub fn validate(&self) -> Result<(), KernelIrError> {
103        if self.values.is_empty() {
104            return validate_graph(&self.values);
105        }
106        validate_root_bounds(&self.values, self.root)?;
107        validate_graph(&self.values)?;
108        validate_scalar_root(
109            &self.values,
110            self.root,
111            "kernel root",
112            "root must be scalar",
113        )
114    }
115
116    /// Returns all IR values in topological order.
117    pub fn values(&self) -> &[KernelValue] {
118        &self.values
119    }
120
121    /// Returns the scalar output identifier.
122    pub fn root(&self) -> KernelValueId {
123        self.root
124    }
125
126    /// Returns a mask of values needed to evaluate the scalar root.
127    pub fn required_values(&self) -> Vec<bool> {
128        required_values(&self.values, std::slice::from_ref(&self.root))
129    }
130}
131
132impl CacheKernelIr {
133    /// Validates values and constructs a cache kernel with the given outputs.
134    ///
135    /// # Errors
136    ///
137    /// Returns [`KernelIrError`] when `outputs` is empty, an output or operand
138    /// is out of bounds, values are not topologically ordered, or a value's
139    /// kind, class, or shape is inconsistent with its instruction.
140    pub fn new(
141        values: Vec<KernelValue>,
142        outputs: Vec<KernelValueId>,
143    ) -> Result<Self, KernelIrError> {
144        if outputs.is_empty() {
145            return Err(KernelIrError::EmptyCacheOutputs);
146        }
147        validate_graph(&values)?;
148        validate_cache_outputs(&values, &outputs)?;
149        Ok(Self { values, outputs })
150    }
151
152    /// Returns all IR values in topological order.
153    pub fn values(&self) -> &[KernelValue] {
154        &self.values
155    }
156
157    /// Returns cache output identifiers in storage order.
158    pub fn outputs(&self) -> &[KernelValueId] {
159        &self.outputs
160    }
161}
162
163impl GradientKernelIr {
164    /// Validates and constructs a gradient kernel.
165    ///
166    /// # Errors
167    ///
168    /// Returns [`KernelIrError`] when the primal IR is invalid, the primal
169    /// root is not scalar, a gradient output is out of bounds, or a gradient
170    /// output is not real-valued.
171    pub fn new(
172        values: Vec<KernelValue>,
173        primal_root: KernelValueId,
174        outputs: Vec<KernelValueId>,
175        component: OutputComponent,
176    ) -> Result<Self, KernelIrError> {
177        let ir = Self {
178            values,
179            primal_root,
180            outputs,
181            component,
182        };
183        ir.validate()?;
184        Ok(ir)
185    }
186
187    /// Revalidates the primal root and real gradient outputs.
188    ///
189    /// # Errors
190    ///
191    /// Returns [`KernelIrError`] when the primal IR is invalid, the primal
192    /// root is not scalar, a gradient output is out of bounds, or a gradient
193    /// output is not real-valued.
194    pub fn validate(&self) -> Result<(), KernelIrError> {
195        if self.values.is_empty() {
196            return validate_graph(&self.values);
197        }
198        validate_root_bounds(&self.values, self.primal_root)?;
199        validate_graph(&self.values)?;
200        validate_scalar_root(
201            &self.values,
202            self.primal_root,
203            "gradient primal root",
204            "primal root must be scalar",
205        )?;
206        validate_gradient_outputs(&self.values, &self.outputs)
207    }
208
209    /// Returns all primal and derivative IR values in topological order.
210    pub fn values(&self) -> &[KernelValue] {
211        &self.values
212    }
213
214    /// Returns the primal scalar output identifier.
215    pub fn primal_root(&self) -> KernelValueId {
216        self.primal_root
217    }
218
219    /// Returns derivative output identifiers.
220    pub fn outputs(&self) -> &[KernelValueId] {
221        &self.outputs
222    }
223
224    /// Returns the differentiated component of the complex primal.
225    pub fn component(&self) -> OutputComponent {
226        self.component
227    }
228
229    /// Returns a mask of values needed to evaluate the derivative outputs.
230    pub fn required_values(&self) -> Vec<bool> {
231        required_values(&self.values, &self.outputs)
232    }
233}