From f96167166c017c31ffebe59a83238139e121770f Mon Sep 17 00:00:00 2001 From: Salvatore1021 <150290043+Salvatore1021@users.noreply.github.com> Date: Sat, 8 Aug 2026 01:18:58 +0530 Subject: [PATCH] Wrapped repeating code in a function --- src/stan/callbacks/logger.hpp | 13 +++++++++++ src/stan/model/test_gradients.hpp | 4 ++-- src/stan/services/optimize/bfgs.hpp | 17 ++++---------- src/stan/services/optimize/lbfgs.hpp | 13 ++++------- src/stan/services/sample/standalone_gqs.hpp | 6 ++--- src/stan/services/util/initialize.hpp | 26 ++++++--------------- src/stan/variational/advi.hpp | 3 +-- 7 files changed, 34 insertions(+), 48 deletions(-) diff --git a/src/stan/callbacks/logger.hpp b/src/stan/callbacks/logger.hpp index 6473a9c5e77..5ed019c80ad 100644 --- a/src/stan/callbacks/logger.hpp +++ b/src/stan/callbacks/logger.hpp @@ -95,6 +95,19 @@ class logger { */ virtual void fatal(const std::stringstream& message) {} }; + +inline void log_if_nonempty(logger& log, const std::stringstream& message) { + if (message.str().length() > 0) { + log.info(message); + } +} + +inline void log_if_nonempty(logger& log, const std::string& message) { + if (!message.empty()) { + log.info(message); + } +} + } // namespace callbacks } // namespace stan diff --git a/src/stan/model/test_gradients.hpp b/src/stan/model/test_gradients.hpp index 663ff895a08..80b2821bf32 100644 --- a/src/stan/model/test_gradients.hpp +++ b/src/stan/model/test_gradients.hpp @@ -48,7 +48,7 @@ int test_gradients(const Model& model, std::vector& params_r, double lp = log_prob_grad( model, params_r, params_i, grad, &msg); if (msg.str().length() > 0) { - logger.info(msg); + log_if_nonempty(logger, msg); parameter_writer(msg.str()); } @@ -56,7 +56,7 @@ int test_gradients(const Model& model, std::vector& params_r, finite_diff_grad(model, interrupt, params_r, params_i, grad_fd, epsilon, &msg); if (msg.str().length() > 0) { - logger.info(msg); + log_if_nonempty(logger, msg); parameter_writer(msg.str()); } diff --git a/src/stan/services/optimize/bfgs.hpp b/src/stan/services/optimize/bfgs.hpp index 6f74388eb11..d74bbbeb7a1 100644 --- a/src/stan/services/optimize/bfgs.hpp +++ b/src/stan/services/optimize/bfgs.hpp @@ -102,14 +102,11 @@ int bfgs(Model& model, const stan::io::var_context& init, model.write_array(rng, cont_vector, disc_vector, values, true, true, &msg); } catch (const std::exception& e) { - if (msg.str().length() > 0) { - logger.info(msg); - } + log_if_nonempty(logger, msg); logger.error(e.what()); return error_codes::SOFTWARE; } - if (msg.str().length() > 0) - logger.info(msg); + log_if_nonempty(logger, msg); values.insert(values.begin(), {lp, static_cast(ret)}); parameter_writer(values); @@ -166,8 +163,7 @@ int bfgs(Model& model, const stan::io::var_context& init, &msg); // This if is here to match the pre-refactor behavior - if (msg.str().length() > 0) - logger.info(msg); + log_if_nonempty(logger, msg); values.insert(values.begin(), {lp, static_cast(ret)}); parameter_writer(values); @@ -185,14 +181,11 @@ int bfgs(Model& model, const stan::io::var_context& init, model.write_array(rng, cont_vector, disc_vector, values, true, true, &msg); } catch (const std::exception& e) { - if (msg.str().length() > 0) { - logger.info(msg); - } + log_if_nonempty(logger, msg); logger.error(e.what()); return error_codes::SOFTWARE; } - if (msg.str().length() > 0) - logger.info(msg); + log_if_nonempty(logger, msg); values.insert(values.begin(), {lp, static_cast(ret)}); parameter_writer(values); } diff --git a/src/stan/services/optimize/lbfgs.hpp b/src/stan/services/optimize/lbfgs.hpp index b38f79c1c23..25429e589c5 100644 --- a/src/stan/services/optimize/lbfgs.hpp +++ b/src/stan/services/optimize/lbfgs.hpp @@ -103,8 +103,7 @@ int lbfgs(Model& model, const stan::io::var_context& init, std::vector values; std::stringstream msg; model.write_array(rng, cont_vector, disc_vector, values, true, true, &msg); - if (msg.str().length() > 0) - logger.info(msg); + log_if_nonempty(logger, msg); values.insert(values.begin(), {lp, static_cast(ret)}); parameter_writer(values); @@ -159,8 +158,7 @@ int lbfgs(Model& model, const stan::io::var_context& init, std::stringstream msg; model.write_array(rng, cont_vector, disc_vector, values, true, true, &msg); - if (msg.str().length() > 0) - logger.info(msg); + log_if_nonempty(logger, msg); values.insert(values.begin(), {lp, static_cast(ret)}); parameter_writer(values); @@ -178,14 +176,11 @@ int lbfgs(Model& model, const stan::io::var_context& init, model.write_array(rng, cont_vector, disc_vector, values, true, true, &msg); } catch (const std::exception& e) { - if (msg.str().length() > 0) { - logger.info(msg); - } + log_if_nonempty(logger, msg); logger.error(e.what()); return error_codes::SOFTWARE; } - if (msg.str().length() > 0) - logger.info(msg); + log_if_nonempty(logger, msg); values.insert(values.begin(), {lp, static_cast(ret)}); parameter_writer(values); diff --git a/src/stan/services/sample/standalone_gqs.hpp b/src/stan/services/sample/standalone_gqs.hpp index f7091be981d..7f92e1ff2f4 100644 --- a/src/stan/services/sample/standalone_gqs.hpp +++ b/src/stan/services/sample/standalone_gqs.hpp @@ -75,8 +75,7 @@ int standalone_generate(const Model &model, const Eigen::MatrixXd &draws, try { model.unconstrain_array(row, unconstrained_params_r, &msg); } catch (const std::exception &e) { - if (msg.str().length() > 0) - logger.error(msg); + log_if_nonempty(logger, msg); logger.error(e.what()); return error_codes::DATAERR; } @@ -178,8 +177,7 @@ int standalone_generate(const Model &model, const int num_chains, row = draws[slice_idx].row(i); model.unconstrain_array(row, unconstrained_params_r, &msg); } catch (const std::domain_error &e) { - if (msg.str().length() > 0) - logger.error(msg); + log_if_nonempty(logger, msg); logger.error(e.what()); error_any = true; return; diff --git a/src/stan/services/util/initialize.hpp b/src/stan/services/util/initialize.hpp index c52c331b449..83487de9e81 100644 --- a/src/stan/services/util/initialize.hpp +++ b/src/stan/services/util/initialize.hpp @@ -105,9 +105,7 @@ std::vector initialize(Model& model, const InitContext& init, RNG& rng, model.transform_inits(context, disc_vector, unconstrained, &msg); } } catch (std::domain_error& e) { - if (msg.str().length() > 0) { - logger.info(msg); - } + log_if_nonempty(logger, msg); logger.warn("Rejecting initial value:"); logger.warn( " Error evaluating the log probability" @@ -115,9 +113,7 @@ std::vector initialize(Model& model, const InitContext& init, RNG& rng, logger.warn(e.what()); continue; } catch (std::exception& e) { - if (msg.str().length() > 0) { - logger.info(msg); - } + log_if_nonempty(logger, msg); logger.error( "Unrecoverable error evaluating the log probability" " at the initial value."); @@ -132,12 +128,9 @@ std::vector initialize(Model& model, const InitContext& init, RNG& rng, // the parameters. log_prob = model.template log_prob(unconstrained, disc_vector, &msg); - if (msg.str().length() > 0) { - logger.info(msg); - } + log_if_nonempty(logger, msg); } catch (std::domain_error& e) { - if (msg.str().length() > 0) - logger.info(msg); + log_if_nonempty(logger, msg); logger.warn("Rejecting initial value:"); logger.warn( " Error evaluating the log probability" @@ -145,9 +138,7 @@ std::vector initialize(Model& model, const InitContext& init, RNG& rng, logger.warn(e.what()); continue; } catch (std::exception& e) { - if (msg.str().length() > 0) { - logger.info(msg); - } + log_if_nonempty(logger, msg); logger.error( "Unrecoverable error evaluating the log probability" " at the initial value."); @@ -172,9 +163,7 @@ std::vector initialize(Model& model, const InitContext& init, RNG& rng, log_prob = stan::model::log_prob_grad( model, unconstrained, disc_vector, gradient, &log_prob_msg); } catch (const std::exception& e) { - if (log_prob_msg.str().length() > 0) { - logger.info(log_prob_msg); - } + log_if_nonempty(logger, log_prob_msg); logger.error(e.what()); throw; } @@ -183,8 +172,7 @@ std::vector initialize(Model& model, const InitContext& init, RNG& rng, = std::chrono::duration_cast(end - start) .count() / 1000000.0; - if (log_prob_msg.str().length() > 0) - logger.info(log_prob_msg); + log_if_nonempty(logger, log_prob_msg); bool gradient_ok = std::isfinite(stan::math::sum(gradient)); diff --git a/src/stan/variational/advi.hpp b/src/stan/variational/advi.hpp index 681ce65e82c..45a35368bc1 100644 --- a/src/stan/variational/advi.hpp +++ b/src/stan/variational/advi.hpp @@ -484,8 +484,7 @@ class advi { std::stringstream msg; model_.write_array(rng_, cont_vector, disc_vector, values, true, true, &msg); - if (msg.str().length() > 0) - logger.info(msg); + log_if_nonempty(logger, msg); // The first row of lp_, log_p, and log_g. values.insert(values.begin(), {0, 0, 0});