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 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 for (i, p) in params.iter().enumerate() {
115 stack[i] = p.clone();
116 }
117
118 for (i, c) in self.constants.iter().enumerate() {
120 stack[self.param_count + i] = c.clone();
121 }
122
123 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 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#[cfg(feature = "simd")]
193impl ExpressionEvaluator<f64> {
194 pub fn compile_vector_evaluator(&self) -> Result<crate::simd::VectorEvaluator> {
204 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 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 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 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 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 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}