Skip to main content

laddu_kernel/ir/
instruction.rs

1use super::*;
2
3impl KernelInstruction {
4    /// Returns the stable instruction name used in diagnostics and support reporting.
5    pub fn diagnostic_name(&self) -> &'static str {
6        match self {
7            Self::Cached(_) => "Cached",
8            Self::RealConstant(_) => "RealConstant",
9            Self::ComplexConstant(_) => "ComplexConstant",
10            Self::Parameter(_) => "Parameter",
11            Self::Unary { .. } => "Unary",
12            Self::Binary { .. } => "Binary",
13            Self::Add(_) => "Add",
14            Self::Mul(_) => "Mul",
15            Self::Complex { .. } => "Complex",
16            Self::Vector(_) => "Vector",
17            Self::Matrix { .. } => "Matrix",
18            Self::Component { .. } => "Component",
19            Self::MatrixElement { .. } => "MatrixElement",
20            Self::MatMul { .. } => "MatMul",
21            Self::MatVec { .. } => "MatVec",
22            Self::Dot { .. } => "Dot",
23            Self::Solve { .. } => "Solve",
24            Self::SolveRow { .. } => "SolveRow",
25            Self::SolveRowAdjointElement { .. } => "SolveRowAdjointElement",
26        }
27    }
28
29    /// Returns the backend-independent event-dependence rule for this instruction.
30    pub fn event_dependence(&self) -> KernelEventDependence {
31        match self {
32            Self::Cached(_) | Self::SolveRow { .. } | Self::SolveRowAdjointElement { .. } => {
33                KernelEventDependence::Event
34            }
35            Self::RealConstant(_) | Self::ComplexConstant(_) | Self::Parameter(_) => {
36                KernelEventDependence::Invariant
37            }
38            _ => KernelEventDependence::Operands,
39        }
40    }
41
42    /// Returns the direct input value identifiers.
43    pub fn operands(&self) -> Vec<KernelValueId> {
44        let mut operands = Vec::new();
45        self.for_each_operand(|operand| operands.push(operand));
46        operands
47    }
48
49    /// Visits direct operands in their instruction order without allocating.
50    ///
51    /// The owned [`Self::operands`] method remains available for callers that
52    /// need a collected list. Backend dependency and liveness analysis should
53    /// prefer this callback when it only needs to traverse the operands.
54    pub fn for_each_operand(&self, mut visit: impl FnMut(KernelValueId)) {
55        match self {
56            Self::Cached(_)
57            | Self::RealConstant(_)
58            | Self::ComplexConstant(_)
59            | Self::Parameter(_) => {}
60            Self::Unary { input, .. }
61            | Self::Component { input, .. }
62            | Self::MatrixElement { input, .. }
63            | Self::SolveRowAdjointElement { adjoint: input, .. } => visit(*input),
64            Self::Binary { lhs, rhs, .. } | Self::MatMul { lhs, rhs } | Self::Dot { lhs, rhs } => {
65                visit(*lhs);
66                visit(*rhs);
67            }
68            Self::MatVec { matrix, vector }
69            | Self::Solve {
70                matrix,
71                rhs: vector,
72            } => {
73                visit(*matrix);
74                visit(*vector);
75            }
76            Self::Add(values)
77            | Self::Mul(values)
78            | Self::Vector(values)
79            | Self::SolveRow { rhs: values, .. } => values.iter().copied().for_each(visit),
80            Self::Complex { re, im } => {
81                visit(*re);
82                visit(*im);
83            }
84            Self::Matrix { elements, .. } => elements.iter().copied().for_each(visit),
85        }
86    }
87
88    pub(super) fn validate_operand_order(&self, value: usize) -> Result<(), KernelIrError> {
89        let mut invalid_operand = None;
90        self.for_each_operand(|operand| {
91            if invalid_operand.is_none() && operand.index() >= value {
92                invalid_operand = Some(operand.index());
93            }
94        });
95        invalid_operand.map_or(Ok(()), |operand| {
96            Err(KernelIrError::InvalidOperand { value, operand })
97        })
98    }
99}