@@ -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