1use crate::domain::{EvaluationDomain, PowfExtension};
7use crate::error::{EvaluationError, Result};
8use crate::function_map::FunctionMap;
9use crate::instruction::Instr;
10
11pub struct ExpressionEvaluator<T: EvaluationDomain> {
17 instructions: Vec<Instr>,
19 param_count: usize,
21 #[allow(dead_code)]
23 const_count: usize,
24 stack_size: usize,
26 result_indices: Vec<usize>,
28 constants: Vec<T>,
30 function_map: Option<FunctionMap<T>>,
32}
33
34impl<T: EvaluationDomain + PowfExtension> ExpressionEvaluator<T> {
35 #[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 #[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 pub fn param_count(&self) -> usize {
83 self.param_count
84 }
85
86 pub fn result_count(&self) -> usize {
88 self.result_indices.len()
89 }
90
91 pub fn stack_size(&self) -> usize {
96 self.stack_size
97 }
98
99 pub fn evaluate(&self, params: &[T]) -> Result<Vec<T>> {
116 let mut stack: Vec<T> = Vec::with_capacity(self.stack_size);
117 let mut results: Vec<T> = Vec::with_capacity(self.result_indices.len());
118 self.evaluate_with_stack(params, &mut stack, &mut results)?;
119 Ok(results)
120 }
121
122 pub fn evaluate_with_stack(
135 &self,
136 params: &[T],
137 stack: &mut Vec<T>,
138 results: &mut Vec<T>,
139 ) -> Result<()> {
140 if params.len() != self.param_count {
141 return Err(EvaluationError::WrongArity {
142 name: "<expr>".into(),
143 expected: self.param_count,
144 got: params.len(),
145 });
146 }
147
148 stack.clear();
149 stack.resize(self.stack_size, T::zero());
150
151 for (i, p) in params.iter().enumerate() {
153 stack[i] = p.clone();
154 }
155
156 for (i, c) in self.constants.iter().enumerate() {
158 stack[self.param_count + i] = c.clone();
159 }
160
161 for instr in &self.instructions {
163 match instr {
164 Instr::Add { dst, srcs } => {
165 let mut sum = stack[srcs[0]].clone();
166 for idx in &srcs[1..] {
167 sum = sum.add_ref(&stack[*idx]);
168 }
169 stack[*dst] = sum;
170 }
171 Instr::Mul { dst, srcs } => {
172 let mut prod = stack[srcs[0]].clone();
173 for idx in &srcs[1..] {
174 prod = prod.mul_ref(&stack[*idx]);
175 }
176 stack[*dst] = prod;
177 }
178 Instr::Pow { dst, base, exp } => {
179 stack[*dst] = stack[*base].powi_ref(*exp);
180 }
181 Instr::Powf { dst, base, exp } => {
182 let result = stack[*base].powf_ref(&stack[*exp])?;
183 stack[*dst] = result;
184 }
185 Instr::BuiltinOp { dst, op, src } => {
186 let name = match op {
187 crate::instruction::BuiltinOp::Sin => "sin",
188 crate::instruction::BuiltinOp::Cos => "cos",
189 crate::instruction::BuiltinOp::Tan => "tan",
190 crate::instruction::BuiltinOp::Sec => "sec",
191 crate::instruction::BuiltinOp::Csc => "csc",
192 crate::instruction::BuiltinOp::Cot => "cot",
193 crate::instruction::BuiltinOp::Exp => "exp",
194 crate::instruction::BuiltinOp::Log => "log",
195 crate::instruction::BuiltinOp::Sqrt => "sqrt",
196 crate::instruction::BuiltinOp::Abs => "abs",
197 };
198 let result = T::resolve_builtin(name, &stack[*src])?;
199 stack[*dst] = result;
200 }
201 Instr::ExternalFun { dst, fn_idx, srcs } => {
202 let args: Vec<T> = srcs.iter().map(|&i| stack[i].clone()).collect();
203 let result = self
204 .function_map
205 .as_ref()
206 .and_then(|fm| fm.call_by_index(*fn_idx, &args))
207 .ok_or_else(|| EvaluationError::FunctionNotFound {
208 name: format!("external function at index {fn_idx}"),
209 })?;
210 stack[*dst] = result;
211 }
212 Instr::Copy { dst, src } => {
213 stack[*dst] = stack[*src].clone();
214 }
215 }
216 }
217
218 results.clear();
220 results.extend(self.result_indices.iter().map(|&i| stack[i].clone()));
221
222 Ok(())
223 }
224}
225
226#[cfg(feature = "jit")]
228impl ExpressionEvaluator<f64> {
229 pub fn compile_jit(&self) -> Result<crate::jit::JitCompiledFunction> {
242 let constants: Vec<f64> = self.constants.clone();
243 crate::jit::JitEngine::compile(
244 &self.instructions,
245 self.param_count,
246 &constants,
247 &self.result_indices,
248 )
249 }
250
251 pub fn compile_jit_f32(&self) -> Result<crate::jit::JitCompiledF32> {
259 let constants: Vec<f32> = self.constants.iter().map(|&c| c as f32).collect();
260 crate::jit::JitEngine::compile_f32(
261 &self.instructions,
262 self.param_count,
263 &constants,
264 &self.result_indices,
265 )
266 }
267}
268
269#[cfg(feature = "simd")]
271impl ExpressionEvaluator<f64> {
272 pub fn compile_vector_evaluator(&self) -> Result<crate::simd::VectorEvaluator> {
282 for instr in &self.instructions {
284 if let Instr::ExternalFun { .. } = instr {
285 return Err(EvaluationError::UnsupportedOperation {
286 message: "external functions not supported in SIMD mode".into(),
287 });
288 }
289 }
290
291 Ok(crate::simd::VectorEvaluator::new(
292 self.instructions.clone(),
293 self.param_count,
294 self.const_count,
295 self.stack_size,
296 self.result_indices.clone(),
297 self.constants.clone(),
298 ))
299 }
300
301 pub fn compile_vector_evaluator_f32(&self) -> Result<crate::simd::VectorEvaluatorF32> {
311 for instr in &self.instructions {
313 if let Instr::ExternalFun { .. } = instr {
314 return Err(EvaluationError::UnsupportedOperation {
315 message: "external functions not supported in SIMD mode".into(),
316 });
317 }
318 }
319
320 let constants: Vec<f32> = self.constants.iter().map(|&c| c as f32).collect();
321 Ok(crate::simd::VectorEvaluatorF32::new(
322 self.instructions.clone(),
323 self.param_count,
324 self.const_count,
325 self.stack_size,
326 self.result_indices.clone(),
327 constants,
328 ))
329 }
330}
331
332#[cfg(test)]
333mod tests {
334 use super::*;
335
336 fn make_simple_evaluator() -> ExpressionEvaluator<f64> {
337 let instructions = vec![Instr::Add {
341 dst: 2,
342 srcs: vec![0, 1],
343 }];
344 let constants = vec![1.0f64];
345 ExpressionEvaluator::new(instructions, 1, 1, 3, vec![2], constants)
346 }
347
348 #[test]
349 fn simple_add() {
350 let eval = make_simple_evaluator();
351 assert_eq!(eval.param_count(), 1);
352 let result = eval.evaluate(&[2.0]).unwrap();
353 assert!((result[0] - 3.0).abs() < 1e-10);
354 }
355
356 #[test]
357 fn wrong_param_count() {
358 let eval = make_simple_evaluator();
359 assert!(eval.evaluate(&[1.0, 2.0]).is_err());
360 assert!(eval.evaluate(&[]).is_err());
361 }
362
363 #[test]
364 fn mul_expression() {
365 let instructions = vec![Instr::Mul {
367 dst: 2,
368 srcs: vec![0, 1],
369 }];
370 let constants = vec![2.0f64];
371 let eval = ExpressionEvaluator::new(instructions, 1, 1, 3, vec![2], constants);
372 assert!((eval.evaluate(&[3.0]).unwrap()[0] - 6.0).abs() < 1e-10);
373 }
374
375 #[test]
376 fn pow_expression() {
377 let instructions = vec![Instr::Pow {
379 dst: 1,
380 base: 0,
381 exp: 3,
382 }];
383 let eval = ExpressionEvaluator::new(instructions, 1, 0, 2, vec![1], vec![]);
384 assert!((eval.evaluate(&[2.0]).unwrap()[0] - 8.0).abs() < 1e-10);
385 }
386
387 #[test]
388 fn builtin_sin() {
389 let instructions = vec![Instr::BuiltinOp {
391 dst: 1,
392 op: crate::instruction::BuiltinOp::Sin,
393 src: 0,
394 }];
395 let eval = ExpressionEvaluator::new(instructions, 1, 0, 2, vec![1], vec![]);
396 let result = eval.evaluate(&[std::f64::consts::FRAC_PI_2]).unwrap();
397 assert!((result[0] - 1.0).abs() < 1e-10);
398 }
399
400 #[test]
401 fn copy_instruction() {
402 let instructions = vec![Instr::Copy { dst: 1, src: 0 }];
404 let eval = ExpressionEvaluator::new(instructions, 1, 0, 2, vec![1], vec![]);
405 assert!((eval.evaluate(&[42.0]).unwrap()[0] - 42.0).abs() < 1e-10);
406 }
407
408 #[test]
409 fn evaluate_with_stack_reuses_buffers() {
410 let eval = make_simple_evaluator();
411 let mut stack: Vec<f64> = Vec::with_capacity(eval.stack_size());
412 let mut results: Vec<f64> = Vec::with_capacity(eval.result_count());
413
414 for x in 0..100 {
415 eval.evaluate_with_stack(&[x as f64], &mut stack, &mut results)
416 .unwrap();
417 assert!((results[0] - (x as f64 + 1.0)).abs() < 1e-10);
418 }
419 assert!(stack.capacity() >= eval.stack_size());
421 }
422
423 #[test]
424 fn evaluate_with_stack_wrong_arity() {
425 let eval = make_simple_evaluator();
426 let mut stack = Vec::new();
427 let mut results = Vec::new();
428 assert!(
429 eval.evaluate_with_stack(&[1.0, 2.0], &mut stack, &mut results)
430 .is_err()
431 );
432 }
433
434 #[test]
435 fn result_and_stack_getters() {
436 let eval = make_simple_evaluator();
437 assert_eq!(eval.result_count(), 1);
438 assert_eq!(eval.stack_size(), 3);
439 }
440}