qualia_core_db/specialized_libs/
symbolic_trig.rs1use super::symbolic_algebra::{add, c, cos, div, mul, neg, pow, simplify, sin, sub, tan, Expr};
15
16pub fn simplify_trig(expr: &Expr) -> Expr {
18 let mut cur = simplify(expr);
19 for _ in 0..16 {
20 let next = simplify(&rewrite(&cur));
21 if next == cur {
22 break;
23 }
24 cur = next;
25 }
26 cur
27}
28
29fn as_scaled_square(e: &Expr) -> Option<(Expr, bool, Expr)> {
31 let classify = |b: &Expr| -> Option<(bool, Expr)> {
32 match b {
33 Expr::Sin(u) => Some((true, (**u).clone())),
34 Expr::Cos(u) => Some((false, (**u).clone())),
35 _ => None,
36 }
37 };
38 match e {
39 Expr::Pow(base, 2) => classify(base).map(|(is_sin, u)| (c(1.0), is_sin, u)),
40 Expr::Mul(a, b) => {
41 if let (Expr::Const(_), Expr::Pow(base, 2)) = (&**a, &**b) {
42 classify(base).map(|(is_sin, u)| ((**a).clone(), is_sin, u))
43 } else if let (Expr::Pow(base, 2), Expr::Const(_)) = (&**a, &**b) {
44 classify(base).map(|(is_sin, u)| ((**b).clone(), is_sin, u))
45 } else {
46 None
47 }
48 }
49 _ => None,
50 }
51}
52
53fn pythagorean(a: &Expr, b: &Expr) -> Option<Expr> {
55 let (ka, sin_a, ua) = as_scaled_square(a)?;
56 let (kb, sin_b, ub) = as_scaled_square(b)?;
57 if ka == kb && ua == ub && sin_a != sin_b {
58 Some(ka)
59 } else {
60 None
61 }
62}
63
64fn rewrite(e: &Expr) -> Expr {
65 let e = match e {
67 Expr::Add(a, b) => add(rewrite(a), rewrite(b)),
68 Expr::Sub(a, b) => sub(rewrite(a), rewrite(b)),
69 Expr::Mul(a, b) => mul(rewrite(a), rewrite(b)),
70 Expr::Div(a, b) => div(rewrite(a), rewrite(b)),
71 Expr::Pow(a, n) => pow(rewrite(a), *n),
72 Expr::Neg(a) => neg(rewrite(a)),
73 Expr::Sqrt(a) => Expr::Sqrt(Box::new(rewrite(a))),
74 Expr::Exp(a) => Expr::Exp(Box::new(rewrite(a))),
75 Expr::Ln(a) => Expr::Ln(Box::new(rewrite(a))),
76 Expr::Sin(a) => sin(rewrite(a)),
77 Expr::Cos(a) => cos(rewrite(a)),
78 Expr::Tan(a) => tan(rewrite(a)),
79 Expr::Const(_) | Expr::Var(_) => e.clone(),
80 };
81
82 match &e {
83 Expr::Sin(u) => {
85 if let Expr::Neg(inner) = &**u {
86 return neg(sin((**inner).clone()));
87 }
88 e
89 }
90 Expr::Cos(u) => {
91 if let Expr::Neg(inner) = &**u {
92 return cos((**inner).clone());
93 }
94 e
95 }
96 Expr::Tan(u) => {
97 if let Expr::Neg(inner) = &**u {
98 return neg(tan((**inner).clone()));
99 }
100 e
101 }
102 Expr::Div(n, d) => {
104 if let (Expr::Sin(un), Expr::Cos(ud)) = (&**n, &**d) {
105 if un == ud {
106 return tan((**un).clone());
107 }
108 }
109 e
110 }
111 Expr::Add(a, b) => {
113 if let Some(r) = pythagorean(a, b).or_else(|| pythagorean(b, a)) {
114 return r;
115 }
116 e
117 }
118 Expr::Sub(a, b) => {
120 if let Expr::Const(one) = &**a {
121 if *one == 1.0 {
122 if let Expr::Pow(base, 2) = &**b {
123 match &**base {
124 Expr::Sin(u) => return pow(cos((**u).clone()), 2),
125 Expr::Cos(u) => return pow(sin((**u).clone()), 2),
126 _ => {}
127 }
128 }
129 }
130 }
131 e
132 }
133 _ => e,
134 }
135}
136
137#[cfg(test)]
138mod tests {
139 use super::super::symbolic_algebra::{var, Expr};
140 use super::*;
141 use std::collections::HashMap;
142
143 fn val(e: &Expr, x: f64) -> f64 {
144 let mut env = HashMap::new();
145 env.insert("x".to_string(), x);
146 e.eval(&env).unwrap()
147 }
148
149 #[test]
150 fn pythagorean_identity_collapses() {
151 let e = add(pow(sin(var("x")), 2), pow(cos(var("x")), 2));
153 assert_eq!(simplify_trig(&e), c(1.0));
154 let e2 = add(pow(cos(var("x")), 2), pow(sin(var("x")), 2));
156 assert_eq!(simplify_trig(&e2), c(1.0));
157 let e3 = add(
159 mul(c(3.0), pow(sin(var("x")), 2)),
160 mul(c(3.0), pow(cos(var("x")), 2)),
161 );
162 assert_eq!(simplify_trig(&e3), c(3.0));
163 }
164
165 #[test]
166 fn pythagorean_in_a_larger_sum() {
167 let e = add(add(pow(sin(var("x")), 2), pow(cos(var("x")), 2)), var("x"));
169 let s = simplify_trig(&e);
170 for &x in &[0.3, 1.2, 2.7] {
171 assert!((val(&s, x) - (1.0 + x)).abs() < 1e-9);
172 }
173 }
174
175 #[test]
176 fn one_minus_square() {
177 let s = simplify_trig(&sub(c(1.0), pow(sin(var("x")), 2)));
179 for &x in &[0.3, 1.2, 2.7] {
180 assert!((val(&s, x) - x.cos().powi(2)).abs() < 1e-9);
181 }
182 }
183
184 #[test]
185 fn parity_and_quotient() {
186 assert_eq!(simplify_trig(&cos(neg(var("x")))), cos(var("x")));
188 assert_eq!(simplify_trig(&sin(neg(var("x")))), neg(sin(var("x"))));
189 assert_eq!(
191 simplify_trig(&div(sin(var("x")), cos(var("x")))),
192 tan(var("x"))
193 );
194 }
195}