Skip to main content

ocas_eval/
compile.rs

1//! AST-to-instruction compiler.
2//!
3//! Transforms an [`EvalTree`](crate::EvalTree) into an
4//! [`ExpressionEvaluator`](crate::ExpressionEvaluator) by generating
5//! a sequence of [`Instr`](crate::Instr)s.
6//!
7//! The compiler walks the tree in post-order, assigning stack slots
8//! and emitting instructions for each node.
9
10use std::collections::HashSet;
11
12use ocas_atom::Atom;
13use ocas_core::FastHashMap as HashMap;
14
15use crate::domain::{EvaluationDomain, PowfExtension};
16use crate::error::{EvaluationError, Result};
17use crate::evaluator::ExpressionEvaluator;
18use crate::function_map::FunctionMap;
19use crate::instruction::Instr;
20use crate::optimize;
21use crate::tree::EvalTree;
22
23/// Compile an [`Atom`] into an [`ExpressionEvaluator`].
24pub fn compile_atom<T: EvaluationDomain + PowfExtension>(
25    atom: Atom<'_>,
26) -> Result<ExpressionEvaluator<T>> {
27    compile_atom_with(atom, None)
28}
29
30/// Compile an [`Atom`] into an [`ExpressionEvaluator`] with a function map.
31pub fn compile_atom_with<T: EvaluationDomain + PowfExtension>(
32    atom: Atom<'_>,
33    function_map: Option<FunctionMap<T>>,
34) -> Result<ExpressionEvaluator<T>> {
35    let tree = EvalTree::from_atom(atom);
36    compile_tree_with(&tree, function_map)
37}
38
39/// Compile an [`EvalTree`] into an [`ExpressionEvaluator`].
40#[allow(dead_code)]
41pub fn compile_tree<T: EvaluationDomain + PowfExtension>(
42    tree: &EvalTree,
43) -> Result<ExpressionEvaluator<T>> {
44    compile_tree_with(tree, None)
45}
46
47/// Compile an [`EvalTree`] with an optional function map.
48pub fn compile_tree_with<T: EvaluationDomain + PowfExtension>(
49    tree: &EvalTree,
50    function_map: Option<FunctionMap<T>>,
51) -> Result<ExpressionEvaluator<T>> {
52    compile_trees(&[tree], function_map)
53}
54
55/// Compile multiple [`Atom`]s into a single multi-output
56/// [`ExpressionEvaluator`], sharing constants, CSE, and stack slots
57/// across all outputs.
58pub fn compile_atoms_multi<T: EvaluationDomain + PowfExtension>(
59    atoms: &[Atom<'_>],
60) -> Result<ExpressionEvaluator<T>> {
61    compile_atoms_multi_with(atoms, None)
62}
63
64/// Compile multiple [`Atom`]s with a function map into a single
65/// multi-output [`ExpressionEvaluator`].
66pub fn compile_atoms_multi_with<T: EvaluationDomain + PowfExtension>(
67    atoms: &[Atom<'_>],
68    function_map: Option<FunctionMap<T>>,
69) -> Result<ExpressionEvaluator<T>> {
70    let trees: Vec<EvalTree> = atoms.iter().map(|a| EvalTree::from_atom(*a)).collect();
71    let refs: Vec<&EvalTree> = trees.iter().collect();
72    compile_trees(&refs, function_map)
73}
74
75/// Compile multiple [`EvalTree`]s into a single multi-output
76/// [`ExpressionEvaluator`].
77pub fn compile_trees_multi<T: EvaluationDomain + PowfExtension>(
78    trees: &[&EvalTree],
79) -> Result<ExpressionEvaluator<T>> {
80    compile_trees(trees, None)
81}
82
83/// Shared compiler core: compile a set of trees into one evaluator
84/// with multiple result slots.
85fn compile_trees<T: EvaluationDomain + PowfExtension>(
86    trees: &[&EvalTree],
87    function_map: Option<FunctionMap<T>>,
88) -> Result<ExpressionEvaluator<T>> {
89    // Constant folding and algebraic simplification at the tree level.
90    let folded: Vec<EvalTree> = trees.iter().map(|t| t.fold_constants()).collect();
91
92    // Pass 1: collect all variable names and count constants
93    let mut var_names = HashSet::new();
94    let mut const_count = 0usize;
95    for tree in &folded {
96        scan_tree(tree, &mut var_names, &mut const_count);
97    }
98    let param_count = var_names.len();
99
100    // Assign parameter slots: sort for deterministic ordering
101    let mut sorted_vars: Vec<String> = var_names.into_iter().collect();
102    sorted_vars.sort();
103    let var_to_param: HashMap<String, usize> = sorted_vars
104        .iter()
105        .enumerate()
106        .map(|(i, v)| (v.clone(), i))
107        .collect();
108
109    // Pass 2: compile all trees with one shared context so that CSE
110    // deduplicates subexpressions across outputs.
111    let temp_base = param_count + const_count;
112    let (instructions, next_temp, constants, result_slots) = {
113        let mut ctx =
114            CompileContext::<T>::new(param_count, temp_base, var_to_param, function_map.as_ref());
115        let mut result_slots = Vec::with_capacity(folded.len());
116        for tree in &folded {
117            result_slots.push(ctx.compile_node(tree)?);
118        }
119        (ctx.instructions, ctx.next_temp, ctx.constants, result_slots)
120    };
121
122    let actual_const_count = constants.len();
123    let (instructions, temp_count, result_indices) =
124        optimize::optimize(instructions, temp_base, next_temp, &result_slots);
125    let stack_size = temp_base + temp_count;
126
127    match function_map {
128        Some(fm) => Ok(ExpressionEvaluator::new_with_functions(
129            instructions,
130            param_count,
131            actual_const_count,
132            stack_size,
133            result_indices,
134            constants,
135            fm,
136        )),
137        None => Ok(ExpressionEvaluator::new(
138            instructions,
139            param_count,
140            actual_const_count,
141            stack_size,
142            result_indices,
143            constants,
144        )),
145    }
146}
147
148/// Pre-scan tree to count variables and constants.
149fn scan_tree(tree: &EvalTree, vars: &mut HashSet<String>, const_count: &mut usize) {
150    match tree {
151        EvalTree::Num(_) => {
152            *const_count += 1;
153        }
154        EvalTree::Var(name) => {
155            vars.insert(name.clone());
156        }
157        EvalTree::Add(terms) | EvalTree::Mul(terms) => {
158            for t in terms {
159                scan_tree(t, vars, const_count);
160            }
161        }
162        EvalTree::Pow(base, exp) => {
163            scan_tree(base, vars, const_count);
164            scan_tree(exp, vars, const_count);
165        }
166        EvalTree::Fun(_, args) => {
167            for a in args {
168                scan_tree(a, vars, const_count);
169            }
170        }
171    }
172}
173
174struct CompileContext<'a, T: EvaluationDomain> {
175    instructions: Vec<Instr>,
176    /// Next available temp slot index. Temps start at `temp_base` in the actual stack.
177    next_temp: usize,
178    /// Parameter slots occupy stack[0..param_count].
179    param_count: usize,
180    /// Base index for temp slots in the actual stack (= param_count + estimated_const_count).
181    temp_base: usize,
182    /// variable name → stack slot index (0..param_count-1)
183    variables: HashMap<String, usize>,
184    /// constant values in order
185    constants: Vec<T>,
186    /// Optional function map for resolving external functions
187    function_map: Option<&'a FunctionMap<T>>,
188}
189
190impl<'a, T: EvaluationDomain> CompileContext<'a, T> {
191    fn new(
192        param_count: usize,
193        temp_base: usize,
194        variables: HashMap<String, usize>,
195        function_map: Option<&'a FunctionMap<T>>,
196    ) -> Self {
197        Self {
198            instructions: Vec::new(),
199            next_temp: 0,
200            param_count,
201            temp_base,
202            variables,
203            constants: Vec::new(),
204            function_map,
205        }
206    }
207
208    fn alloc_temp(&mut self) -> usize {
209        let slot = self.next_temp + self.temp_base;
210        self.next_temp += 1;
211        slot
212    }
213
214    fn param_slot(&self, name: &str) -> usize {
215        self.variables[name]
216    }
217
218    fn const_slot(&mut self, value: T) -> usize {
219        let idx = self.constants.len();
220        self.constants.push(value);
221        self.param_count + idx
222    }
223
224    fn compile_node(&mut self, node: &EvalTree) -> Result<usize> {
225        match node {
226            EvalTree::Num(n) => {
227                let dst = self.alloc_temp();
228                let const_slot = self.const_slot(T::from_f64(*n));
229                self.instructions.push(Instr::Copy {
230                    dst,
231                    src: const_slot,
232                });
233                Ok(dst)
234            }
235            EvalTree::Var(name) => {
236                let dst = self.alloc_temp();
237                let param_slot = self.param_slot(name);
238                self.instructions.push(Instr::Copy {
239                    dst,
240                    src: param_slot,
241                });
242                Ok(dst)
243            }
244            EvalTree::Add(terms) => {
245                let dst = self.alloc_temp();
246                let mut srcs = Vec::with_capacity(terms.len());
247                for term in terms {
248                    srcs.push(self.compile_node(term)?);
249                }
250                self.instructions.push(Instr::Add { dst, srcs });
251                Ok(dst)
252            }
253            EvalTree::Mul(factors) => {
254                let dst = self.alloc_temp();
255                let mut srcs = Vec::with_capacity(factors.len());
256                for factor in factors {
257                    srcs.push(self.compile_node(factor)?);
258                }
259                self.instructions.push(Instr::Mul { dst, srcs });
260                Ok(dst)
261            }
262            EvalTree::Pow(base, exp) => {
263                let base_slot = self.compile_node(base)?;
264                let dst = self.alloc_temp();
265                if let EvalTree::Num(n) = exp.as_ref()
266                    && n.fract() == 0.0
267                    && *n >= i64::MIN as f64
268                    && *n <= i64::MAX as f64
269                {
270                    self.instructions.push(Instr::Pow {
271                        dst,
272                        base: base_slot,
273                        exp: *n as i64,
274                    });
275                    return Ok(dst);
276                }
277                let exp_slot = self.compile_node(exp)?;
278                self.instructions.push(Instr::Powf {
279                    dst,
280                    base: base_slot,
281                    exp: exp_slot,
282                });
283                Ok(dst)
284            }
285            EvalTree::Fun(name, args) => {
286                if is_builtin(name) && args.len() == 1 {
287                    let arg_slot = self.compile_node(&args[0])?;
288                    let dst = self.alloc_temp();
289                    let op = crate::instruction::BuiltinOp::from_name(name)
290                        .expect("is_builtin guarantees known name");
291                    self.instructions.push(Instr::BuiltinOp {
292                        dst,
293                        op,
294                        src: arg_slot,
295                    });
296                    Ok(dst)
297                } else if let Some(fm) = self.function_map {
298                    // Look up in function map
299                    if let Some(_entry) = fm.resolve(name) {
300                        let mut srcs = Vec::with_capacity(args.len());
301                        for arg in args {
302                            srcs.push(self.compile_node(arg)?);
303                        }
304                        let dst = self.alloc_temp();
305                        // Find the function index in the map
306                        let fn_idx =
307                            fm.index_of(name)
308                                .ok_or_else(|| EvaluationError::FunctionNotFound {
309                                    name: name.clone(),
310                                })?;
311                        self.instructions
312                            .push(Instr::ExternalFun { dst, fn_idx, srcs });
313                        Ok(dst)
314                    } else {
315                        Err(EvaluationError::FunctionNotFound { name: name.clone() })
316                    }
317                } else {
318                    Err(EvaluationError::FunctionNotFound { name: name.clone() })
319                }
320            }
321        }
322    }
323}
324
325fn is_builtin(name: &str) -> bool {
326    matches!(
327        name.to_lowercase().as_str(),
328        "sin" | "cos" | "tan" | "sec" | "csc" | "cot" | "exp" | "log" | "sqrt" | "abs"
329    )
330}
331
332// ---------------------------------------------------------------------------
333// ExpressionEvaluator::compile
334// ---------------------------------------------------------------------------
335
336impl<T: EvaluationDomain + PowfExtension> ExpressionEvaluator<T> {
337    /// Compile an [`Atom`] into an executable evaluator.
338    pub fn compile(atom: Atom<'_>) -> Result<Self> {
339        compile_atom(atom)
340    }
341
342    /// Compile an [`Atom`] with a [`FunctionMap`] for user-defined functions.
343    pub fn compile_with(atom: Atom<'_>, map: FunctionMap<T>) -> Result<Self> {
344        compile_atom_with(atom, Some(map))
345    }
346
347    /// Compile multiple [`Atom`]s into one multi-output evaluator,
348    /// sharing common subexpressions across all outputs.
349    pub fn compile_multi(atoms: &[Atom<'_>]) -> Result<Self> {
350        compile_atoms_multi(atoms)
351    }
352
353    /// Compile multiple [`Atom`]s with a [`FunctionMap`] into one
354    /// multi-output evaluator.
355    pub fn compile_multi_with(atoms: &[Atom<'_>], map: FunctionMap<T>) -> Result<Self> {
356        compile_atoms_multi_with(atoms, Some(map))
357    }
358}
359
360#[cfg(test)]
361mod tests {
362    use super::*;
363    use ocas_atom::AtomArena;
364    use ocas_core::arena::Arena;
365
366    #[test]
367    fn compile_constant() {
368        let arena = Arena::new();
369        let ctx = AtomArena::new(&arena);
370        let expr = ctx.num(42);
371        let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
372        let result = eval.evaluate(&[]).unwrap();
373        assert!((result[0] - 42.0).abs() < 1e-10);
374    }
375
376    #[test]
377    fn compile_single_var() {
378        let arena = Arena::new();
379        let ctx = AtomArena::new(&arena);
380        let expr = ctx.var("x");
381        let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
382        assert_eq!(eval.param_count(), 1);
383        let result = eval.evaluate(&[7.0]).unwrap();
384        assert!((result[0] - 7.0).abs() < 1e-10);
385    }
386
387    #[test]
388    fn compile_add_two_vars() {
389        let arena = Arena::new();
390        let ctx = AtomArena::new(&arena);
391        let expr = ctx.add(&[ctx.var("x"), ctx.var("y")]);
392        let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
393        let result = eval.evaluate(&[2.0, 3.0]).unwrap();
394        assert!((result[0] - 5.0).abs() < 1e-10);
395    }
396
397    #[test]
398    fn compile_mul_var_const() {
399        let arena = Arena::new();
400        let ctx = AtomArena::new(&arena);
401        let expr = ctx.mul(&[ctx.var("x"), ctx.num(3)]);
402        let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
403        let result = eval.evaluate(&[4.0]).unwrap();
404        assert!((result[0] - 12.0).abs() < 1e-10);
405    }
406
407    #[test]
408    fn compile_pow_integer_exp() {
409        let arena = Arena::new();
410        let ctx = AtomArena::new(&arena);
411        let expr = ctx.pow(ctx.var("x"), ctx.num(3));
412        let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
413        let result = eval.evaluate(&[2.0]).unwrap();
414        assert!((result[0] - 8.0).abs() < 1e-10);
415    }
416
417    #[test]
418    fn compile_sin() {
419        let arena = Arena::new();
420        let ctx = AtomArena::new(&arena);
421        let expr = ctx.fun("sin", &[ctx.var("x")]);
422        let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
423        let result = eval.evaluate(&[std::f64::consts::FRAC_PI_2]).unwrap();
424        assert!((result[0] - 1.0).abs() < 1e-10);
425    }
426
427    #[test]
428    fn compile_cos() {
429        let arena = Arena::new();
430        let ctx = AtomArena::new(&arena);
431        let expr = ctx.fun("cos", &[ctx.var("x")]);
432        let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
433        let result = eval.evaluate(&[std::f64::consts::PI]).unwrap();
434        assert!((result[0] + 1.0).abs() < 1e-10);
435    }
436
437    #[test]
438    fn compile_exp_log_roundtrip() {
439        let arena = Arena::new();
440        let ctx = AtomArena::new(&arena);
441        let exp_x = ctx.fun("exp", &[ctx.var("x")]);
442        let expr = ctx.fun("log", &[exp_x]);
443        let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
444        let result = eval.evaluate(&[2.0]).unwrap();
445        assert!((result[0] - 2.0).abs() < 1e-10);
446    }
447
448    #[test]
449    fn compile_nested_expression() {
450        // (x + 1) * (x - 1) = x^2 - 1
451        let arena = Arena::new();
452        let ctx = AtomArena::new(&arena);
453        let x = ctx.var("x");
454        let x_plus_1 = ctx.add(&[x, ctx.num(1)]);
455        let x_minus_1 = ctx.add(&[x, ctx.num(-1)]);
456        let expr = ctx.mul(&[x_plus_1, x_minus_1]);
457        let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
458        let result = eval.evaluate(&[3.0]).unwrap();
459        assert!((result[0] - 8.0).abs() < 1e-10);
460    }
461
462    #[test]
463    fn compile_sqrt() {
464        let arena = Arena::new();
465        let ctx = AtomArena::new(&arena);
466        let expr = ctx.fun("sqrt", &[ctx.num(16)]);
467        let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
468        let result = eval.evaluate(&[]).unwrap();
469        assert!((result[0] - 4.0).abs() < 1e-10);
470    }
471
472    #[test]
473    fn compile_zero_params() {
474        let arena = Arena::new();
475        let ctx = AtomArena::new(&arena);
476        let expr = ctx.fun("sin", &[ctx.num(1)]); // sin(1 rad)
477        let eval: ExpressionEvaluator<f64> = ExpressionEvaluator::compile(expr).unwrap();
478        assert_eq!(eval.param_count(), 0);
479        let result = eval.evaluate(&[]).unwrap();
480        assert!((result[0] - 1.0f64.sin()).abs() < 1e-10);
481    }
482
483    #[test]
484    fn compile_with_external_function() {
485        let arena = Arena::new();
486        let ctx = AtomArena::new(&arena);
487        let expr = ctx.fun("square", &[ctx.var("x")]);
488
489        let mut map = FunctionMap::<f64>::new();
490        map.register("square", 1, Box::new(|args| args[0] * args[0]));
491
492        let eval = ExpressionEvaluator::compile_with(expr, map).unwrap();
493        let result = eval.evaluate(&[3.0]).unwrap();
494        assert!((result[0] - 9.0).abs() < 1e-10);
495    }
496
497    #[test]
498    fn compile_external_function_not_registered() {
499        let arena = Arena::new();
500        let ctx = AtomArena::new(&arena);
501        let expr = ctx.fun("missing_fn", &[ctx.var("x")]);
502
503        let result: Result<ExpressionEvaluator<f64>> = ExpressionEvaluator::compile(expr);
504        assert!(result.is_err());
505    }
506
507    #[test]
508    fn compile_with_case_insensitive_external() {
509        let arena = Arena::new();
510        let ctx = AtomArena::new(&arena);
511        let expr = ctx.fun("Square", &[ctx.num(4)]);
512
513        let mut map = FunctionMap::<f64>::new();
514        map.register("square", 1, Box::new(|args| args[0] * args[0]));
515
516        let eval = ExpressionEvaluator::compile_with(expr, map).unwrap();
517        let result = eval.evaluate(&[]).unwrap();
518        assert!((result[0] - 16.0).abs() < 1e-10);
519    }
520
521    #[test]
522    fn compile_multi_two_outputs() {
523        let arena = Arena::new();
524        let ctx = AtomArena::new(&arena);
525        let sum = ctx.add(&[ctx.var("x"), ctx.var("y")]);
526        let prod = ctx.mul(&[ctx.var("x"), ctx.var("y")]);
527        let eval: ExpressionEvaluator<f64> =
528            ExpressionEvaluator::compile_multi(&[sum, prod]).unwrap();
529        assert_eq!(eval.result_count(), 2);
530        assert_eq!(eval.param_count(), 2);
531        let result = eval.evaluate(&[2.0, 3.0]).unwrap();
532        assert!((result[0] - 5.0).abs() < 1e-10);
533        assert!((result[1] - 6.0).abs() < 1e-10);
534    }
535
536    #[test]
537    fn compile_multi_shared_subexpression() {
538        // outputs: sin(x) + 1, sin(x) * 2 — sin(x) shared via CSE
539        let arena = Arena::new();
540        let ctx = AtomArena::new(&arena);
541        let sin_x = ctx.fun("sin", &[ctx.var("x")]);
542        let out0 = ctx.add(&[sin_x, ctx.num(1)]);
543        let out1 = ctx.mul(&[sin_x, ctx.num(2)]);
544        let eval: ExpressionEvaluator<f64> =
545            ExpressionEvaluator::compile_multi(&[out0, out1]).unwrap();
546        let result = eval.evaluate(&[std::f64::consts::FRAC_PI_2]).unwrap();
547        assert!((result[0] - 2.0).abs() < 1e-10);
548        assert!((result[1] - 2.0).abs() < 1e-10);
549    }
550
551    #[test]
552    fn compile_multi_constant_folding() {
553        // (2 + 3) * x and x^1 — folding reduces both
554        let arena = Arena::new();
555        let ctx = AtomArena::new(&arena);
556        let five_x = ctx.mul(&[ctx.add(&[ctx.num(2), ctx.num(3)]), ctx.var("x")]);
557        let x_pow_1 = ctx.pow(ctx.var("x"), ctx.num(1));
558        let eval: ExpressionEvaluator<f64> =
559            ExpressionEvaluator::compile_multi(&[five_x, x_pow_1]).unwrap();
560        let result = eval.evaluate(&[4.0]).unwrap();
561        assert!((result[0] - 20.0).abs() < 1e-10);
562        assert!((result[1] - 4.0).abs() < 1e-10);
563    }
564
565    #[test]
566    fn compile_multi_with_external_function() {
567        let arena = Arena::new();
568        let ctx = AtomArena::new(&arena);
569        let sq = ctx.fun("square", &[ctx.var("x")]);
570        let cube_arg = ctx.fun("square", &[ctx.var("x")]);
571        let out1 = ctx.mul(&[cube_arg, ctx.var("x")]);
572
573        let mut map = FunctionMap::<f64>::new();
574        map.register("square", 1, Box::new(|args| args[0] * args[0]));
575
576        let eval = ExpressionEvaluator::compile_multi_with(&[sq, out1], map).unwrap();
577        let result = eval.evaluate(&[3.0]).unwrap();
578        assert!((result[0] - 9.0).abs() < 1e-10);
579        assert!((result[1] - 27.0).abs() < 1e-10);
580    }
581}