Skip to main content

ocas_eval/
evaluator.rs

1//! Stack-based expression evaluator.
2//!
3//! The [`ExpressionEvaluator`] executes a pre-compiled instruction sequence
4//! on a flat stack, producing numeric results from input parameters.
5
6use crate::domain::{EvaluationDomain, PowfExtension};
7use crate::error::{EvaluationError, Result};
8use crate::function_map::FunctionMap;
9use crate::instruction::Instr;
10
11/// A compiled expression ready for numeric evaluation.
12///
13/// The evaluator holds a sequence of [`Instr`]s and a pre-allocated stack.
14/// Call [`evaluate`](ExpressionEvaluator::evaluate) with parameter values
15/// to compute the result.
16pub struct ExpressionEvaluator<T: EvaluationDomain> {
17    /// The instruction sequence to execute.
18    instructions: Vec<Instr>,
19    /// Number of parameter slots at the start of the stack.
20    param_count: usize,
21    /// Number of constant slots (after params).
22    #[allow(dead_code)]
23    const_count: usize,
24    /// Total stack size (params + constants + temporaries + outputs).
25    stack_size: usize,
26    /// Indices of result slots in the stack.
27    result_indices: Vec<usize>,
28    /// Pre-computed constant values.
29    constants: Vec<T>,
30    /// Optional user-defined function registry.
31    function_map: Option<FunctionMap<T>>,
32}
33
34impl<T: EvaluationDomain + PowfExtension> ExpressionEvaluator<T> {
35    /// Create an evaluator from compiled instruction data.
36    ///
37    /// This is a low-level constructor used by the compiler. Most users
38    /// should use [`ExpressionEvaluator::compile`] instead.
39    #[allow(dead_code)]
40    pub(crate) fn new(
41        instructions: Vec<Instr>,
42        param_count: usize,
43        const_count: usize,
44        stack_size: usize,
45        result_indices: Vec<usize>,
46        constants: Vec<T>,
47    ) -> Self {
48        Self {
49            instructions,
50            param_count,
51            const_count,
52            stack_size,
53            result_indices,
54            constants,
55            function_map: None,
56        }
57    }
58
59    /// Create an evaluator with a function map for user-defined functions.
60    #[allow(dead_code)]
61    pub(crate) fn new_with_functions(
62        instructions: Vec<Instr>,
63        param_count: usize,
64        const_count: usize,
65        stack_size: usize,
66        result_indices: Vec<usize>,
67        constants: Vec<T>,
68        function_map: FunctionMap<T>,
69    ) -> Self {
70        Self {
71            instructions,
72            param_count,
73            const_count,
74            stack_size,
75            result_indices,
76            constants,
77            function_map: Some(function_map),
78        }
79    }
80
81    /// Return the number of parameters expected by this evaluator.
82    pub fn param_count(&self) -> usize {
83        self.param_count
84    }
85
86    /// Evaluate the expression with the given parameter values.
87    ///
88    /// Returns a vector of result values. The number of results equals
89    /// `result_indices.len()`.
90    ///
91    /// # Errors
92    ///
93    /// Returns [`EvaluationError`] if the number of parameters does not
94    /// match, or if an arithmetic error occurs (e.g. division by zero).
95    ///
96    /// # Example
97    ///
98    /// ```ignore
99    /// let result = evaluator.evaluate(&[2.0_f64])?;
100    /// assert_eq!(result.len(), 1);
101    /// ```
102    pub fn evaluate(&self, params: &[T]) -> Result<Vec<T>> {
103        if params.len() != self.param_count {
104            return Err(EvaluationError::WrongArity {
105                name: "<expr>".into(),
106                expected: self.param_count,
107                got: params.len(),
108            });
109        }
110
111        let mut stack: Vec<T> = vec![T::zero(); self.stack_size];
112
113        // Fill parameters
114        for (i, p) in params.iter().enumerate() {
115            stack[i] = p.clone();
116        }
117
118        // Fill constants
119        for (i, c) in self.constants.iter().enumerate() {
120            stack[self.param_count + i] = c.clone();
121        }
122
123        // Execute instructions
124        for instr in &self.instructions {
125            match instr {
126                Instr::Add { dst, srcs } => {
127                    let mut sum = stack[srcs[0]].clone();
128                    for idx in &srcs[1..] {
129                        sum = sum.add_ref(&stack[*idx]);
130                    }
131                    stack[*dst] = sum;
132                }
133                Instr::Mul { dst, srcs } => {
134                    let mut prod = stack[srcs[0]].clone();
135                    for idx in &srcs[1..] {
136                        prod = prod.mul_ref(&stack[*idx]);
137                    }
138                    stack[*dst] = prod;
139                }
140                Instr::Pow { dst, base, exp } => {
141                    stack[*dst] = stack[*base].powi_ref(*exp);
142                }
143                Instr::Powf { dst, base, exp } => {
144                    let result = stack[*base].powf_ref(&stack[*exp])?;
145                    stack[*dst] = result;
146                }
147                Instr::BuiltinOp { dst, op, src } => {
148                    let name = match op {
149                        crate::instruction::BuiltinOp::Sin => "sin",
150                        crate::instruction::BuiltinOp::Cos => "cos",
151                        crate::instruction::BuiltinOp::Tan => "tan",
152                        crate::instruction::BuiltinOp::Sec => "sec",
153                        crate::instruction::BuiltinOp::Csc => "csc",
154                        crate::instruction::BuiltinOp::Cot => "cot",
155                        crate::instruction::BuiltinOp::Exp => "exp",
156                        crate::instruction::BuiltinOp::Log => "log",
157                        crate::instruction::BuiltinOp::Sqrt => "sqrt",
158                        crate::instruction::BuiltinOp::Abs => "abs",
159                    };
160                    let result = T::resolve_builtin(name, &stack[*src])?;
161                    stack[*dst] = result;
162                }
163                Instr::ExternalFun { dst, fn_idx, srcs } => {
164                    let args: Vec<T> = srcs.iter().map(|&i| stack[i].clone()).collect();
165                    let result = self
166                        .function_map
167                        .as_ref()
168                        .and_then(|fm| fm.call_by_index(*fn_idx, &args))
169                        .ok_or_else(|| EvaluationError::FunctionNotFound {
170                            name: format!("external function at index {fn_idx}"),
171                        })?;
172                    stack[*dst] = result;
173                }
174                Instr::Copy { dst, src } => {
175                    stack[*dst] = stack[*src].clone();
176                }
177            }
178        }
179
180        // Collect results
181        let results: Vec<T> = self
182            .result_indices
183            .iter()
184            .map(|&i| stack[i].clone())
185            .collect();
186
187        Ok(results)
188    }
189}
190
191/// SIMD batch evaluation support.
192#[cfg(feature = "simd")]
193impl ExpressionEvaluator<f64> {
194    /// Compile this evaluator into a [`VectorEvaluator`] for batch SIMD evaluation.
195    ///
196    /// The resulting evaluator processes multiple input values simultaneously
197    /// using the best available SIMD width (SSE2/AVX2/AVX-512).
198    ///
199    /// # Errors
200    ///
201    /// Returns [`EvaluationError::UnsupportedOperation`] if the expression
202    /// contains external functions, which are not supported in SIMD mode.
203    pub fn compile_vector_evaluator(&self) -> Result<crate::simd::VectorEvaluator> {
204        // Check for unsupported instructions
205        for instr in &self.instructions {
206            if let Instr::ExternalFun { .. } = instr {
207                return Err(EvaluationError::UnsupportedOperation {
208                    message: "external functions not supported in SIMD mode".into(),
209                });
210            }
211        }
212
213        Ok(crate::simd::VectorEvaluator::new(
214            self.instructions.clone(),
215            self.param_count,
216            self.const_count,
217            self.stack_size,
218            self.result_indices.clone(),
219            self.constants.clone(),
220        ))
221    }
222}
223
224#[cfg(test)]
225mod tests {
226    use super::*;
227
228    fn make_simple_evaluator() -> ExpressionEvaluator<f64> {
229        // Evaluate: x + 1
230        // Stack layout: [param(0)=x] [const(0)=1.0] [temp(0)=result]
231        // Instructions: Add(temp(0), param(0), const(0))
232        let instructions = vec![Instr::Add {
233            dst: 2,
234            srcs: vec![0, 1],
235        }];
236        let constants = vec![1.0f64];
237        ExpressionEvaluator::new(instructions, 1, 1, 3, vec![2], constants)
238    }
239
240    #[test]
241    fn simple_add() {
242        let eval = make_simple_evaluator();
243        assert_eq!(eval.param_count(), 1);
244        let result = eval.evaluate(&[2.0]).unwrap();
245        assert!((result[0] - 3.0).abs() < 1e-10);
246    }
247
248    #[test]
249    fn wrong_param_count() {
250        let eval = make_simple_evaluator();
251        assert!(eval.evaluate(&[1.0, 2.0]).is_err());
252        assert!(eval.evaluate(&[]).is_err());
253    }
254
255    #[test]
256    fn mul_expression() {
257        // Evaluate: x * 2
258        let instructions = vec![Instr::Mul {
259            dst: 2,
260            srcs: vec![0, 1],
261        }];
262        let constants = vec![2.0f64];
263        let eval = ExpressionEvaluator::new(instructions, 1, 1, 3, vec![2], constants);
264        assert!((eval.evaluate(&[3.0]).unwrap()[0] - 6.0).abs() < 1e-10);
265    }
266
267    #[test]
268    fn pow_expression() {
269        // Evaluate: x^3
270        let instructions = vec![Instr::Pow {
271            dst: 1,
272            base: 0,
273            exp: 3,
274        }];
275        let eval = ExpressionEvaluator::new(instructions, 1, 0, 2, vec![1], vec![]);
276        assert!((eval.evaluate(&[2.0]).unwrap()[0] - 8.0).abs() < 1e-10);
277    }
278
279    #[test]
280    fn builtin_sin() {
281        // Evaluate: sin(x)
282        let instructions = vec![Instr::BuiltinOp {
283            dst: 1,
284            op: crate::instruction::BuiltinOp::Sin,
285            src: 0,
286        }];
287        let eval = ExpressionEvaluator::new(instructions, 1, 0, 2, vec![1], vec![]);
288        let result = eval.evaluate(&[std::f64::consts::FRAC_PI_2]).unwrap();
289        assert!((result[0] - 1.0).abs() < 1e-10);
290    }
291
292    #[test]
293    fn copy_instruction() {
294        // Evaluate: x (identity)
295        let instructions = vec![Instr::Copy { dst: 1, src: 0 }];
296        let eval = ExpressionEvaluator::new(instructions, 1, 0, 2, vec![1], vec![]);
297        assert!((eval.evaluate(&[42.0]).unwrap()[0] - 42.0).abs() < 1e-10);
298    }
299}