Skip to content

Commit 73cb113

Browse files
committed
adaptive step + retry for num grad
1 parent ad609d2 commit 73cb113

3 files changed

Lines changed: 267 additions & 53 deletions

File tree

src/fit.rs

Lines changed: 69 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -138,6 +138,8 @@ const MAX_POLYNOMIAL_PARAMS: usize = 10;
138138
const STEEPEST_DESCENT_GRAD_TOL: f64 = 1e-12;
139139
const HESSIAN_FD_REL_STEP: f64 = 1e-4;
140140
const HESSIAN_FD_MIN_STEP: f64 = 1e-6;
141+
// Пробуем базовый шаг, затем уменьшаем/увеличиваем его, чтобы переживать локальные NaN/Inf.
142+
const FD_STEP_RETRY_FACTORS: [f64; 5] = [1.0, 0.5, 2.0, 0.25, 4.0];
141143
pub(crate) const HESSIAN_DIAGONAL_JITTER: f64 = models::HESSIAN_DIAGONAL_JITTER;
142144

143145
fn positive_x(value: f64) -> f64 {
@@ -153,6 +155,24 @@ fn gradient_l2_norm(values: &[f64]) -> f64 {
153155
values.iter().map(|value| value * value).sum::<f64>().sqrt()
154156
}
155157

158+
#[inline]
159+
fn finite_central_difference(value_plus: f64, value_minus: f64, step: f64) -> Option<f64> {
160+
if !value_plus.is_finite() || !value_minus.is_finite() {
161+
return None;
162+
}
163+
let derivative = (value_plus - value_minus) / (2.0 * step);
164+
if derivative.is_finite() {
165+
Some(derivative)
166+
} else {
167+
None
168+
}
169+
}
170+
171+
#[inline]
172+
fn finite_array1(values: &Array1<f64>) -> bool {
173+
values.iter().all(|value| value.is_finite())
174+
}
175+
156176
fn vec_to_array1(values: &[f64]) -> Array1<f64> {
157177
Array1::from_vec(values.to_vec())
158178
}
@@ -194,20 +214,48 @@ where
194214
let dimension = param_slice.len();
195215
let mut hessian = Array2::zeros((dimension, dimension));
196216
let mut probe = param.clone();
217+
let mut column_values = vec![0.0; dimension];
197218

198219
for column in 0..dimension {
199-
let step =
220+
let base_step =
200221
((param_slice[column].abs() + 1.0) * HESSIAN_FD_REL_STEP).max(HESSIAN_FD_MIN_STEP);
201-
probe[column] = param[column] + step;
202-
let grad_plus = problem.gradient(&probe)?;
203-
probe[column] = param[column] - step;
204-
let grad_minus = problem.gradient(&probe)?;
205-
probe[column] = param[column];
222+
let mut computed = false;
223+
for factor in FD_STEP_RETRY_FACTORS {
224+
let step = base_step * factor;
225+
probe[column] = param[column] + step;
226+
let grad_plus = problem.gradient(&probe)?;
227+
probe[column] = param[column] - step;
228+
let grad_minus = problem.gradient(&probe)?;
229+
probe[column] = param[column];
230+
231+
if !finite_array1(&grad_plus) || !finite_array1(&grad_minus) {
232+
continue;
233+
}
206234

207-
let denom = 2.0 * step;
208-
for row in 0..dimension {
209-
let value = (grad_plus[row] - grad_minus[row]) / denom;
210-
hessian[[row, column]] = if value.is_finite() { value } else { 0.0 };
235+
let denom = 2.0 * step;
236+
let mut column_is_finite = true;
237+
for row in 0..dimension {
238+
let value = (grad_plus[row] - grad_minus[row]) / denom;
239+
if !value.is_finite() {
240+
column_is_finite = false;
241+
break;
242+
}
243+
column_values[row] = value;
244+
}
245+
246+
if column_is_finite {
247+
for row in 0..dimension {
248+
hessian[[row, column]] = column_values[row];
249+
}
250+
computed = true;
251+
break;
252+
}
253+
}
254+
255+
if !computed {
256+
for row in 0..dimension {
257+
hessian[[row, column]] = 0.0;
258+
}
211259
}
212260
}
213261

@@ -2023,19 +2071,18 @@ impl Gradient for SplineProblem {
20232071
let mut gradient = Array1::zeros(param_slice.len());
20242072
for (index, gradient_value) in gradient.iter_mut().enumerate() {
20252073
// Численный градиент по центральной схеме конечной разности.
2026-
let step =
2074+
let base_step =
20272075
((param_slice[index].abs() + 1.0) * SPLINE_FD_REL_STEP).max(SPLINE_FD_MIN_STEP);
2028-
probe[index] = param_slice[index] + step;
2029-
let cost_plus = self.evaluate_objective(array1_as_slice(&probe));
2030-
probe[index] = param_slice[index] - step;
2031-
let cost_minus = self.evaluate_objective(array1_as_slice(&probe));
2032-
probe[index] = param_slice[index];
2033-
let derivative = (cost_plus - cost_minus) / (2.0 * step);
2034-
*gradient_value = if derivative.is_finite() {
2035-
derivative
2036-
} else {
2037-
LARGE_COST
2038-
};
2076+
let derivative = FD_STEP_RETRY_FACTORS.iter().copied().find_map(|factor| {
2077+
let step = base_step * factor;
2078+
probe[index] = param_slice[index] + step;
2079+
let cost_plus = self.evaluate_objective(array1_as_slice(&probe));
2080+
probe[index] = param_slice[index] - step;
2081+
let cost_minus = self.evaluate_objective(array1_as_slice(&probe));
2082+
probe[index] = param_slice[index];
2083+
finite_central_difference(cost_plus, cost_minus, step)
2084+
});
2085+
*gradient_value = derivative.unwrap_or(LARGE_COST);
20392086
}
20402087
Ok(gradient)
20412088
}

src/fit/tests.rs

Lines changed: 121 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -1,24 +1,26 @@
11
use super::simd;
22
use super::{
33
CurveProblem, CurveProblemPredictionLoss, DEFAULT_SPLINE_KNOTS, FitError,
4-
IncrementalSplineFitRunner, IncrementalSplineFitStep, MetricQuantization,
5-
MetricQuantizationDecimalPlaces, OptimizationLossMetric, SplineConfig, SplineDuplicateXPolicy,
6-
SplineExtrapolation, SplineFamilyKind, SplineFinalizeContext, SplineKnotStrategy,
7-
approximate_spline_knots, build_spline_initial_curve_from_knot_y,
8-
build_spline_result_from_knot_y, calculate_iteration_metrics,
9-
calculate_iteration_metrics_with_quantization, calculate_metrics, evaluate_linear_spline,
10-
expanded_spline_curve_x_bounds, fit_akima_spline, fit_akima_spline_with_config, fit_curve,
11-
fit_curve_with_optimizer_config, fit_curve_with_progress,
12-
fit_curve_with_progress_and_optimizer_config,
4+
HESSIAN_DIAGONAL_JITTER, HESSIAN_FD_MIN_STEP, HESSIAN_FD_REL_STEP, IncrementalSplineFitRunner,
5+
IncrementalSplineFitStep, MetricQuantization, MetricQuantizationDecimalPlaces,
6+
OptimizationLossMetric, SplineConfig, SplineDuplicateXPolicy, SplineExtrapolation,
7+
SplineFamilyKind, SplineFinalizeContext, SplineKnotStrategy, approximate_spline_knots,
8+
build_spline_initial_curve_from_knot_y, build_spline_result_from_knot_y,
9+
calculate_iteration_metrics, calculate_iteration_metrics_with_quantization, calculate_metrics,
10+
evaluate_linear_spline, expanded_spline_curve_x_bounds, fit_akima_spline,
11+
fit_akima_spline_with_config, fit_curve, fit_curve_with_optimizer_config,
12+
fit_curve_with_progress, fit_curve_with_progress_and_optimizer_config,
1313
fit_curve_with_progress_and_optimizer_config_and_loss_metric, fit_linear_spline,
1414
fit_linear_spline_with_config, fit_monotone_cubic_spline, fit_natural_cubic_spline,
15-
sorted_points_with_duplicate_policy,
15+
numerical_hessian_from_gradient, sorted_points_with_duplicate_policy,
1616
};
1717
use crate::domain::{
1818
AdamConfig, CurveFamily, CurveParams, InputError, LbfgsConfig, NelderMeadConfig,
1919
NewtonCgConfig, OptimizerConfig, Point, Points, SgdConfig, SteepestDescentConfig,
2020
};
2121
use crate::models::{self, ObjectiveGrad, ObjectiveHessian, ObjectiveValue, PredictionLoss};
22+
use argmin::core::Gradient;
23+
use ndarray::Array1;
2224

2325
#[derive(Clone, Copy)]
2426
struct MsePredictionLoss;
@@ -65,6 +67,115 @@ fn quantization(decimal_places: u8) -> MetricQuantization {
6567
)
6668
}
6769

70+
struct RetryGradientProblem {
71+
center: f64,
72+
invalid_step: f64,
73+
}
74+
75+
struct AlwaysInvalidGradientProblem;
76+
77+
impl Gradient for RetryGradientProblem {
78+
type Param = Array1<f64>;
79+
type Gradient = Array1<f64>;
80+
81+
fn gradient(&self, param: &Self::Param) -> Result<Self::Gradient, argmin::core::Error> {
82+
let delta = (param[0] - self.center).abs();
83+
if (delta - self.invalid_step).abs() <= 1e-14 {
84+
return Ok(Array1::from_vec(vec![f64::NAN]));
85+
}
86+
Ok(Array1::from_vec(vec![2.0 * param[0]]))
87+
}
88+
}
89+
90+
impl Gradient for AlwaysInvalidGradientProblem {
91+
type Param = Array1<f64>;
92+
type Gradient = Array1<f64>;
93+
94+
fn gradient(&self, param: &Self::Param) -> Result<Self::Gradient, argmin::core::Error> {
95+
Ok(Array1::from_vec(vec![f64::NAN; param.len()]))
96+
}
97+
}
98+
99+
#[test]
100+
fn central_diff_gradient_retries_step_when_primary_step_is_invalid() {
101+
let param = [1.0_f64];
102+
let rel_step = 1e-3;
103+
let min_step = 1e-3;
104+
let base_step = ((param[0].abs() + 1.0) * rel_step).max(min_step);
105+
let mut gradient = [0.0];
106+
models::central_diff_gradient_from_value(
107+
&param,
108+
rel_step,
109+
min_step,
110+
|probe| {
111+
let delta = (probe[0] - param[0]).abs();
112+
if (delta - base_step).abs() <= 1e-14 {
113+
f64::NAN
114+
} else {
115+
probe[0] * probe[0]
116+
}
117+
},
118+
&mut gradient,
119+
);
120+
121+
assert_near(gradient[0], 2.0, 1e-10);
122+
}
123+
124+
#[test]
125+
fn central_diff_hessian_retries_step_when_primary_step_is_invalid() {
126+
let param = [1.0_f64];
127+
let rel_step = 1e-3;
128+
let min_step = 1e-3;
129+
let base_step = ((param[0].abs() + 1.0) * rel_step).max(min_step);
130+
let hessian = models::central_diff_hessian_from_gradient(
131+
&param,
132+
rel_step,
133+
min_step,
134+
|probe, gradient_out| {
135+
let delta = (probe[0] - param[0]).abs();
136+
gradient_out[0] = if (delta - base_step).abs() <= 1e-14 {
137+
f64::NAN
138+
} else {
139+
2.0 * probe[0]
140+
};
141+
},
142+
);
143+
144+
assert_near(hessian[[0, 0]], 2.0, 1e-10);
145+
}
146+
147+
#[test]
148+
fn fit_numerical_hessian_retries_step_when_primary_step_is_invalid() {
149+
let param = Array1::from_vec(vec![1.0_f64]);
150+
let base_step = ((param[0].abs() + 1.0) * HESSIAN_FD_REL_STEP).max(HESSIAN_FD_MIN_STEP);
151+
let problem = RetryGradientProblem {
152+
center: param[0],
153+
invalid_step: base_step,
154+
};
155+
let hessian =
156+
numerical_hessian_from_gradient(&problem, &param).expect("hessian must be computed");
157+
158+
assert_near(hessian[[0, 0]], 2.0 + HESSIAN_DIAGONAL_JITTER, 1e-10);
159+
}
160+
161+
#[test]
162+
fn central_diff_gradient_falls_back_to_zero_when_all_retry_steps_invalid() {
163+
let param = [1.0_f64];
164+
let mut gradient = [123.0];
165+
models::central_diff_gradient_from_value(&param, 1e-3, 1e-3, |_probe| f64::NAN, &mut gradient);
166+
167+
assert_eq!(gradient[0], 0.0);
168+
}
169+
170+
#[test]
171+
fn fit_numerical_hessian_falls_back_to_diagonal_jitter_when_all_retry_steps_invalid() {
172+
let param = Array1::from_vec(vec![1.0_f64]);
173+
let hessian = numerical_hessian_from_gradient(&AlwaysInvalidGradientProblem, &param)
174+
.expect("hessian must be computed even with invalid gradients");
175+
176+
assert_near(hessian[[0, 0]], HESSIAN_DIAGONAL_JITTER, 1e-15);
177+
}
178+
68179
#[test]
69180
fn curve_objective_arrhenius_is_consistent_across_levels() {
70181
let x_values = [0.4, 0.8, 1.4, 2.5, 4.0];

0 commit comments

Comments
 (0)