diff --git a/lib/gm_to_dirac/gm_to_dirac_short.cpp b/lib/gm_to_dirac/gm_to_dirac_short.cpp index 55556a7..8671cc7 100644 --- a/lib/gm_to_dirac/gm_to_dirac_short.cpp +++ b/lib/gm_to_dirac/gm_to_dirac_short.cpp @@ -234,7 +234,8 @@ void gm_to_dirac_short::modified_van_mises_distance_sq( GMToDiracConstWeightOptimizationParams optiParams = GMToDiracConstWeightOptimizationParams(covDiag, wXHelper, N, L, bMax, c_b(bMax)); - *distance = modified_van_mises_distance_sq(x, &optiParams); + *distance = modified_van_mises_distance_sq(x, &optiParams) + + constantOffset(&optiParams); } template <> diff --git a/lib/gm_to_dirac/gm_to_dirac_short.h b/lib/gm_to_dirac/gm_to_dirac_short.h index fda9b77..9d5ab70 100644 --- a/lib/gm_to_dirac/gm_to_dirac_short.h +++ b/lib/gm_to_dirac/gm_to_dirac_short.h @@ -124,6 +124,8 @@ class gm_to_dirac_short : public gm_to_dirac_approx_i { GMToDiracConstWeightOptimizationParams* params, double* f, gsl_vector* grad); + static inline double constantOffset(const GMToDiracConstWeightOptimizationParams* params); + static inline void correctMean(gsl_vector* x, const gsl_vector* wX, size_t L, size_t N); diff --git a/lib/gm_to_dirac/gm_to_dirac_short.tpp b/lib/gm_to_dirac/gm_to_dirac_short.tpp index eecbe43..66f1e5b 100644 --- a/lib/gm_to_dirac/gm_to_dirac_short.tpp +++ b/lib/gm_to_dirac/gm_to_dirac_short.tpp @@ -44,11 +44,12 @@ double gm_to_dirac_short::calculateP2(double b, void* params) { const double xikSqrd = x->data[i * N + k] * x->data[i * N + k]; innerSum += xikSqrd / (covDiagSqrd->data[k] + twoBSqrd); } - const double expValuei = wXi * std::exp(-0.5 * innerSum); + const double expm1Valuei = wXi * std::expm1(-0.5 * innerSum); #ifdef USE_CACHE_MANAGER - if (cacheManagerInnerSum) cacheManagerInnerSum->set(b, i, expValuei); + if (cacheManagerInnerSum) + cacheManagerInnerSum->set(b, i, expm1Valuei + wXi); #endif - sum += expValuei; + sum += expm1Valuei; } const double ret = prefactor * sum; @@ -154,8 +155,6 @@ inline void gm_to_dirac_short::calculateD3( const gsl_vector* wX = params->wX; const size_t bMax = params->bMax; gsl_vector* localDist = params->vecN; - const double bMaxSqrd = static_cast(bMax * bMax); - const double bMaxSqrdHalf = bMaxSqrd / 2.00; const double cB = params->cB; (void)cB; @@ -171,22 +170,19 @@ inline void gm_to_dirac_short::calculateD3( const double wXiwXj = wXi * wX->data[j]; double localDistSq = 0.00; if (i == j) { - if (f) d3[(size_t)tid] += bMaxSqrdHalf * wXiwXj; - continue; + continue; // constant-only contribution } for (size_t k = 0; k < N; k++) { localDist->data[k] = x->data[i * N + k] - x->data[j * N + k]; localDistSq += localDist->data[k] * localDist->data[k]; } if (localDistSq <= 0.0) { - if (f) d3[(size_t)tid] += bMaxSqrdHalf * wXiwXj; - continue; + continue; // coincident points: constant-only contribution } const double logLocalDistSq = std::log(localDistSq); if (f) { d3[(size_t)tid] += 0.125 * wXiwXj * - (4.00 * bMaxSqrd + - (logApprox - 2.00 * std::log(bMax)) * localDistSq + + ((logApprox - 2.00 * std::log(bMax)) * localDistSq + (localDistSq * logLocalDistSq)); } if (grad) { @@ -207,6 +203,18 @@ inline void gm_to_dirac_short::calculateD3( } } +template +inline double gm_to_dirac_short::constantOffset( + const GMToDiracConstWeightOptimizationParams* params) { + const double bMax = static_cast(params->bMax); + double w = 0.00; + for (size_t i = 0; i < params->L; ++i) + w += params->wX->data[i]; + + return -2.00 * w * params->twoPiNHalf * params->D1 + + 0.5 * bMax * bMax * w * w; +} + template inline void gm_to_dirac_short::correctMean(gsl_vector* x, const gsl_vector* wX, size_t L, diff --git a/lib/gm_to_dirac/tests/unit_tests/gm_to_dirac_short_test_derivative.cpp b/lib/gm_to_dirac/tests/unit_tests/gm_to_dirac_short_test_derivative.cpp index e546dce..9d27cb5 100644 --- a/lib/gm_to_dirac/tests/unit_tests/gm_to_dirac_short_test_derivative.cpp +++ b/lib/gm_to_dirac/tests/unit_tests/gm_to_dirac_short_test_derivative.cpp @@ -92,9 +92,10 @@ TEST_P( // internal impl double distance_internal = 1; distance_internal = - gm_to_dirac_short::modified_van_mises_distance_sq(x, ¶ms); + gm_to_dirac_short::modified_van_mises_distance_sq(x, ¶ms) + + gm_to_dirac_short::constantOffset(¶ms); - ASSERT_TRUE(distance_wrapper == distance_internal); + ASSERT_DOUBLE_EQ(distance_wrapper, distance_internal); } TEST_P(