1use arrayvec::ArrayVec;
2
3use crate::{Cubic, Quadratic, different_signs};
4
5impl Cubic {
6 #[doc(hidden)]
11 pub fn eval_opt(&self, x: f64) -> f64 {
12 let [c0, c1, c2, c3] = self.coeffs;
13 let xx = x * x;
14 let xxx = xx * x;
15 c0 + c1 * x + c2 * xx + c3 * xxx
16 }
17
18 #[doc(hidden)]
26 pub fn eval_with_deriv_opt(&self, deriv: &Quadratic, x: f64) -> (f64, f64) {
27 let [c0, c1, c2, c3] = self.coeffs;
28 let [d0, d1, d2] = deriv.coeffs;
29 let xx = x * x;
30 let xxx = xx * x;
31 (c0 + c1 * x + c2 * xx + c3 * xxx, d0 + d1 * x + d2 * xx)
32 }
33
34 fn critical_points(&self) -> Option<(f64, f64)> {
45 let a = 3.0 * self.coeffs[3];
46 let b_2 = self.coeffs[2];
47 let c = self.coeffs[1];
48 let disc_4 = b_2 * b_2 - a * c;
49
50 if !disc_4.is_finite() {
51 return self.rescaled_critical_points();
52 }
53
54 if disc_4 > 0.0 {
55 let q = -(b_2 + disc_4.sqrt().copysign(b_2));
56 let r0 = q / a;
57 let r1 = c / q;
58 Some((r0.min(r1), r0.max(r1)))
59 } else {
60 None
61 }
62 }
63
64 #[cold]
65 fn rescaled_critical_points(&self) -> Option<(f64, f64)> {
66 let scale = 2.0f64.powi(-515);
67 (*self * scale).critical_points()
68 }
69
70 fn one_root(
71 &self,
72 lower: f64,
73 upper: f64,
74 lower_val: f64,
75 upper_val: f64,
76 x_error: f64,
77 ) -> f64 {
78 let deriv = self.deriv();
79 if !deriv.is_finite() {
80 return f64::NAN;
81 }
82 crate::yuksel::find_root(
83 |x| self.eval_opt(x),
84 |x| deriv.eval_opt(x),
85 lower,
86 upper,
87 lower_val,
88 upper_val,
89 x_error,
90 )
91 }
92
93 #[doc(hidden)]
95 pub fn root_between(self, lower: f64, upper: f64, x_error: f64) -> f64 {
96 self.one_root(lower, upper, self.eval(lower), self.eval(upper), x_error)
97 }
98
99 fn first_root(self, lower: f64, upper: f64, x_error: f64) -> Option<f64> {
100 if let Some((x0, x1)) = self.critical_points() {
101 let possible_endpoints: [f64; 3] = [x0, x1, upper];
102 let mut last = lower;
103 let mut last_val = self.eval(last);
104 for x in possible_endpoints {
105 if x > last && x <= upper {
106 let val = self.eval(x);
107 if different_signs(last_val, val) {
108 return Some(self.one_root(last, x, last_val, val, x_error));
109 }
110
111 last = x;
112 last_val = val;
113 }
114 }
115 None
116 } else {
117 let lower_val = self.eval(lower);
118 let upper_val = self.eval(upper);
119 if different_signs(lower_val, upper_val) {
120 Some(self.one_root(lower, upper, lower_val, upper_val, x_error))
121 } else {
122 None
123 }
124 }
125 }
126
127 pub fn roots_between(self, lower: f64, upper: f64, x_error: f64) -> ArrayVec<f64, 3> {
135 let mut ret = ArrayVec::new();
136 let mut scratch = ArrayVec::new();
137 self.roots_between_with_buffer(lower, upper, x_error, &mut scratch, &mut ret);
138 ret
139 }
140
141 pub(crate) fn roots_between_with_buffer<const M: usize>(
142 self,
143 lower: f64,
144 upper: f64,
145 x_error: f64,
146 _scratch: &mut ArrayVec<f64, M>,
147 out: &mut ArrayVec<f64, M>,
148 ) {
149 if let Some(r) = self.first_root(lower, upper, x_error) {
150 out.push(r);
151 let quad = self.deflate(r);
152 if let Some((x0, x1)) = quad.positive_discriminant_roots() {
153 if lower <= x0 && x0 <= upper {
154 out.push(x0);
155 }
156 if lower <= x1 && x1 <= upper {
157 out.push(x1);
158 }
159
160 if lower <= x0 && x0 < r {
164 out.sort_by(|x, y| x.partial_cmp(y).unwrap());
165 }
166 }
167 }
168 }
169
170 #[doc(hidden)]
171 pub fn precondition(&self) -> Cubic {
172 let min_coeff = 2.0f64.powi(-256);
175 let truncate = |x: &mut f64| {
176 if x.abs() <= min_coeff {
177 *x = 0.0
178 }
179 };
180
181 let small_coeff = 2.0f64.powi(-64);
185
186 let large_coeff = 2.0f64.powi(64);
187
188 let mut c = *self;
189 if (self.magnitude() != 0.0 && self.magnitude() <= small_coeff)
190 || self.magnitude() >= large_coeff
191 {
192 c /= self.magnitude();
193 }
194
195 truncate(&mut c.coeffs[0]);
196 truncate(&mut c.coeffs[1]);
197 truncate(&mut c.coeffs[2]);
198 truncate(&mut c.coeffs[3]);
199 c
200 }
201
202 #[doc(hidden)]
207 pub fn roots_blinn(&self) -> ArrayVec<f64, 3> {
208 let mut ret = ArrayVec::new();
209 let a = self.coeffs[3];
210 let b = self.coeffs[2] * (1.0 / 3.0);
211 let c = self.coeffs[1] * (1.0 / 3.0);
212 let d = self.coeffs[0];
213
214 let delta_1 = a * c - b * b;
215 let delta_2 = a * d - b * c;
216 let delta_3 = b * d - c * c;
217 let disc = 4.0 * delta_1 * delta_3 - delta_2 * delta_2;
218
219 if !disc.is_finite() {
220 return ret;
221 }
222 let mut push = |x: f64| {
230 if x.is_finite() {
231 ret.push(x);
232 }
233 };
234 if disc <= 0.0 {
235 let (tilde_a, tilde_c, tilde_d) = if b * b * b * d >= a * c * c * c {
237 (a, delta_1, -2.0 * b * delta_1 + a * delta_2)
238 } else {
239 (d, delta_3, -d * delta_2 + 2.0 * c * delta_3)
240 };
241 let t_0 = -tilde_a.copysign(tilde_d) * (-disc).sqrt();
243 let t_1 = -tilde_d + t_0;
244 let p = (t_1 / 2.0).cbrt();
245
246 let q = if t_0 == t_1 { -p } else { -tilde_c / p };
247 let tilde_x = if tilde_c <= 0.0 {
249 p + q
250 } else {
251 -tilde_d / (p * p + q * q + tilde_c)
252 };
253
254 let (x, w) = if b * b * b * d >= a * c * c * c {
255 (tilde_x - b, a)
256 } else {
257 (-d, tilde_x + c)
258 };
259
260 push(x / w);
261 } else {
262 fn one_root(a_or_d: f64, disc: f64, bar_c: f64, bar_d: f64) -> (f64, f64) {
264 let sqrt_c = (-bar_c).sqrt();
265 let theta = (1.0 / 3.0) * (a_or_d * disc.sqrt()).atan2(-bar_d).abs();
266 let (sin_theta, cos_theta) = theta.sin_cos();
267 let tilde_x_1 = 2.0 * sqrt_c * cos_theta;
269 let tilde_x_3 = sqrt_c * (-cos_theta - 3.0f64.sqrt() * sin_theta);
270 (tilde_x_1, tilde_x_3)
271 }
272
273 let bar_c_a = delta_1;
274 let bar_d_a = -2.0 * b * delta_1 + a * delta_2;
275 let (tilde_x_1_a, tilde_x_3_a) = one_root(a, disc, bar_c_a, bar_d_a);
276
277 let bar_c_d = delta_3;
278 let bar_d_d = -d * delta_2 + 2.0 * c * delta_3;
279 let (tilde_x_1_d, tilde_x_3_d) = one_root(d, disc, bar_c_d, bar_d_d);
280
281 let tilde_x_l = if tilde_x_1_a + tilde_x_3_a > 2.0 * b {
282 tilde_x_1_a
283 } else {
284 tilde_x_3_a
285 };
286 let tilde_x_s = if tilde_x_1_d + tilde_x_3_d < 2.0 * c {
287 tilde_x_1_d
288 } else {
289 tilde_x_3_d
290 };
291
292 let (x_l, w_l) = (tilde_x_l - b, a);
293 let (x_s, w_s) = (-d, tilde_x_s + c);
294
295 let e = w_l * w_s;
296 let f = -x_l * w_s - w_l * x_s;
297 let g = x_l * x_s;
298
299 let (x_m, w_m) = (c * f - b * g, c * e - b * f);
300
301 push(x_l / w_l);
302 push(x_s / w_s);
303 push(x_m / w_m);
304 }
305 ret
306 }
307
308 #[doc(hidden)]
313 pub fn roots_blinn_and_deflate(&self) -> ArrayVec<f64, 3> {
314 let mut ret = ArrayVec::new();
315 let a = self.coeffs[3];
316 let b = self.coeffs[2] * (1.0 / 3.0);
317 let c = self.coeffs[1] * (1.0 / 3.0);
318 let d = self.coeffs[0];
319
320 let delta_1 = a * c - b * b;
321 let delta_2 = a * d - b * c;
322 let delta_3 = b * d - c * c;
323 let disc = 4.0 * delta_1 * delta_3 - delta_2 * delta_2;
324
325 if !disc.is_finite() {
326 return ret;
327 }
328 if disc <= 0.0 {
332 let (tilde_a, tilde_c, tilde_d) = if b * b * b * d >= a * c * c * c {
333 (a, delta_1, -2.0 * b * delta_1 + a * delta_2)
334 } else {
335 (d, delta_3, -d * delta_2 + 2.0 * c * delta_3)
336 };
337 let t_0 = -tilde_a.copysign(tilde_d) * (-disc).sqrt();
338 let t_1 = -tilde_d + t_0;
339 let p = (t_1 / 2.0).cbrt();
340
341 let q = if t_0 == t_1 { -p } else { -tilde_c / p };
342 let tilde_x = if tilde_c <= 0.0 {
343 p + q
344 } else {
345 -tilde_d / (p * p + q * q + tilde_c)
346 };
347
348 let (x, w) = if b * b * b * d >= a * c * c * c {
349 (tilde_x - b, a)
350 } else {
351 (-d, tilde_x + c)
352 };
353
354 if x.is_finite() && w.is_finite() {
355 ret.push(x / w);
356 }
357 } else {
358 fn one_root(a_or_d: f64, disc: f64, bar_c: f64, bar_d: f64) -> (f64, f64) {
359 let sqrt_c = (-bar_c).sqrt();
360 let theta = (1.0 / 3.0) * (a_or_d * disc.sqrt()).atan2(-bar_d).abs();
361 let (sin_theta, cos_theta) = theta.sin_cos();
362 let tilde_x_1 = 2.0 * sqrt_c * cos_theta;
363 let tilde_x_3 = sqrt_c * (-theta.cos() - 3.0f64.sqrt() * sin_theta);
364 (tilde_x_1, tilde_x_3)
365 }
366
367 let bar_c_a = delta_1;
370 let bar_d_a = -2.0 * b * delta_1 + a * delta_2;
371 let (tilde_x_1_a, tilde_x_3_a) = one_root(a, disc, bar_c_a, bar_d_a);
372
373 let bar_c_d = delta_3;
374 let bar_d_d = -d * delta_2 + 2.0 * c * delta_3;
375 let (tilde_x_1_d, tilde_x_3_d) = one_root(d, disc, bar_c_d, bar_d_d);
376
377 let tilde_x_l = if tilde_x_1_a + tilde_x_3_a > 2.0 * b {
378 tilde_x_1_a
379 } else {
380 tilde_x_3_a
381 };
382 let tilde_x_s = if tilde_x_1_d + tilde_x_3_d < 2.0 * c {
384 tilde_x_1_d
385 } else {
386 tilde_x_3_d
387 };
388
389 let (x_l, w_l) = (tilde_x_l - b, a);
390 let (x_s, w_s) = (-d, tilde_x_s + c);
391
392 let x = if (x_l * w_s).abs() <= (x_s * w_l).abs() {
393 x_l / w_l
394 } else {
395 x_s / w_s
396 };
397 if x.is_finite() {
398 ret.push(x);
399 }
400 let q = self.deflate(x);
401 if q.is_finite() {
403 ret.extend(q.roots());
404 ret.sort_by(|x, y| x.partial_cmp(y).unwrap());
405 }
406 }
407 ret
408 }
409}
410
411#[cfg(test)]
412mod tests {
413 use crate::Cubic;
414
415 const TRICKY_CUBICS: [Cubic; 5] = [
416 Cubic::new([
418 1.6149620090145706e-94,
419 1.6149620090145634e-94,
420 1.6149620090145663e-94,
421 9.66803867245343e272,
422 ]),
423 Cubic::new([
428 -6.323283382275869e98,
429 3.0957754283429482e-307,
430 3.095775428342964e-307,
431 3.095775428342951e-307,
432 ]),
433 Cubic::new([
436 -8.522348907129e-161,
437 4.471145208374078e-67,
438 -0.052026185927646074,
439 -2.9441090045938734e-57,
440 ]),
441 Cubic::new([
444 -2.5162489269306657e-175,
445 -2.516248926930655e-175,
446 -2.5162489269306522e-175,
447 -0.39205037382350466,
448 ]),
449 Cubic::new([
450 -6.428720163649757e103,
451 -6.428720163649766e103,
452 -3.3646756114322413e-74,
453 -3.3646756114322547e-74,
454 ]),
455 ];
456
457 #[test]
458 fn smoke() {
459 let poly = Cubic::new([
482 -6.428720163649757e103,
483 -6.428720163649766e103,
484 -3.3646756114322413e-74,
485 -3.3646756114322547e-74,
486 ]);
487
488 let roots = poly.precondition().roots_blinn();
489 dbg!(&roots);
491 for r in roots {
492 dbg!(poly.eval(r));
493 }
494 }
495
496 #[test]
497 fn bad_for_blinn() {
498 for c in TRICKY_CUBICS {
499 dbg!(c.roots_blinn());
500 dbg!(c.roots_blinn_and_deflate());
501 }
502 }
503
504 fn check_root_values(c: &Cubic, roots: &[f64]) {
508 let magnitude = c.magnitude().max(1.0);
511 let accuracy = magnitude * 1e-12;
512
513 for r in roots {
514 let accuracy = accuracy * r.abs().powi(3).max(1.0);
518 let y = c.eval(*r);
519 if y.is_finite() {
520 assert!(
521 y.abs() <= accuracy,
522 "cubic {c:?} had root {r} evaluate to {y:?}, but expected {accuracy:?}"
523 );
524 }
525 }
526 }
527
528 #[test]
529 fn root_evaluation() {
530 arbtest::arbtest(|u| {
531 let c = crate::arbitrary::cubic(u)?;
532
533 let roots = c.roots_between(-10.0, 10.0, 1e-13);
537 if roots.iter().all(|r| r.is_finite()) {
538 assert!(roots.is_sorted());
539 }
540 check_root_values(&c, &roots);
541
542 let preconditioned = c.precondition();
543 check_root_values(&c, &preconditioned.roots_blinn());
544 check_root_values(&c, &preconditioned.roots_blinn_and_deflate());
545
546 Ok(())
552 })
553 .budget_ms(5_000);
554 }
555
556 #[test]
557 #[ignore]
558 fn root_evaluation_kurbo() {
559 arbtest::arbtest(|u| {
560 let c = crate::arbitrary::cubic(u)?;
561 let magnitude = c.magnitude().max(1.0);
564 let accuracy = magnitude * 1e-12;
565
566 let &[c0, c1, c2, c3] = c.coeffs();
570 for r in kurbo::common::solve_cubic(c0, c1, c2, c3) {
571 let y = c.eval(r);
572 if y.is_finite() {
573 assert!(y.abs() <= accuracy);
574 }
575 }
576 Ok(())
577 })
578 .budget_ms(5_000);
579 }
580}