laddu_kernel/ir/
wrappers.rs1use 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 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 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 pub fn values(&self) -> &[KernelValue] {
118 &self.values
119 }
120
121 pub fn root(&self) -> KernelValueId {
123 self.root
124 }
125
126 pub fn required_values(&self) -> Vec<bool> {
128 required_values(&self.values, std::slice::from_ref(&self.root))
129 }
130}
131
132impl CacheKernelIr {
133 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 pub fn values(&self) -> &[KernelValue] {
154 &self.values
155 }
156
157 pub fn outputs(&self) -> &[KernelValueId] {
159 &self.outputs
160 }
161}
162
163impl GradientKernelIr {
164 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 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 pub fn values(&self) -> &[KernelValue] {
211 &self.values
212 }
213
214 pub fn primal_root(&self) -> KernelValueId {
216 self.primal_root
217 }
218
219 pub fn outputs(&self) -> &[KernelValueId] {
221 &self.outputs
222 }
223
224 pub fn component(&self) -> OutputComponent {
226 self.component
227 }
228
229 pub fn required_values(&self) -> Vec<bool> {
231 required_values(&self.values, &self.outputs)
232 }
233}