From 7ba6a45001a252fcc18eb029e3d48496ae4f8407 Mon Sep 17 00:00:00 2001 From: myerspat Date: Thu, 18 May 2023 15:15:25 -0400 Subject: [PATCH] rejection sampling for lin-log interpolation --- src/distribution.cpp | 23 ++++++++++++++++++++++- 1 file changed, 22 insertions(+), 1 deletion(-) diff --git a/src/distribution.cpp b/src/distribution.cpp index 10d45cba1..5c4050763 100644 --- a/src/distribution.cpp +++ b/src/distribution.cpp @@ -304,7 +304,7 @@ void Tabular::init( c_[i] = c_[i - 1] + p_[i - 1] * (x_[i] - x_[i - 1]) / std::log(p_[i] / p_[i - 1]) * (p_[i] / p_[i - 1] - 1); - } else { + } else if (interp_ == Interpolation::log_log) { double m = std::log(p_[i] / p_[i - 1] - x_[i] / x_[i - 1]); c_[i] = c_[i - 1] + p_[i - 1] / (std::pow(x_[i - 1], m) * (m + 1)) * @@ -359,6 +359,27 @@ double Tabular::sample(uint64_t* seed) const (std::sqrt(std::max(0.0, p_i * p_i + 2 * m * (c - c_i))) - p_i) / m; } + } else if (interp_ == Interpolation::lin_log) { + // Linear-log interpolation + // Inverse transform sampling with linear-log interpolation is not possible + // so this uses rejection sampling + double x_i1 = x_[i + 1]; + double p_i1 = p_[i + 1]; + + if (p_i1 == p_i) { + return x_i + (c - c_i) / p_i; + } + + while (true) { + double x = x_i + prn(seed) * (x_i1 - x_i); + double p_unif = p_i + prn(seed) * (p_i1 - p_i); + + // Sample linear-log PDF + double p = p_i + std::log(x / x_i - x_i1 / x_i) * (p_i1 - p_i); + if (p_unif < p) { + return x; + } + } } else if (interp_ == Interpolation::log_lin) { // Log-linear interpolation double x_i1 = x_[i + 1];