@@ -128,6 +128,9 @@ def _loss_kern_1d(
128128 ssp_halpha_luminosity = None ,
129129 lg_halpha_LF_target = None ,
130130 lg_halpha_Lbin_edges = None ,
131+ halpha_LF_z = None ,
132+ halpha_LF_delta_z = None ,
133+ halpha_LF_delta_z_vol_Mpc3 = None ,
131134):
132135 # The if structure below assumes that if len(u_theta)==1, then it is just diffstarpop params
133136 if len (u_theta ) == 3 :
@@ -278,6 +281,9 @@ def fit_n_1d(
278281 ssp_halpha_luminosity = None ,
279282 lg_halpha_LF_target = None ,
280283 lg_halpha_Lbin_edges = None ,
284+ halpha_LF_z = None ,
285+ halpha_LF_delta_z = None ,
286+ halpha_LF_delta_z_vol_Mpc3 = None ,
281287):
282288 opt_init , opt_update , get_params = jax_opt .adam (step_size )
283289 opt_state = opt_init (u_theta_init )
@@ -311,6 +317,9 @@ def fit_n_1d(
311317 ssp_halpha_luminosity ,
312318 lg_halpha_LF_target ,
313319 lg_halpha_Lbin_edges ,
320+ halpha_LF_z ,
321+ halpha_LF_delta_z ,
322+ halpha_LF_delta_z_vol_Mpc3 ,
314323 )
315324
316325 def _opt_update (opt_state , i ):
@@ -357,6 +366,9 @@ def _opt_update(opt_state, i):
357366 None ,
358367 0 ,
359368 0 ,
369+ 0 ,
370+ None ,
371+ 0 ,
360372)
361373_loss_kern_1d_multi_z = jjit (
362374 vmap (
@@ -408,6 +420,9 @@ def fit_n_1d_multi_z(
408420 ssp_halpha_luminosity = None ,
409421 lg_halpha_LF_target = None ,
410422 lg_halpha_Lbin_edges = None ,
423+ halpha_LF_z = None ,
424+ halpha_LF_delta_z = None ,
425+ halpha_LF_delta_z_vol_Mpc3 = None ,
411426):
412427 opt_init , opt_update , get_params = jax_opt .adam (step_size )
413428 opt_state = opt_init (u_theta_init )
@@ -441,6 +456,9 @@ def fit_n_1d_multi_z(
441456 ssp_halpha_luminosity ,
442457 lg_halpha_LF_target ,
443458 lg_halpha_Lbin_edges ,
459+ halpha_LF_z ,
460+ halpha_LF_delta_z ,
461+ halpha_LF_delta_z_vol_Mpc3 ,
444462 )
445463
446464 def _opt_update (opt_state , i ):
0 commit comments