From cebc58bcc3528224851dc3103fb20753546e4732 Mon Sep 17 00:00:00 2001 From: Alex Alarcon Date: Tue, 5 May 2026 18:42:35 +0200 Subject: [PATCH 1/2] Fixing behaviour of satquench model, which was currently not having an effect --- diffstar/diffstarpop/kernels/diffstarpop_mgash.py | 15 +++++++++++---- 1 file changed, 11 insertions(+), 4 deletions(-) diff --git a/diffstar/diffstarpop/kernels/diffstarpop_mgash.py b/diffstar/diffstarpop/kernels/diffstarpop_mgash.py index 6554d16..1892b68 100644 --- a/diffstar/diffstarpop/kernels/diffstarpop_mgash.py +++ b/diffstar/diffstarpop/kernels/diffstarpop_mgash.py @@ -96,14 +96,21 @@ def _diffstarpop_means_covs( means_covs = _sfh_pdf_scalar_kernel(sfh_pdf_cens_params, logmp0, tpeak) # Modify frac_q for satellites - frac_q = means_covs[0] + ( + frac_quench_cen, + frac_quench_sat, + ) = means_covs[:2] satquench_params = DEFAULT_SATQUENCHPOP_PARAMS._make( [getattr(diffstarpop_params, x) for x in DEFAULT_SATQUENCHPOP_PARAMS._fields] ) - frac_q = get_qprob_sat( - satquench_params, lgmu_infall, logmhost_infall, gyr_since_infall, frac_q + frac_quench_sat_updated = get_qprob_sat( + satquench_params, + lgmu_infall, + logmhost_infall, + gyr_since_infall, + frac_quench_sat, ) - means_covs = (frac_q, *means_covs[1:]) + means_covs = (frac_quench_cen, frac_quench_sat_updated, *means_covs[2:]) return means_covs From 24ca0723a26ab7254325ed45bbea9dfaff7aaba7 Mon Sep 17 00:00:00 2001 From: Alex Alarcon Date: Tue, 5 May 2026 20:25:20 +0200 Subject: [PATCH 2/2] Updating test_gradients to correctly test the satquench params --- diffstar/diffstarpop/tests/test_gradients.py | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/diffstar/diffstarpop/tests/test_gradients.py b/diffstar/diffstarpop/tests/test_gradients.py index 3434ac1..f8a0b64 100644 --- a/diffstar/diffstarpop/tests/test_gradients.py +++ b/diffstar/diffstarpop/tests/test_gradients.py @@ -202,15 +202,15 @@ def test_gradients_of_diffstarpop_pdf_satquench_params_are_nonzero(): gyr_since_infall, ) _res = _diffstarpop_means_covs(*args) - frac_quench = _res[0] ( + frac_quench_cen, frac_quench_sat, mu_mseq, mu_qseq, cov_mseq_ms_block, cov_qseq_ms_block, cov_qseq_q_block, - ) = _res[1:] + ) = _res # Generate an alternate galpop at some other point in param space ran_params_key, ran_key = jran.split(ran_key, 2) @@ -225,22 +225,22 @@ def test_gradients_of_diffstarpop_pdf_satquench_params_are_nonzero(): gyr_since_infall, ) _res = _diffstarpop_means_covs(*args) - frac_quench2 = _res[0] ( + frac_quench_cen2, frac_quench_sat2, mu_mseq2, mu_qseq2, cov_mseq_ms_block2, cov_qseq_ms_block2, cov_qseq_q_block2, - ) = _res[1:] + ) = _res assert not np.allclose(mu_mseq2, mu_mseq) assert not np.allclose(cov_qseq_ms_block2, cov_qseq_ms_block) assert not np.allclose(cov_qseq_q_block2, cov_qseq_q_block) - assert not np.allclose(frac_quench2, frac_quench) + assert not np.allclose(frac_quench_sat, frac_quench_sat2) - frac_q_target = np.copy(frac_quench) + frac_q_target_sat = np.copy(frac_quench_sat) @jjit def _loss(u_params): @@ -254,8 +254,8 @@ def _loss(u_params): gyr_since_infall, ) _res = _diffstarpop_means_covs(*args) - frac_q_pred = _res[0] - return _mse(frac_q_pred, frac_q_target) + frac_q_pred_sat = _res[1] + return _mse(frac_q_pred_sat, frac_q_target_sat) frac_q_loss, frac_q_grads = value_and_grad(_loss)(alt_dpp_u_params) assert np.isfinite(frac_q_loss)