Skip to content

Commit 7b02a9c

Browse files
SteveBronderMarton A. Varga
authored andcommitted
Fix Wiener gradient term count narrowing
1 parent 96e5acd commit 7b02a9c

1 file changed

Lines changed: 9 additions & 13 deletions

File tree

stan/math/prim/prob/wiener5_lpdf.hpp

Lines changed: 9 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -542,17 +542,6 @@ inline auto wiener5_grad_w(const T_y& y, const T_a& a, const T_v& v,
542542
const auto n_large_density
543543
= wiener5_density_large_reaction_time_terms(y, a, w, log_error);
544544

545-
const auto n_small_grad
546-
= wiener5_n_terms_small_t<false, GradientCalc::ON>(y, a, w, log_error);
547-
const auto n_large_grad
548-
= wiener5_gradient_large_reaction_time_terms<GradientCalc::ON>(y, a, w,
549-
log_error);
550-
551-
const int n_small = static_cast<int>(
552-
fmax(value_of_rec(n_small_density), value_of_rec(n_small_grad)));
553-
const int n_large = static_cast<int>(
554-
fmax(value_of_rec(n_large_density), value_of_rec(n_large_grad)));
555-
556545
ret_t series_grad_w = 0.0;
557546

558547
if (2.0 * n_small_density <= n_large_density) {
@@ -564,7 +553,10 @@ inline auto wiener5_grad_w(const T_y& y, const T_a& a, const T_v& v,
564553
// dR_s/dw = sum_k (z_k^2 / t* - 1)
565554
// exp(-z_k^2 / (2 t*)).
566555
ret_t max_log = NEGATIVE_INFTY;
567-
556+
const auto n_small_grad
557+
= wiener5_n_terms_small_t<false, GradientCalc::ON>(y, a, w, log_error);
558+
const int n_small = static_cast<int>(
559+
fmax(value_of_rec(n_small_density), value_of_rec(n_small_grad)));
568560
for (int k = -n_small; k <= n_small; ++k) {
569561
const double kd = static_cast<double>(k);
570562
const auto z = q + 2.0 * kd;
@@ -596,9 +588,13 @@ inline auto wiener5_grad_w(const T_y& y, const T_a& a, const T_v& v,
596588
// dR_l/dw = -sum_{k=1}^{K}
597589
// k^2 pi cos(k pi q)
598590
// exp(-(k^2 - 1) pi^2 t* / 2).
591+
const auto n_large_grad
592+
= wiener5_gradient_large_reaction_time_terms<GradientCalc::ON>(
593+
y, a, w, log_error);
594+
const int n_large = static_cast<int>(
595+
fmax(value_of_rec(n_large_density), value_of_rec(n_large_grad)));
599596
ret_t raw = 0.0;
600597
ret_t draw_dw = 0.0;
601-
602598
const auto half_pi2_y = 0.5 * square(pi()) * y_asq;
603599

604600
for (int k = 1; k <= n_large; ++k) {

0 commit comments

Comments
 (0)