1use rustc_hash::{FxBuildHasher, FxHashMap};
2use smallvec::smallvec;
3
4use crate::arena::{ExprArena, ExprId, ExprNode, VarId};
5
6#[derive(Clone, Debug, Default)]
8pub struct LinearTerms {
9 pub coeffs: Vec<(VarId, f64)>,
10 pub constant: f64,
11}
12
13struct CoeffAccum {
16 coeffs: Vec<(VarId, f64)>,
17 slot: FxHashMap<VarId, usize>,
18}
19
20impl CoeffAccum {
21 fn with_capacity(n: usize) -> Self {
22 Self {
23 coeffs: Vec::with_capacity(n),
24 slot: FxHashMap::with_capacity_and_hasher(n, FxBuildHasher),
25 }
26 }
27
28 fn add(&mut self, v: VarId, c: f64) {
31 if let Some(&i) = self.slot.get(&v) {
32 self.coeffs[i].1 += c;
33 } else {
34 self.slot.insert(v, self.coeffs.len());
35 self.coeffs.push((v, c));
36 }
37 }
38
39 fn extend(&mut self, terms: impl IntoIterator<Item = (VarId, f64)>) {
40 for (v, c) in terms {
41 self.add(v, c);
42 }
43 }
44
45 fn into_coeffs(self) -> Vec<(VarId, f64)> {
46 self.coeffs
47 }
48}
49
50fn as_linear(arena: &ExprArena, id: ExprId, resolve_params: bool) -> Option<LinearTerms> {
56 match arena.get(id) {
57 ExprNode::Const(c) => Some(LinearTerms { coeffs: Vec::new(), constant: *c }),
58 ExprNode::Param(p) if resolve_params => {
59 Some(LinearTerms { coeffs: Vec::new(), constant: arena.param_value(*p) })
60 }
61 ExprNode::Var(v) => Some(LinearTerms { coeffs: vec![(*v, 1.0)], constant: 0.0 }),
62 ExprNode::Linear { coeffs, constant } => {
63 Some(LinearTerms { coeffs: coeffs.clone(), constant: *constant })
64 }
65 ExprNode::Neg(inner) => {
66 let inner = *inner;
67 as_linear(arena, inner, resolve_params).map(|mut t| {
68 t.coeffs.iter_mut().for_each(|(_, c)| *c = -*c);
69 t.constant = -t.constant;
70 t
71 })
72 }
73 ExprNode::Add(children) => {
74 let children: smallvec::SmallVec<[ExprId; 4]> = children.iter().copied().collect();
75 let mut acc = CoeffAccum::with_capacity(children.len() * 4);
76 let mut constant = 0.0;
77 for child in children {
78 let t = as_linear(arena, child, resolve_params)?;
79 acc.extend(t.coeffs);
80 constant += t.constant;
81 }
82 Some(LinearTerms { coeffs: acc.into_coeffs(), constant })
83 }
84 ExprNode::Mul(children) => {
85 let children: smallvec::SmallVec<[ExprId; 4]> = children.iter().copied().collect();
87 let mut scalar = 1.0;
88 let mut linear: Option<LinearTerms> = None;
89 for child in children {
90 match arena.get(child) {
91 ExprNode::Const(c) => scalar *= c,
92 ExprNode::Param(p) if resolve_params => scalar *= arena.param_value(*p),
93 _ if linear.is_none() => {
94 linear = Some(as_linear(arena, child, resolve_params)?);
95 }
96 _ => return None,
97 }
98 }
99 Some(match linear {
100 None => LinearTerms { coeffs: Vec::new(), constant: scalar },
101 Some(mut t) => {
102 t.coeffs.iter_mut().for_each(|(_, c)| *c *= scalar);
103 t.constant *= scalar;
104 t
105 }
106 })
107 }
108 _ => None,
109 }
110}
111
112fn push_linear(arena: &mut ExprArena, mut t: LinearTerms) -> ExprId {
114 t.coeffs.retain(|(_, c)| *c != 0.0);
115 arena.push(ExprNode::Linear { coeffs: t.coeffs, constant: t.constant })
116}
117
118pub(crate) fn add_into(arena: &mut ExprArena, lhs: ExprId, rhs: ExprId) -> ExprId {
121 if let (Some(lt), Some(rt)) = (as_linear(arena, lhs, false), as_linear(arena, rhs, false)) {
122 let mut acc = CoeffAccum::with_capacity(lt.coeffs.len() + rt.coeffs.len());
123 acc.extend(lt.coeffs);
124 acc.extend(rt.coeffs);
125 return push_linear(
126 arena,
127 LinearTerms { coeffs: acc.into_coeffs(), constant: lt.constant + rt.constant },
128 );
129 }
130 arena.push(ExprNode::Add(smallvec![lhs, rhs]))
131}
132
133pub(crate) fn add_n(arena: &mut ExprArena, ids: &[ExprId]) -> ExprId {
140 match ids {
141 [] => panic!("add_n on an empty term list"),
142 [one] => *one,
143 _ => arena.push(ExprNode::Add(ids.iter().copied().collect())),
144 }
145}
146
147pub(crate) fn sub_into(arena: &mut ExprArena, lhs: ExprId, rhs: ExprId) -> ExprId {
149 let neg = neg_into(arena, rhs);
150 add_into(arena, lhs, neg)
151}
152
153pub(crate) fn mul_into(arena: &mut ExprArena, lhs: ExprId, rhs: ExprId) -> ExprId {
156 if let ExprNode::Const(c) = *arena.get(lhs) {
157 if let Some(mut t) = as_linear(arena, rhs, false) {
158 t.coeffs.iter_mut().for_each(|(_, co)| *co *= c);
159 t.constant *= c;
160 return push_linear(arena, t);
161 }
162 }
163 if let ExprNode::Const(c) = *arena.get(rhs) {
164 if let Some(mut t) = as_linear(arena, lhs, false) {
165 t.coeffs.iter_mut().for_each(|(_, co)| *co *= c);
166 t.constant *= c;
167 return push_linear(arena, t);
168 }
169 }
170 arena.push(ExprNode::Mul(smallvec![lhs, rhs]))
171}
172
173pub(crate) fn div_into(arena: &mut ExprArena, num: ExprId, den: ExprId) -> ExprId {
177 if let ExprNode::Const(c) = *arena.get(den) {
178 if c != 0.0 {
179 if let Some(mut t) = as_linear(arena, num, false) {
180 let inv = 1.0 / c;
181 t.coeffs.iter_mut().for_each(|(_, co)| *co *= inv);
182 t.constant *= inv;
183 return push_linear(arena, t);
184 }
185 let inv = arena.push(ExprNode::Const(1.0 / c));
186 return mul_into(arena, num, inv);
187 }
188 }
189 arena.push(ExprNode::Div(num, den))
190}
191
192pub(crate) fn neg_into(arena: &mut ExprArena, rhs: ExprId) -> ExprId {
194 if let Some(mut t) = as_linear(arena, rhs, false) {
195 t.coeffs.iter_mut().for_each(|(_, c)| *c = -*c);
196 t.constant = -t.constant;
197 return push_linear(arena, t);
198 }
199 arena.push(ExprNode::Neg(rhs))
200}
201
202pub fn extract_linear(arena: &ExprArena, id: ExprId) -> Option<LinearTerms> {
210 as_linear(arena, id, true)
211}
212
213#[derive(Copy, Clone, Debug, PartialEq, Eq)]
217pub struct SignedExpr {
218 pub id: ExprId,
219 pub neg: bool,
220}
221
222pub fn split_linear(arena: &ExprArena, id: ExprId) -> (LinearTerms, Vec<SignedExpr>) {
234 if let Some(lt) = as_linear(arena, id, true) {
235 return (lt, Vec::new());
236 }
237 let mut lin = CoeffAccum::with_capacity(0);
238 let mut constant = 0.0;
239 let mut residual: Vec<SignedExpr> = Vec::new();
240 let mut sign_stack: smallvec::SmallVec<[(ExprId, f64); 8]> = smallvec![(id, 1.0)];
241 while let Some((cur, sign)) = sign_stack.pop() {
242 match arena.get(cur) {
243 ExprNode::Add(children) => {
244 for c in children.iter().copied() {
245 sign_stack.push((c, sign));
246 }
247 }
248 ExprNode::Neg(inner) => sign_stack.push((*inner, -sign)),
249 _ => {
250 if let Some(mut t) = as_linear(arena, cur, true) {
251 if (sign - 1.0).abs() > 0.0 {
252 t.coeffs.iter_mut().for_each(|(_, c)| *c *= sign);
253 t.constant *= sign;
254 }
255 lin.extend(t.coeffs);
256 constant += t.constant;
257 } else {
258 residual.push(SignedExpr { id: cur, neg: sign < 0.0 });
259 }
260 }
261 }
262 }
263 let mut coeffs = lin.into_coeffs();
264 coeffs.retain(|(_, c)| *c != 0.0);
265 (LinearTerms { coeffs, constant }, residual)
266}
267
268pub fn describe_nonlinear_term(
272 arena: &ExprArena,
273 id: ExprId,
274 resolve: &impl Fn(VarId) -> String,
275) -> Option<String> {
276 use crate::render::{PREC_ADD, PREC_UNARY, render_node};
277 let (_, residual) = split_linear(arena, id);
278 residual.first().map(|s| {
279 if s.neg {
280 format!("-{}", render_node(arena, s.id, resolve, PREC_UNARY))
281 } else {
282 render_node(arena, s.id, resolve, PREC_ADD)
283 }
284 })
285}
286
287#[cfg(test)]
288mod tests {
289 use super::*;
290 use crate::arena::{ExprArena, ExprNode, VarId};
291
292 #[test]
293 fn param_times_var_stays_symbolic_until_extracted() {
294 let mut arena = ExprArena::new();
297 let pid = arena.new_param(3.0);
298 let price = arena.param(pid);
299 let xnode = arena.push(ExprNode::Var(VarId(0)));
300 let prod = mul_into(&mut arena, price, xnode);
301 assert!(matches!(arena.get(prod), ExprNode::Mul(_)));
302
303 let terms = extract_linear(&arena, prod).expect("linear");
304 assert_eq!(terms.coeffs, vec![(VarId(0), 3.0)]);
305 assert!(terms.constant.abs() < f64::EPSILON);
306 }
307
308 #[test]
309 fn rebinding_param_updates_extracted_coeff() {
310 let mut arena = ExprArena::new();
311 let pid = arena.new_param(3.0);
312 let price = arena.param(pid);
313 let xnode = arena.push(ExprNode::Var(VarId(0)));
314 let prod = mul_into(&mut arena, price, xnode);
315
316 arena.set_param_value(pid, 10.0);
317 let terms = extract_linear(&arena, prod).expect("linear");
318 assert_eq!(terms.coeffs, vec![(VarId(0), 10.0)]);
319 }
320
321 #[test]
322 fn param_plus_var_resolves_constant() {
323 let mut arena = ExprArena::new();
324 let pid = arena.new_param(5.0);
325 let price = arena.param(pid);
326 let xnode = arena.push(ExprNode::Var(VarId(0)));
327 let sum = add_into(&mut arena, price, xnode);
328 let terms = extract_linear(&arena, sum).expect("linear");
329 assert_eq!(terms.coeffs, vec![(VarId(0), 1.0)]);
330 assert!((terms.constant - 5.0).abs() < f64::EPSILON);
331 }
332
333 #[test]
334 fn add_extraction_is_first_seen_ordered_and_merges() {
335 let mut arena = ExprArena::new();
338 let z = arena.push(ExprNode::Var(VarId(2)));
339 let x = arena.push(ExprNode::Var(VarId(0)));
340 let y = arena.push(ExprNode::Var(VarId(1)));
341 let sum = arena.push(ExprNode::Add(smallvec::smallvec![z, x, y, x]));
342
343 let terms = extract_linear(&arena, sum).expect("linear");
344 assert_eq!(terms.coeffs, vec![(VarId(2), 1.0), (VarId(0), 2.0), (VarId(1), 1.0)]);
345 assert!(terms.constant.abs() < f64::EPSILON);
346 assert_eq!(extract_linear(&arena, sum).unwrap().coeffs, terms.coeffs);
347 }
348
349 #[test]
350 fn wide_sum_merges_repeated_vars_in_order() {
351 let mut arena = ExprArena::new();
352 let n = 50u32;
353 let mut ids = Vec::new();
354 for _ in 0..3 {
355 for v in 0..n {
356 ids.push(arena.push(ExprNode::Var(VarId(v))));
357 }
358 }
359 let sum = arena.push(ExprNode::Add(ids.into_iter().collect()));
360 let terms = extract_linear(&arena, sum).expect("linear");
361 let expected: Vec<(VarId, f64)> = (0..n).map(|v| (VarId(v), 3.0)).collect();
362 assert_eq!(terms.coeffs, expected);
363 }
364
365 fn names(v: VarId) -> String {
366 match v.0 {
367 0 => "x".to_string(),
368 1 => "y".to_string(),
369 n => format!("v{n}"),
370 }
371 }
372
373 #[test]
374 fn describe_renders_the_first_nonlinear_summand() {
375 let mut arena = ExprArena::new();
376 let x = arena.push(ExprNode::Var(VarId(0)));
377 let y = arena.push(ExprNode::Var(VarId(1)));
378
379 let prod = arena.push(ExprNode::Mul(smallvec::smallvec![x, y]));
380 assert_eq!(describe_nonlinear_term(&arena, prod, &names).as_deref(), Some("x * y"));
381
382 let two = arena.constant(2.0);
383 let pow = arena.push(ExprNode::Pow(x, two));
384 assert_eq!(describe_nonlinear_term(&arena, pow, &names).as_deref(), Some("x^2"));
385
386 let s = arena.push(ExprNode::Sin(x));
387 assert_eq!(describe_nonlinear_term(&arena, s, &names).as_deref(), Some("sin(x)"));
388
389 let div = arena.push(ExprNode::Div(x, y));
390 assert_eq!(describe_nonlinear_term(&arena, div, &names).as_deref(), Some("x / y"));
391 }
392
393 #[test]
394 fn describe_isolates_the_nonlinear_part_of_a_mixed_expression() {
395 let mut arena = ExprArena::new();
396 let x = arena.push(ExprNode::Var(VarId(0)));
397 let y = arena.push(ExprNode::Var(VarId(1)));
398 let z = arena.push(ExprNode::Var(VarId(2)));
399 let two = arena.constant(2.0);
400 let two_z = arena.push(ExprNode::Mul(smallvec::smallvec![two, z]));
401 let prod = arena.push(ExprNode::Mul(smallvec::smallvec![x, y]));
402 let sum = arena.push(ExprNode::Add(smallvec::smallvec![two_z, prod]));
403 assert_eq!(describe_nonlinear_term(&arena, sum, &names).as_deref(), Some("x * y"));
404 }
405
406 #[test]
407 fn describe_returns_none_for_affine() {
408 let mut arena = ExprArena::new();
409 let x = arena.push(ExprNode::Var(VarId(0)));
410 let three = arena.constant(3.0);
411 let sum = arena.push(ExprNode::Add(smallvec::smallvec![x, three]));
412 assert_eq!(describe_nonlinear_term(&arena, sum, &names), None);
413 }
414
415 #[test]
416 fn describe_falls_back_to_index_for_unknown_var() {
417 let mut arena = ExprArena::new();
418 let a = arena.push(ExprNode::Var(VarId(7)));
419 let b = arena.push(ExprNode::Var(VarId(8)));
420 let prod = arena.push(ExprNode::Mul(smallvec::smallvec![a, b]));
421 assert_eq!(describe_nonlinear_term(&arena, prod, &names).as_deref(), Some("v7 * v8"));
422 }
423}