@@ -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
2427mod curve;
2528mod 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
0 commit comments