Skip to main content

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;