Skip to content

Commit 28f5bec

Browse files
committed
transition to a new curve fitting pipeline
1 parent ce8031b commit 28f5bec

35 files changed

Lines changed: 1462 additions & 1627 deletions

src/domain.rs

Lines changed: 0 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -544,10 +544,6 @@ impl CurveFamily {
544544
},
545545
}
546546
}
547-
548-
pub(crate) fn evaluate_raw(self, params: &[f64], x: f64) -> f64 {
549-
models::evaluate_raw(self, params, x)
550-
}
551547
}
552548

553549
impl fmt::Display for CurveFamily {

src/fit.rs

Lines changed: 84 additions & 52 deletions
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,10 @@ use crate::domain::{
1919
AdamConfig, CurveFamily, CurveParams, FitResult, InputError, LbfgsConfig, NelderMeadConfig,
2020
NewtonCgConfig, OptimizerConfig, Points, SgdConfig, SteepestDescentConfig,
2121
};
22-
use crate::models::{self, ObjectiveGrad, ObjectiveHessian, ObjectiveValue, PredictionLoss};
22+
use crate::models::{
23+
self, ObjectiveGrad, ObjectiveHessian, ObjectiveValue, PredictionLoss, TermGrad, TermHessian,
24+
TermValue,
25+
};
2326

2427
mod curve;
2528
mod simd;
@@ -427,114 +430,143 @@ struct CurveProblemObjective<'a> {
427430
problem: &'a CurveProblem,
428431
}
429432

430-
impl CurveProblemObjective<'_> {
433+
#[derive(Clone, Copy)]
434+
struct CurveProblemTerm<'a> {
435+
problem: &'a CurveProblem,
436+
}
437+
438+
impl CurveProblemTerm<'_> {
431439
fn simd_enabled(&self) -> bool {
432440
matches!(
433441
self.problem.metric_quantization,
434442
MetricQuantization::Disabled
435443
)
436444
}
437445

438-
fn value(&self, param: &[f64]) -> f64 {
446+
fn fallback_term(&self) -> models::DataTerm<'_, CurveProblemPredictionLoss<'_>> {
447+
let loss = CurveProblemPredictionLoss {
448+
problem: self.problem,
449+
};
450+
models::DataTerm::new(
451+
self.problem.family,
452+
self.problem.point_x.as_ref(),
453+
self.problem.point_y.as_ref(),
454+
loss,
455+
)
456+
}
457+
}
458+
459+
impl TermValue for CurveProblemTerm<'_> {
460+
fn add_value(&self, param: &[f64], value: &mut f64) {
439461
if self.simd_enabled() && self.problem.family.is_polynomial() {
440-
return simd::polynomial_cost(
462+
*value += simd::polynomial_cost(
441463
param,
442464
self.problem.point_x.as_ref(),
443465
self.problem.point_y.as_ref(),
444466
self.problem.loss_metric,
445467
);
468+
return;
446469
}
447470
if self.simd_enabled() && self.problem.family == CurveFamily::Inverse {
448-
return simd::inverse_cost(
471+
*value += simd::inverse_cost(
449472
param,
450473
self.problem.point_x.as_ref(),
451474
self.problem.point_y.as_ref(),
452475
self.problem.loss_metric,
453476
);
477+
return;
454478
}
455-
456-
let loss = CurveProblemPredictionLoss {
457-
problem: self.problem,
458-
};
459-
let term = models::DataTerm::new(
460-
self.problem.family,
461-
self.problem.point_x.as_ref(),
462-
self.problem.point_y.as_ref(),
463-
loss,
464-
);
465-
let objective = models::CurveObjective::new(self.problem.family.parameter_count(), term);
466-
objective.value(param)
479+
self.fallback_term().add_value(param, value);
467480
}
481+
}
468482

469-
fn value_grad(&self, param: &[f64]) -> (f64, Vec<f64>) {
483+
impl TermGrad for CurveProblemTerm<'_> {
484+
fn add_value_grad(&self, param: &[f64], value: &mut f64, gradient: &mut [f64]) {
470485
if self.simd_enabled() && self.problem.family.is_polynomial() {
471-
let mut gradient = vec![0.0; self.problem.family.parameter_count()];
486+
let mut local_gradient = vec![0.0; self.problem.family.parameter_count()];
472487
simd::accumulate_polynomial_gradient(
473488
self.problem.point_x.as_ref(),
474489
self.problem.point_y.as_ref(),
475490
param,
476491
self.problem.loss_metric,
477-
&mut gradient,
492+
&mut local_gradient,
478493
);
479494
let sample_scale = 1.0 / self.problem.point_x.len() as f64;
480-
for value in &mut gradient {
481-
*value *= sample_scale;
495+
for local_value in &mut local_gradient {
496+
*local_value *= sample_scale;
482497
}
483-
let value = simd::polynomial_cost(
498+
*value += simd::polynomial_cost(
484499
param,
485500
self.problem.point_x.as_ref(),
486501
self.problem.point_y.as_ref(),
487502
self.problem.loss_metric,
488503
);
489-
return (value, gradient);
504+
for (gradient_value, local_value) in gradient.iter_mut().zip(local_gradient) {
505+
*gradient_value += local_value;
506+
}
507+
return;
490508
}
491509
if self.simd_enabled() && self.problem.family == CurveFamily::Inverse {
492-
let mut gradient = vec![0.0; self.problem.family.parameter_count()];
510+
let mut local_gradient = vec![0.0; self.problem.family.parameter_count()];
493511
simd::accumulate_inverse_gradient(
494512
self.problem.point_x.as_ref(),
495513
self.problem.point_y.as_ref(),
496514
param,
497515
self.problem.loss_metric,
498-
&mut gradient,
516+
&mut local_gradient,
499517
);
500518
let sample_scale = 1.0 / self.problem.point_x.len() as f64;
501-
for value in &mut gradient {
502-
*value *= sample_scale;
519+
for local_value in &mut local_gradient {
520+
*local_value *= sample_scale;
503521
}
504-
let value = simd::inverse_cost(
522+
*value += simd::inverse_cost(
505523
param,
506524
self.problem.point_x.as_ref(),
507525
self.problem.point_y.as_ref(),
508526
self.problem.loss_metric,
509527
);
510-
return (value, gradient);
528+
for (gradient_value, local_value) in gradient.iter_mut().zip(local_gradient) {
529+
*gradient_value += local_value;
530+
}
531+
return;
511532
}
533+
self.fallback_term().add_value_grad(param, value, gradient);
534+
}
535+
}
512536

513-
let loss = CurveProblemPredictionLoss {
514-
problem: self.problem,
515-
};
516-
let term = models::DataTerm::new(
517-
self.problem.family,
518-
self.problem.point_x.as_ref(),
519-
self.problem.point_y.as_ref(),
520-
loss,
521-
);
522-
let objective = models::CurveObjective::new(self.problem.family.parameter_count(), term);
523-
objective.value_grad(param)
537+
impl TermHessian for CurveProblemTerm<'_> {
538+
fn add_value_grad_hessian(
539+
&self,
540+
param: &[f64],
541+
value: &mut f64,
542+
gradient: &mut [f64],
543+
hessian: &mut Array2<f64>,
544+
) {
545+
self.fallback_term()
546+
.add_value_grad_hessian(param, value, gradient, hessian);
547+
}
548+
}
549+
550+
impl CurveProblemObjective<'_> {
551+
fn objective(&self) -> models::CurveObjective<CurveProblemTerm<'_>> {
552+
models::CurveObjective::new(
553+
self.problem.family.parameter_count(),
554+
CurveProblemTerm {
555+
problem: self.problem,
556+
},
557+
)
558+
}
559+
560+
fn value(&self, param: &[f64]) -> f64 {
561+
self.objective().value(param)
562+
}
563+
564+
fn value_grad(&self, param: &[f64]) -> (f64, Vec<f64>) {
565+
self.objective().value_grad(param)
524566
}
525567

526568
fn value_grad_hessian(&self, param: &[f64]) -> (f64, Vec<f64>, Array2<f64>) {
527-
let loss = CurveProblemPredictionLoss {
528-
problem: self.problem,
529-
};
530-
let term = models::DataTerm::new(
531-
self.problem.family,
532-
self.problem.point_x.as_ref(),
533-
self.problem.point_y.as_ref(),
534-
loss,
535-
);
536-
let objective = models::CurveObjective::new(self.problem.family.parameter_count(), term);
537-
objective.value_grad_hessian(param)
569+
self.objective().value_grad_hessian(param)
538570
}
539571
}
540572

src/fit/tests.rs

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -73,7 +73,7 @@ fn curve_objective_arrhenius_is_consistent_across_levels() {
7373
let probe_params = [1.2, 0.5];
7474
let y_values = x_values
7575
.iter()
76-
.map(|&x| models::evaluate_raw(CurveFamily::Arrhenius, &true_params, x))
76+
.map(|&x| models::value_at(CurveFamily::Arrhenius, &true_params, x))
7777
.collect::<Vec<_>>();
7878

7979
let term = models::DataTerm::new(
@@ -136,7 +136,7 @@ fn curve_objective_emg_matches_numerical_derivatives() {
136136
let probe_params = [1.2, 0.1, 0.5, 0.4, 0.0];
137137
let y_values = x_values
138138
.iter()
139-
.map(|&x| models::evaluate_raw(CurveFamily::Emg, &true_params, x))
139+
.map(|&x| models::value_at(CurveFamily::Emg, &true_params, x))
140140
.collect::<Vec<_>>();
141141

142142
let term = models::DataTerm::new(CurveFamily::Emg, &x_values, &y_values, MsePredictionLoss);

0 commit comments

Comments
 (0)