laddu_kernel/ir.rs
1pub use crate::KernelIrError;
2use laddu_expr::{BinaryOp, UnaryOp, parameters::ParamId};
3use num::complex::Complex64;
4
5/// Stable identifier for a value in kernel IR.
6#[derive(Copy, Clone, Debug, PartialEq, Eq)]
7pub struct KernelValueId(usize);
8
9impl KernelValueId {
10 /// Creates an identifier from a zero-based value index.
11 pub fn from_index(index: usize) -> Self {
12 Self(index)
13 }
14
15 /// Returns the zero-based value index.
16 pub fn index(self) -> usize {
17 self.0
18 }
19}
20
21/// Runtime shape and scalar representation of a kernel value.
22#[derive(Copy, Clone, Debug, PartialEq, Eq)]
23pub enum KernelValueKind {
24 /// A real scalar.
25 Real,
26 /// A complex scalar.
27 Complex,
28 /// A complex vector.
29 Vector {
30 /// Number of elements.
31 len: usize,
32 },
33 /// A complex matrix.
34 Matrix {
35 /// Number of rows.
36 rows: usize,
37 /// Number of columns.
38 cols: usize,
39 },
40}
41
42impl KernelValueKind {
43 /// Returns the number of logical scalar elements.
44 ///
45 /// # Panics
46 ///
47 /// Panics when matrix dimensions exceed the addressable `usize` width.
48 /// Validated kernel IR rejects such dimensions during construction.
49 pub fn width(self) -> usize {
50 match self {
51 Self::Real | Self::Complex => 1,
52 Self::Vector { len } => len,
53 Self::Matrix { rows, cols } => checked_matrix_width(rows, cols)
54 .expect("kernel matrix dimensions exceed addressable width"),
55 }
56 }
57
58 /// Returns the checked row-major element index for a matrix value.
59 ///
60 /// Returns `None` when this is not a matrix, either coordinate is out of
61 /// bounds, or the dimensions/index arithmetic cannot be represented by
62 /// `usize`.
63 pub fn checked_row_major_index(self, row: usize, col: usize) -> Option<usize> {
64 let Self::Matrix { rows, cols } = self else {
65 return None;
66 };
67 checked_row_major_index(rows, cols, row, col)
68 }
69
70 fn scalar_combine(self, rhs: Self) -> Option<Self> {
71 match (self, rhs) {
72 (Self::Real, Self::Real) => Some(Self::Real),
73 (Self::Real | Self::Complex, Self::Real | Self::Complex) => Some(Self::Complex),
74 _ => None,
75 }
76 }
77
78 fn is_scalar(self) -> bool {
79 matches!(self, Self::Real | Self::Complex)
80 }
81}
82
83fn checked_matrix_width(rows: usize, cols: usize) -> Option<usize> {
84 rows.checked_mul(cols)
85}
86
87fn checked_row_major_index(rows: usize, cols: usize, row: usize, col: usize) -> Option<usize> {
88 checked_matrix_width(rows, cols)?;
89 if row >= rows || col >= cols {
90 return None;
91 }
92 row.checked_mul(cols)?.checked_add(col)
93}
94
95/// Whether a kernel value is constant across events or event-dependent.
96#[derive(Copy, Clone, Debug, PartialEq, Eq)]
97pub enum KernelValueClass {
98 /// The value depends only on constants and parameters.
99 Invariant,
100 /// The value depends on event data or cache inputs.
101 Event,
102}
103
104/// How an instruction's event dependence is determined.
105#[derive(Copy, Clone, Debug, PartialEq, Eq)]
106pub enum KernelEventDependence {
107 /// The instruction is invariant regardless of its operands.
108 Invariant,
109 /// The instruction is event-dependent regardless of its operands.
110 Event,
111 /// The instruction is event-dependent when any direct operand is event-dependent.
112 Operands,
113}
114
115/// Operation that produces one value in kernel IR.
116#[derive(Clone, Debug)]
117pub enum KernelInstruction {
118 /// Reads a precomputed cache slot.
119 Cached(usize),
120 /// Emits a real constant.
121 RealConstant(f64),
122 /// Emits a complex constant.
123 ComplexConstant(Complex64),
124 /// Reads a scalar parameter.
125 Parameter(ParamId),
126 /// Applies a unary scalar operation.
127 Unary {
128 /// Operation to apply.
129 op: UnaryOp,
130 /// Input value.
131 input: KernelValueId,
132 },
133 /// Applies a binary scalar operation.
134 Binary {
135 /// Operation to apply.
136 op: BinaryOp,
137 /// Left operand.
138 lhs: KernelValueId,
139 /// Right operand.
140 rhs: KernelValueId,
141 },
142 /// Adds scalar operands.
143 Add(Vec<KernelValueId>),
144 /// Multiplies scalar operands.
145 Mul(Vec<KernelValueId>),
146 /// Constructs a complex scalar from real components.
147 Complex {
148 /// Real component.
149 re: KernelValueId,
150 /// Imaginary component.
151 im: KernelValueId,
152 },
153 /// Constructs a vector from scalar elements.
154 Vector(Vec<KernelValueId>),
155 /// Constructs a row-major matrix.
156 Matrix {
157 /// Number of rows.
158 rows: usize,
159 /// Number of columns.
160 cols: usize,
161 /// Row-major scalar elements.
162 elements: Vec<KernelValueId>,
163 },
164 /// Selects a vector component.
165 Component {
166 /// Vector input.
167 input: KernelValueId,
168 /// Zero-based component index.
169 index: usize,
170 },
171 /// Selects a matrix element.
172 MatrixElement {
173 /// Matrix input.
174 input: KernelValueId,
175 /// Zero-based row index.
176 row: usize,
177 /// Zero-based column index.
178 col: usize,
179 },
180 /// Multiplies two matrices.
181 MatMul {
182 /// Left matrix.
183 lhs: KernelValueId,
184 /// Right matrix.
185 rhs: KernelValueId,
186 },
187 /// Multiplies a matrix by a vector.
188 MatVec {
189 /// Matrix operand.
190 matrix: KernelValueId,
191 /// Vector operand.
192 vector: KernelValueId,
193 },
194 /// Computes a vector dot product.
195 Dot {
196 /// Left vector.
197 lhs: KernelValueId,
198 /// Right vector.
199 rhs: KernelValueId,
200 },
201 /// Solves a linear system.
202 Solve {
203 /// Coefficient matrix.
204 matrix: KernelValueId,
205 /// Right-hand-side vector.
206 rhs: KernelValueId,
207 },
208 /// Evaluates one row of a specialized cached solve.
209 SolveRow {
210 /// Cache slot containing the decomposed matrix row data.
211 row_slot: usize,
212 /// Right-hand-side scalar values.
213 rhs: Vec<KernelValueId>,
214 },
215 /// Evaluates one adjoint element for a specialized solve row.
216 SolveRowAdjointElement {
217 /// Cache slot containing the decomposed matrix row data.
218 row_slot: usize,
219 /// Element index within the row.
220 index: usize,
221 /// Row length.
222 len: usize,
223 /// Incoming scalar adjoint.
224 adjoint: KernelValueId,
225 },
226}
227
228/// Typed instruction and evaluation class for one kernel IR value.
229#[derive(Clone, Debug)]
230pub struct KernelValue {
231 /// Value shape and scalar representation.
232 pub kind: KernelValueKind,
233 /// Event-dependency class.
234 pub class: KernelValueClass,
235 /// Instruction that produces the value.
236 pub instruction: KernelInstruction,
237}
238
239/// Validated IR for a kernel with one scalar output.
240#[derive(Clone, Debug)]
241pub struct ScalarKernelIr {
242 values: Vec<KernelValue>,
243 root: KernelValueId,
244}
245
246/// Validated IR for a kernel that populates multiple cache outputs.
247#[derive(Clone, Debug)]
248pub struct CacheKernelIr {
249 values: Vec<KernelValue>,
250 outputs: Vec<KernelValueId>,
251}
252
253/// Scalar component of a complex primal output to differentiate.
254#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash)]
255pub enum OutputComponent {
256 /// Differentiate the real component.
257 Real,
258 /// Differentiate the imaginary component.
259 Imag,
260}
261
262/// Validated IR containing a primal computation and real gradient outputs.
263#[derive(Clone, Debug)]
264pub struct GradientKernelIr {
265 values: Vec<KernelValue>,
266 primal_root: KernelValueId,
267 outputs: Vec<KernelValueId>,
268 component: OutputComponent,
269}
270
271/// Builder for appending type-checked instructions to existing scalar IR.
272#[derive(Clone, Debug)]
273pub struct KernelIrBuilder {
274 values: Vec<KernelValue>,
275}
276
277mod builder;
278mod instruction;
279mod validate;
280mod wrappers;
281
282use validate::validate_graph;
283
284#[cfg(test)]
285mod tests;