Skip to content

Commit 8d072a6

Browse files
authored
Merge pull request #116 from ArgonneCPAC/fixing_satquench
Fixing satquench model
2 parents c6e0028 + 24ca072 commit 8d072a6

2 files changed

Lines changed: 19 additions & 12 deletions

File tree

diffstar/diffstarpop/kernels/diffstarpop_mgash.py

Lines changed: 11 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -96,14 +96,21 @@ def _diffstarpop_means_covs(
9696
means_covs = _sfh_pdf_scalar_kernel(sfh_pdf_cens_params, logmp0, tpeak)
9797

9898
# Modify frac_q for satellites
99-
frac_q = means_covs[0]
99+
(
100+
frac_quench_cen,
101+
frac_quench_sat,
102+
) = means_covs[:2]
100103
satquench_params = DEFAULT_SATQUENCHPOP_PARAMS._make(
101104
[getattr(diffstarpop_params, x) for x in DEFAULT_SATQUENCHPOP_PARAMS._fields]
102105
)
103-
frac_q = get_qprob_sat(
104-
satquench_params, lgmu_infall, logmhost_infall, gyr_since_infall, frac_q
106+
frac_quench_sat_updated = get_qprob_sat(
107+
satquench_params,
108+
lgmu_infall,
109+
logmhost_infall,
110+
gyr_since_infall,
111+
frac_quench_sat,
105112
)
106-
means_covs = (frac_q, *means_covs[1:])
113+
means_covs = (frac_quench_cen, frac_quench_sat_updated, *means_covs[2:])
107114
return means_covs
108115

109116

diffstar/diffstarpop/tests/test_gradients.py

Lines changed: 8 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -202,15 +202,15 @@ def test_gradients_of_diffstarpop_pdf_satquench_params_are_nonzero():
202202
gyr_since_infall,
203203
)
204204
_res = _diffstarpop_means_covs(*args)
205-
frac_quench = _res[0]
206205
(
206+
frac_quench_cen,
207207
frac_quench_sat,
208208
mu_mseq,
209209
mu_qseq,
210210
cov_mseq_ms_block,
211211
cov_qseq_ms_block,
212212
cov_qseq_q_block,
213-
) = _res[1:]
213+
) = _res
214214

215215
# Generate an alternate galpop at some other point in param space
216216
ran_params_key, ran_key = jran.split(ran_key, 2)
@@ -225,22 +225,22 @@ def test_gradients_of_diffstarpop_pdf_satquench_params_are_nonzero():
225225
gyr_since_infall,
226226
)
227227
_res = _diffstarpop_means_covs(*args)
228-
frac_quench2 = _res[0]
229228
(
229+
frac_quench_cen2,
230230
frac_quench_sat2,
231231
mu_mseq2,
232232
mu_qseq2,
233233
cov_mseq_ms_block2,
234234
cov_qseq_ms_block2,
235235
cov_qseq_q_block2,
236-
) = _res[1:]
236+
) = _res
237237

238238
assert not np.allclose(mu_mseq2, mu_mseq)
239239
assert not np.allclose(cov_qseq_ms_block2, cov_qseq_ms_block)
240240
assert not np.allclose(cov_qseq_q_block2, cov_qseq_q_block)
241-
assert not np.allclose(frac_quench2, frac_quench)
241+
assert not np.allclose(frac_quench_sat, frac_quench_sat2)
242242

243-
frac_q_target = np.copy(frac_quench)
243+
frac_q_target_sat = np.copy(frac_quench_sat)
244244

245245
@jjit
246246
def _loss(u_params):
@@ -254,8 +254,8 @@ def _loss(u_params):
254254
gyr_since_infall,
255255
)
256256
_res = _diffstarpop_means_covs(*args)
257-
frac_q_pred = _res[0]
258-
return _mse(frac_q_pred, frac_q_target)
257+
frac_q_pred_sat = _res[1]
258+
return _mse(frac_q_pred_sat, frac_q_target_sat)
259259

260260
frac_q_loss, frac_q_grads = value_and_grad(_loss)(alt_dpp_u_params)
261261
assert np.isfinite(frac_q_loss)

0 commit comments

Comments
 (0)