Skip to content

Commit bd8a4d2

Browse files
committed
added pairwise grad function
1 parent 28f5bec commit bd8a4d2

26 files changed

Lines changed: 725 additions & 324 deletions

src/models/arctangent_step.rs

Lines changed: 29 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -17,28 +17,45 @@ pub(super) fn value_at(param: &[f64], x: f64) -> f64 {
1717
amplitude * (slope * (x - x0)).atan() + offset
1818
}
1919

20+
#[inline]
21+
pub(super) fn value_grad_at(param: &[f64], x: f64, grad: &mut [f64]) -> f64 {
22+
debug_assert_eq!(grad.len(), 4);
23+
24+
let amplitude = param[0];
25+
let slope = param[1];
26+
let x0 = param[2];
27+
let offset = param[3];
28+
let z = slope * (x - x0);
29+
let atan_z = z.atan();
30+
let inv_den = 1.0 / (1.0 + z * z);
31+
32+
grad[0] = atan_z;
33+
grad[1] = amplitude * (x - x0) * inv_den;
34+
grad[2] = -amplitude * slope * inv_den;
35+
grad[3] = 1.0;
36+
37+
amplitude * atan_z + offset
38+
}
39+
2040
pub(super) fn add_value_grad(
2141
x_values: &[f64],
2242
param: &[f64],
2343
value_first: &[f64],
2444
gradient: &mut [f64],
2545
) {
26-
let amplitude = param[0];
27-
let slope = param[1];
28-
let x0 = param[2];
46+
debug_assert_eq!(x_values.len(), value_first.len());
47+
debug_assert_eq!(gradient.len(), param.len());
2948

49+
let mut point_grad = [0.0; 4];
3050
let mut index = 0;
3151
while index < x_values.len() {
32-
let x = x_values[index];
33-
let z = slope * (x - x0);
34-
let atan_z = z.atan();
35-
let inv_den = 1.0 / (1.0 + z * z);
36-
let residual = value_first[index];
52+
let upstream = value_first[index];
53+
value_grad_at(param, x_values[index], &mut point_grad);
3754

38-
gradient[0] += residual * atan_z;
39-
gradient[1] += residual * (amplitude * (x - x0) * inv_den);
40-
gradient[2] += residual * (-amplitude * slope * inv_den);
41-
gradient[3] += residual;
55+
gradient[0] += upstream * point_grad[0];
56+
gradient[1] += upstream * point_grad[1];
57+
gradient[2] += upstream * point_grad[2];
58+
gradient[3] += upstream * point_grad[3];
4259
index += 1;
4360
}
4461
}

src/models/arrhenius.rs

Lines changed: 22 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -16,22 +16,37 @@ pub(super) fn value_at(param: &[f64], x: f64) -> f64 {
1616
prefactor * (temp_coeff / x).exp()
1717
}
1818

19+
#[inline]
20+
pub(super) fn value_grad_at(param: &[f64], x: f64, grad: &mut [f64]) -> f64 {
21+
debug_assert_eq!(grad.len(), 2);
22+
23+
let x = positive_x(x);
24+
let prefactor = param[0];
25+
let temp_coeff = param[1];
26+
let exp_term = (temp_coeff / x).exp();
27+
28+
grad[0] = exp_term;
29+
grad[1] = prefactor * exp_term / x;
30+
31+
prefactor * exp_term
32+
}
33+
1934
pub(super) fn add_value_grad(
2035
x_values: &[f64],
2136
param: &[f64],
2237
value_first: &[f64],
2338
gradient: &mut [f64],
2439
) {
25-
let prefactor = param[0];
26-
let temp_coeff = param[1];
40+
debug_assert_eq!(x_values.len(), value_first.len());
41+
debug_assert_eq!(gradient.len(), param.len());
2742

43+
let mut point_grad = [0.0; 2];
2844
let mut index = 0;
2945
while index < x_values.len() {
30-
let x = positive_x(x_values[index]);
31-
let exp_term = (temp_coeff / x).exp();
32-
let residual = value_first[index];
33-
gradient[0] += residual * exp_term;
34-
gradient[1] += residual * (prefactor * exp_term / x);
46+
let upstream = value_first[index];
47+
value_grad_at(param, x_values[index], &mut point_grad);
48+
gradient[0] += upstream * point_grad[0];
49+
gradient[1] += upstream * point_grad[1];
3550
index += 1;
3651
}
3752
}

src/models/bi_exponential.rs

Lines changed: 32 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -17,29 +17,47 @@ pub(super) fn value_at(param: &[f64], x: f64) -> f64 {
1717
a1 * (-k1 * x).exp() + a2 * (-k2 * x).exp() + offset
1818
}
1919

20+
#[inline]
21+
pub(super) fn value_grad_at(param: &[f64], x: f64, grad: &mut [f64]) -> f64 {
22+
debug_assert_eq!(grad.len(), 5);
23+
24+
let a1 = param[0];
25+
let k1 = param[1];
26+
let a2 = param[2];
27+
let k2 = param[3];
28+
let offset = param[4];
29+
let exp1 = (-k1 * x).exp();
30+
let exp2 = (-k2 * x).exp();
31+
32+
grad[0] = exp1;
33+
grad[1] = -a1 * x * exp1;
34+
grad[2] = exp2;
35+
grad[3] = -a2 * x * exp2;
36+
grad[4] = 1.0;
37+
38+
a1 * exp1 + a2 * exp2 + offset
39+
}
40+
2041
pub(super) fn add_value_grad(
2142
x_values: &[f64],
2243
param: &[f64],
2344
value_first: &[f64],
2445
gradient: &mut [f64],
2546
) {
26-
let a1 = param[0];
27-
let k1 = param[1];
28-
let a2 = param[2];
29-
let k2 = param[3];
47+
debug_assert_eq!(x_values.len(), value_first.len());
48+
debug_assert_eq!(gradient.len(), param.len());
3049

50+
let mut point_grad = [0.0; 5];
3151
let mut index = 0;
3252
while index < x_values.len() {
33-
let x = x_values[index];
34-
let exp1 = (-k1 * x).exp();
35-
let exp2 = (-k2 * x).exp();
36-
let residual = value_first[index];
37-
38-
gradient[0] += residual * exp1;
39-
gradient[1] += residual * (-a1 * x * exp1);
40-
gradient[2] += residual * exp2;
41-
gradient[3] += residual * (-a2 * x * exp2);
42-
gradient[4] += residual;
53+
let upstream = value_first[index];
54+
value_grad_at(param, x_values[index], &mut point_grad);
55+
56+
gradient[0] += upstream * point_grad[0];
57+
gradient[1] += upstream * point_grad[1];
58+
gradient[2] += upstream * point_grad[2];
59+
gradient[3] += upstream * point_grad[3];
60+
gradient[4] += upstream * point_grad[4];
4361
index += 1;
4462
}
4563
}

src/models/damped_sinusoid.rs

Lines changed: 34 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -18,31 +18,49 @@ pub(super) fn value_at(param: &[f64], x: f64) -> f64 {
1818
amplitude * (-damping * x).exp() * (omega * x + phi).sin() + offset
1919
}
2020

21+
#[inline]
22+
pub(super) fn value_grad_at(param: &[f64], x: f64, grad: &mut [f64]) -> f64 {
23+
debug_assert_eq!(grad.len(), 5);
24+
25+
let amplitude = param[0];
26+
let damping = param[1];
27+
let omega = param[2];
28+
let phi = param[3];
29+
let offset = param[4];
30+
let exp_part = (-damping * x).exp();
31+
let angle = omega * x + phi;
32+
let sin_part = angle.sin();
33+
let cos_part = angle.cos();
34+
35+
grad[0] = exp_part * sin_part;
36+
grad[1] = -amplitude * x * exp_part * sin_part;
37+
grad[2] = amplitude * exp_part * cos_part * x;
38+
grad[3] = amplitude * exp_part * cos_part;
39+
grad[4] = 1.0;
40+
41+
amplitude * exp_part * sin_part + offset
42+
}
43+
2144
pub(super) fn add_value_grad(
2245
x_values: &[f64],
2346
param: &[f64],
2447
value_first: &[f64],
2548
gradient: &mut [f64],
2649
) {
27-
let amplitude = param[0];
28-
let damping = param[1];
29-
let omega = param[2];
30-
let phi = param[3];
50+
debug_assert_eq!(x_values.len(), value_first.len());
51+
debug_assert_eq!(gradient.len(), param.len());
3152

53+
let mut point_grad = [0.0; 5];
3254
let mut index = 0;
3355
while index < x_values.len() {
34-
let x = x_values[index];
35-
let exp_part = (-damping * x).exp();
36-
let angle = omega * x + phi;
37-
let sin_part = angle.sin();
38-
let cos_part = angle.cos();
39-
let residual = value_first[index];
40-
41-
gradient[0] += residual * exp_part * sin_part;
42-
gradient[1] += residual * (-amplitude * x * exp_part * sin_part);
43-
gradient[2] += residual * (amplitude * exp_part * cos_part * x);
44-
gradient[3] += residual * (amplitude * exp_part * cos_part);
45-
gradient[4] += residual;
56+
let upstream = value_first[index];
57+
value_grad_at(param, x_values[index], &mut point_grad);
58+
59+
gradient[0] += upstream * point_grad[0];
60+
gradient[1] += upstream * point_grad[1];
61+
gradient[2] += upstream * point_grad[2];
62+
gradient[3] += upstream * point_grad[3];
63+
gradient[4] += upstream * point_grad[4];
4664
index += 1;
4765
}
4866
}

src/models/dispatch.rs

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -92,6 +92,7 @@ pub(crate) fn add_model_grad_unscaled(
9292
gradient: &mut [f64],
9393
) {
9494
debug_assert_eq!(x_values.len(), value_first.len());
95+
debug_assert_eq!(gradient.len(), param.len());
9596

9697
if family.is_polynomial() {
9798
polynomial::add_value_grad(x_values, param, value_first, gradient);

src/models/emg.rs

Lines changed: 27 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -37,12 +37,35 @@ pub(super) fn value_at(param: &[f64], x: f64) -> f64 {
3737
}
3838
}
3939

40+
#[inline]
41+
pub(super) fn value_grad_at(param: &[f64], x: f64, grad: &mut [f64]) -> f64 {
42+
debug_assert_eq!(grad.len(), 5);
43+
grad.fill(0.0);
44+
value_at(param, x)
45+
}
46+
4047
pub(super) fn add_value_grad(
41-
_x_values: &[f64],
42-
_param: &[f64],
43-
_value_first: &[f64],
44-
_gradient: &mut [f64],
48+
x_values: &[f64],
49+
param: &[f64],
50+
value_first: &[f64],
51+
gradient: &mut [f64],
4552
) {
53+
debug_assert_eq!(x_values.len(), value_first.len());
54+
debug_assert_eq!(gradient.len(), param.len());
55+
56+
let mut point_grad = [0.0; 5];
57+
let mut index = 0;
58+
while index < x_values.len() {
59+
let upstream = value_first[index];
60+
value_grad_at(param, x_values[index], &mut point_grad);
61+
62+
gradient[0] += upstream * point_grad[0];
63+
gradient[1] += upstream * point_grad[1];
64+
gradient[2] += upstream * point_grad[2];
65+
gradient[3] += upstream * point_grad[3];
66+
gradient[4] += upstream * point_grad[4];
67+
index += 1;
68+
}
4669
}
4770

4871
pub(super) fn add_value_grad_raw_hessian(

src/models/exponential_basic.rs

Lines changed: 25 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -15,23 +15,40 @@ pub(super) fn value_at(param: &[f64], x: f64) -> f64 {
1515
offset + amplitude * (-decay_rate * x).exp()
1616
}
1717

18+
#[inline]
19+
pub(super) fn value_grad_at(param: &[f64], x: f64, grad: &mut [f64]) -> f64 {
20+
debug_assert_eq!(grad.len(), 3);
21+
22+
let offset = param[0];
23+
let amplitude = param[1];
24+
let decay_rate = param[2];
25+
let exp_part = (-decay_rate * x).exp();
26+
27+
grad[0] = 1.0;
28+
grad[1] = exp_part;
29+
grad[2] = -amplitude * x * exp_part;
30+
31+
offset + amplitude * exp_part
32+
}
33+
1834
pub(super) fn add_value_grad(
1935
x_values: &[f64],
2036
param: &[f64],
2137
value_first: &[f64],
2238
gradient: &mut [f64],
2339
) {
24-
let amplitude = param[1];
25-
let decay_rate = param[2];
40+
debug_assert_eq!(x_values.len(), value_first.len());
41+
debug_assert_eq!(gradient.len(), param.len());
2642

43+
let mut point_grad = [0.0; 3];
2744
let mut index = 0;
2845
while index < x_values.len() {
29-
let x = x_values[index];
30-
let exp_part = (-decay_rate * x).exp();
31-
let residual = value_first[index];
32-
gradient[0] += residual;
33-
gradient[1] += residual * exp_part;
34-
gradient[2] += residual * (-amplitude * x * exp_part);
46+
let upstream = value_first[index];
47+
value_grad_at(param, x_values[index], &mut point_grad);
48+
49+
gradient[0] += upstream * point_grad[0];
50+
gradient[1] += upstream * point_grad[1];
51+
gradient[2] += upstream * point_grad[2];
3552
index += 1;
3653
}
3754
}

src/models/exponential_half_life.rs

Lines changed: 28 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -19,27 +19,43 @@ pub(super) fn value_at(param: &[f64], x: f64) -> f64 {
1919
offset + amplitude * exponent.exp()
2020
}
2121

22+
#[inline]
23+
pub(super) fn value_grad_at(param: &[f64], x: f64, grad: &mut [f64]) -> f64 {
24+
debug_assert_eq!(grad.len(), 3);
25+
26+
let offset = param[0];
27+
let amplitude = param[1];
28+
let half_life_raw = param[2];
29+
let (half_life, d_c_raw) = positive_param_with_derivative(half_life_raw);
30+
let exponent = -LN_2 * x / half_life;
31+
let pow = exponent.exp();
32+
let d_model_d_c = amplitude * pow * LN_2 * x / (half_life * half_life);
33+
34+
grad[0] = 1.0;
35+
grad[1] = pow;
36+
grad[2] = d_model_d_c * d_c_raw;
37+
38+
offset + amplitude * pow
39+
}
40+
2241
pub(super) fn add_value_grad(
2342
x_values: &[f64],
2443
param: &[f64],
2544
value_first: &[f64],
2645
gradient: &mut [f64],
2746
) {
28-
let amplitude = param[1];
29-
let half_life_raw = param[2];
30-
let (half_life, d_c_raw) = positive_param_with_derivative(half_life_raw);
47+
debug_assert_eq!(x_values.len(), value_first.len());
48+
debug_assert_eq!(gradient.len(), param.len());
3149

50+
let mut point_grad = [0.0; 3];
3251
let mut index = 0;
3352
while index < x_values.len() {
34-
let x = x_values[index];
35-
let exponent = -LN_2 * x / half_life;
36-
let pow = exponent.exp();
37-
let residual = value_first[index];
38-
let d_model_d_c = amplitude * pow * LN_2 * x / (half_life * half_life);
39-
40-
gradient[0] += residual;
41-
gradient[1] += residual * pow;
42-
gradient[2] += residual * d_model_d_c * d_c_raw;
53+
let upstream = value_first[index];
54+
value_grad_at(param, x_values[index], &mut point_grad);
55+
56+
gradient[0] += upstream * point_grad[0];
57+
gradient[1] += upstream * point_grad[1];
58+
gradient[2] += upstream * point_grad[2];
4359
index += 1;
4460
}
4561
}

0 commit comments

Comments
 (0)