diff --git a/linx/background.py b/linx/background.py index 404fce5..4e07c21 100644 --- a/linx/background.py +++ b/linx/background.py @@ -242,7 +242,7 @@ def dY(self, t, Y, args): H = thermo.Hubble(rho_EM + rho_nu + rho_extra) - C_rho_nue, C_rho_numu, _, _ = thermo.collision_terms_std( + C_rho_nue, C_rho_numu, C_n_nue, C_n_numu = thermo.collision_terms_std( T_g, T_nu, T_nu, me=me, decoupled=self.decoupled, use_FD=self.use_FD, collision_me=self.collision_me ) @@ -252,4 +252,16 @@ def dY(self, t, Y, args): dT_g_dt = drho_EM_dt / drho_EM_dT_g dT_nu_dt = drho_nu_dt / drho_nu_dT_nu + #P_e = thermo.rho_nue_std(T_g, me=me) / 3 + #dndt_e = C_n_nue(T_g, T_nu, T_nu, me=me, decoupled=self.decoupled, use_FD=self.use_FD, collision_me=self.collision_me) + + #dmue_dt = -(-3.* H *( (thermo.rho_nue_std+P_e) * thermo.dn_nue_dT_nue_std - thermo.n_nue_std * thermo.drho_nue_dT_nue_std) + thermo.dn_nue_dT_nue_std * thermo.drho_nue_dT_nue_std - thermo.drho_nue_dT_nue_std * dndt_e )/(thermo.dn_nue_dmu_nue_std * thermo.drho_nue_dT_nue_std - thermo.dn_nue_dT_nue_std * thermo.drho_nue_dmu_nue_std) + + + #P_mu = thermo.rho_numt_std(T_nu, me=me) / 3 + #dndt_mu = C_n_numu(T_g, T_nu, T_nu, me=me, decoupled=self.decoupled, use_FD=self.use_FD, collision_me=self.collision_me) + + #dmunu_dt = -(-3.*H*( (thermo.rho_numt_std+P_mu)* thermo.dn_numt_dT_numt_std - thermo.n_numt.std * thermo.drho_numt_dT_numt_std) + thermo.dn_numt_dT_numt_std * thermo.drho_numt_dT_numt_std - thermo.drho_numt_dT_numt_std * dndt_mu )/(thermo.dn_numt_dmu_numt_std * thermo.drho_numt_dT_numt_std - thermo.dn_numt_dT_numt_std * thermo.drho_numt_dmu_numt_std) + return H, dT_g_dt, dT_nu_dt + #return H, dT_g_dt, dT_nu_dt, dmue_dt, dmunu_dt diff --git a/linx/const.py b/linx/const.py index 39d1a75..9fb4f99 100644 --- a/linx/const.py +++ b/linx/const.py @@ -124,4 +124,4 @@ # T_Middle to T_switch: small network always sufficient. T_switch = 1.25e9 * kB # MeV # T_switch to T_end: turn on full network if desired. -T_end = 6.e7 * kB # MeV \ No newline at end of file +T_end = 6.e7 * kB # MeV diff --git a/linx/thermo.py b/linx/thermo.py index ffb94e6..9ca2f56 100644 --- a/linx/thermo.py +++ b/linx/thermo.py @@ -1023,13 +1023,13 @@ def collision_terms_std( """ - f_n, f_a, f_s = lax.cond( - decoupled, - lambda _: (0., 0., 0.), + f_n, f_a, f_s, f_a3, f_n3 = lax.cond( + decoupled, + lambda _: (0., 0., 0., 0., 0.), lambda _: lax.cond( - use_FD, lambda _: (0.852, 0.884, 0.829), - lambda _: (1., 1., 1.), 0. - ), + use_FD, lambda _: (0.852, 0.884, 0.829, 0.031837, 0.041088), + lambda _: (1., 1., 1., 0., 0.), 0. + ), 0. ) @@ -1045,8 +1045,19 @@ def G(T_1, mu_1, T_2, mu_2): T_1**9 * jnp.exp(2 * mu_1 / T_1) - T_2**9 * jnp.exp(2 * mu_2 / T_2) ) - + 56 * f_s * jnp.exp(2 * mu_1 / T_1) * jnp.exp(2 * mu_2 / T_2) *( - T_1**4 * T_2**4 * (T_1 - T_2) + + 32 * f_a3 * (T_1*T_2)**4.5 * (jnp.exp(2 * mu_2 / T_2) - jnp.exp(2 * mu_1 / T_1)) + + 56 * f_s * jnp.exp(mu_1 / T_1) * jnp.exp(mu_2 / T_2) *(T_1**4 * T_2**4 * (T_1 - T_2)) + ) + + def N(T_1, mu_1, T_2, mu_2): + + return ( + f_n * ( + T_1**8 * jnp.exp(2 * mu_1 / T_1) + - T_2**8 * jnp.exp(2 * mu_2 / T_2) + ) + + f_n3 * (T_1*T_2)**4 * ( + jnp.exp(2 * mu_2 / T_2) - jnp.exp(2 * mu_1 / T_1) ) ) @@ -1064,6 +1075,18 @@ def interp_fa2(f_tab): me/T_1, f_tab[:,0], f_tab[:,index], left=f_tab[0,index], right=f_tab[-1,index] ) + def interp_fa3(f_tab): + index = 3 + return jnp.interp( + me/T_1, f_tab[:,0], f_tab[:,index], left=f_tab[0,index], right=f_tab[-1,index] + ) + + def interp_fa4(f_tab): + index = 4 + return jnp.interp( + me/T_1, f_tab[:,0], f_tab[:,index], left=f_tab[0,index], right=f_tab[-1,index] + ) + def interp_fs1(f_tab): index = 5 return jnp.interp( @@ -1075,25 +1098,8 @@ def interp_fs2(f_tab): return jnp.interp( me/T_1, f_tab[:,0], f_tab[:,index], left=f_tab[0,index], right=f_tab[-1,index] ) -# def interp_f(f_tab): -# # Tables have boundary values 0.0 (low T) and 1.0 (high T) -# return interpax.interp1d( -# T_1, f_tab[:,0], f_tab[:,1], extrap=(0.0, 1.0) -# ) - - - # def interp_f(f_tab): - - # return jnp.interp( - # T_1, f_tab[:,0], f_tab[:,1], left=f_tab[0,1], right=f_tab[-1,1] - # ) - - # f_nue_ann = lax.cond( - # collision_me, interp_f, lambda _: 1., f_nue_ann_tab - # ) - # f_nue_scat = lax.cond( - # collision_me, interp_f, lambda _: 1., f_nue_scat_tab - # ) + + f_ann_1 = lax.cond( collision_me, interp_fa1, lambda _: 1., f_coeffs @@ -1108,25 +1114,86 @@ def interp_fs2(f_tab): f_scat_2 = lax.cond( collision_me, interp_fs2, lambda _: 1., f_coeffs ) + f_ann_3 = lax.cond( + collision_me, interp_fa3, lambda _: 0., f_coeffs + ) + f_ann_4 = lax.cond( + collision_me, interp_fa4, lambda _: 0., f_coeffs + ) return ( # note f_a and f_s are now folded into f_nue_ann/scat 4 * (geL**2 + geR**2) * (32 * f_ann_1 * ( T_1**9 * jnp.exp(2 * mu_1 / T_1) - T_2**9 * jnp.exp(2 * mu_2 / T_2) - ) - + 56 * f_scat_1 * ( - jnp.exp(2 * mu_1 / T_1) * jnp.exp(2 * mu_2 / T_2) - * T_1**4 * T_2**4 * (T_1 - T_2) - ) + ) + + 32 * f_ann_3 * (T_1*T_2)**4.5 * (jnp.exp(2 * mu_2 / T_2) - jnp.exp(2 * mu_1 / T_1)) + + 56 * f_scat_1 * (jnp.exp(mu_1 / T_1) * jnp.exp(mu_2 / T_2)* T_1**4 * T_2**4 * (T_1 - T_2)) ) # new terms (previously baked into tabulated rates) + 4 * geL*geR * (f_ann_2 * 32 * ( T_1**9 * jnp.exp(2 * mu_1 / T_1) - T_2**9 * jnp.exp(2 * mu_2 / T_2) ) - + 56 * f_scat_2 * ( - jnp.exp(2 * mu_1 / T_1) * jnp.exp(2 * mu_2 / T_2) - * T_1**4 * T_2**4 * (T_1 - T_2) + + 32 * f_ann_4 * (T_1*T_2)**4.5 * (jnp.exp(2*mu_2/T_2) - jnp.exp(2*mu_1/T_1)) + + 56 * f_scat_2 * (jnp.exp(mu_1/T_1) * jnp.exp(mu_2/T_2) * T_1**4 * T_2**4 * (T_1-T_2)) + ) + ) + def N_nue_with_me(T_1, mu_1, T_2, mu_2, me): + + def interp_fn1(f_tab): + index = 9 + return jnp.interp( + me/T_1, f_tab[:,0], f_tab[:,index], left=f_tab[0,index], right=f_tab[-1,index] + ) + + def interp_fn2(f_tab): + index = 10 + return jnp.interp( + me/T_1, f_tab[:,0], f_tab[:,index], left=f_tab[0,index], right=f_tab[-1,index] + ) + + def interp_fn3(f_tab): + index = 11 + return jnp.interp( + me/T_1, f_tab[:,0], f_tab[:,index], left=f_tab[0,index], right=f_tab[-1,index] + ) + + def interp_fn4(f_tab): + index = 12 + return jnp.interp( + me/T_1, f_tab[:,0], f_tab[:,index], left=f_tab[0,index], right=f_tab[-1,index] + ) + + f_num_1 = lax.cond( + collision_me, interp_fn1, lambda _: 1., f_coeffs + ) + f_num_2 = lax.cond( + collision_me, interp_fn2, lambda _: 0., f_coeffs + ) + f_num_3 = lax.cond( + collision_me, interp_fn3, lambda _: 0., f_coeffs + ) + f_num_4 = lax.cond( + collision_me, interp_fn4, lambda _: 0., f_coeffs + ) + + return ( + 4 * (geL**2 + geR**2) * ( + f_num_1 * ( + T_1**8 * jnp.exp(2 * mu_1 / T_1) + - T_2**8 * jnp.exp(2 * mu_2 / T_2) + ) + + f_num_3 * (T_1*T_2)**4 * ( + jnp.exp(2 * mu_2 / T_2) - jnp.exp(2 * mu_1 / T_1) + ) + ) + + 4 * geL*geR * ( + f_num_2 * ( + T_1**8 * jnp.exp(2 * mu_1 / T_1) + - T_2**8 * jnp.exp(2 * mu_2 / T_2) + ) + + f_num_4 * (T_1*T_2)**4 * ( + jnp.exp(2 * mu_2 / T_2) - jnp.exp(2 * mu_1 / T_1) ) ) ) @@ -1144,6 +1211,17 @@ def interp_fa2(f_tab): return jnp.interp( me/T_1, f_tab[:,0], f_tab[:,index], left=f_tab[0,index], right=f_tab[-1,index] ) + def interp_fa3(f_tab): + index = 3 + return jnp.interp( + me/T_1, f_tab[:,0], f_tab[:,index], left=f_tab[0,index], right=f_tab[-1,index] + ) + + def interp_fa4(f_tab): + index = 4 + return jnp.interp( + me/T_1, f_tab[:,0], f_tab[:,index], left=f_tab[0,index], right=f_tab[-1,index] + ) def interp_fs1(f_tab): index = 5 @@ -1155,26 +1233,7 @@ def interp_fs2(f_tab): index = 6 return jnp.interp( me/T_1, f_tab[:,0], f_tab[:,index], left=f_tab[0,index], right=f_tab[-1,index]) -# def G_numt_with_me(T_1, mu_1, T_2, mu_2): - -# def interp_f(f_tab): -# # Tables have boundary values 0.0 (low T) and 1.0 (high T) -# return interpax.interp1d( -# T_1, f_tab[:,0], f_tab[:,1], extrap=(0.0, 1.0) -# ) - # def interp_f(f_tab): - - # return jnp.interp( - # T_1, f_tab[:,0], f_tab[:,1], left=f_tab[0,1], right=f_tab[-1,1] - # ) - - # f_numt_ann = lax.cond( - # collision_me, interp_f, lambda _: 1., f_numu_ann_tab - # ) - # f_numt_scat = lax.cond( - # collision_me, interp_f, lambda _: 1., f_numu_scat_tab - # ) f_ann_1 = lax.cond( collision_me, interp_fa1, lambda _: 1., f_coeffs @@ -1189,29 +1248,89 @@ def interp_fs2(f_tab): f_scat_2 = lax.cond( collision_me, interp_fs2, lambda _: 1., f_coeffs ) - + f_ann_3 = lax.cond( + collision_me, interp_fa3, lambda _: 0., f_coeffs + ) + f_ann_4 = lax.cond( + collision_me, interp_fa4, lambda _: 0., f_coeffs + ) return ( # f_s, f_a now folded into f_ann and f_scat 4 * (gmuL**2 + gmuR**2) * (32 * f_ann_1 * ( T_1**9 * jnp.exp(2 * mu_1 / T_1) - T_2**9 * jnp.exp(2 * mu_2 / T_2) ) - + 56 * f_scat_1 * ( - jnp.exp(2 * mu_1 / T_1) * jnp.exp(2 * mu_2 / T_2) - * T_1**4 * T_2**4 * (T_1 - T_2) - ) + + 32 * f_ann_3 * (T_1*T_2)**4.5 * (jnp.exp(2 * mu_2 / T_2) - jnp.exp(2 * mu_1 / T_1)) + + 56 * f_scat_1 * (jnp.exp(mu_1 / T_1) * jnp.exp(mu_2 / T_2)* T_1**4 * T_2**4 * (T_1 - T_2)) ) # new terms (previously baked into tabulated rates) + 4 * gmuL*gmuR * (f_ann_2 * 32 * ( T_1**9 * jnp.exp(2 * mu_1 / T_1) - T_2**9 * jnp.exp(2 * mu_2 / T_2) ) - + 56 * f_scat_2 * ( - jnp.exp(2 * mu_1 / T_1) * jnp.exp(2 * mu_2 / T_2) - * T_1**4 * T_2**4 * (T_1 - T_2) - ) + + 32 * f_ann_4 * (T_1*T_2)**4.5 * (jnp.exp(2 * mu_2 / T_2) - jnp.exp(2 * mu_1 / T_1)) + + 56 * f_scat_2 * (jnp.exp(mu_1 / T_1) * jnp.exp(mu_2 / T_2) * T_1**4 * T_2**4 * (T_1 - T_2)) ) ) + + def N_numt_with_me(T_1, mu_1, T_2, mu_2, me): + def interp_fn1(f_tab): + index = 9 + return jnp.interp( + me/T_1, f_tab[:,0], f_tab[:,index], left=f_tab[0,index], right=f_tab[-1,index] + ) + + def interp_fn2(f_tab): + index = 10 + return jnp.interp( + me/T_1, f_tab[:,0], f_tab[:,index], left=f_tab[0,index], right=f_tab[-1,index] + ) + + def interp_fn3(f_tab): + index = 11 + return jnp.interp( + me/T_1, f_tab[:,0], f_tab[:,index], left=f_tab[0,index], right=f_tab[-1,index] + ) + + def interp_fn4(f_tab): + index = 12 + return jnp.interp( + me/T_1, f_tab[:,0], f_tab[:,index], left=f_tab[0,index], right=f_tab[-1,index] + ) + + f_num_1 = lax.cond( + collision_me, interp_fn1, lambda _: 1., f_coeffs + ) + f_num_2 = lax.cond( + collision_me, interp_fn2, lambda _: 0., f_coeffs + ) + f_num_3 = lax.cond( + collision_me, interp_fn3, lambda _: 0., f_coeffs + ) + f_num_4 = lax.cond( + collision_me, interp_fn4, lambda _: 0., f_coeffs + ) + + return ( + 4 * (gmuL**2 + gmuR**2) * ( + f_num_1 * ( + T_1**8 * jnp.exp(2 * mu_1 / T_1) + - T_2**8 * jnp.exp(2 * mu_2 / T_2) + ) + + f_num_3 * (T_1*T_2)**4 * ( + jnp.exp(2 * mu_2 / T_2) - jnp.exp(2 * mu_1 / T_1) + ) + ) + + 4 * gmuL*gmuR * ( + f_num_2 * ( + T_1**8 * jnp.exp(2 * mu_1 / T_1) + - T_2**8 * jnp.exp(2 * mu_2 / T_2) + ) + + f_num_4 * (T_1*T_2)**4 * ( + jnp.exp(2 * mu_2 / T_2) - jnp.exp(2 * mu_1 / T_1) + ) + ) + ) # Units MeV^4 s^-1 C_rho_nue = const.GF**2 / jnp.pi**5 * ( # prev coeff now in G def G_nue_with_me(T_g, 0., T_nue, mu_nue, me) @@ -1225,23 +1344,25 @@ def interp_fs2(f_tab): ) / const.hbar # Units MeV^3 s^-1 - C_n_nue = 8 * f_n * const.GF**2 / jnp.pi**5 * ( - 4 * (geL**2 + geR**2) - * (T_g**8 - T_nue**8 * jnp.exp(2 * mu_nue / T_nue)) - + 2 * ( - T_numt**8 * jnp.exp(2 * mu_numt / T_numt) - - T_nue**8 * jnp.exp(2 * mu_nue / T_nue) - ) - ) / const.hbar + C_n_nue = lax.cond( + decoupled, + lambda _: 0., + lambda _: 8 * const.GF**2 / jnp.pi**5 * ( + N_nue_with_me(T_g, 0., T_nue, mu_nue, me) + + 2 * N(T_numt, mu_numt, T_nue, mu_nue) + ) / const.hbar, + 0. + ) # Units MeV^3 s^-1 - C_n_numu = 8 * f_n * const.GF**2 / jnp.pi**5 * ( - 4 * (gmuL**2 + gmuR**2) - * (T_g**8 - T_nue**8 * jnp.exp(2 * mu_numt / T_numt)) - - ( - T_numt**8 * jnp.exp(2 * mu_numt / T_numt) - - T_nue**8 * jnp.exp(2 * mu_nue / T_nue) - ) - ) / const.hbar + C_n_numu = lax.cond( + decoupled, + lambda _: 0., + lambda _: 8 * const.GF**2 / jnp.pi**5 * ( + N_numt_with_me(T_g, 0., T_numt, mu_numt, me) + - N(T_numt, mu_numt, T_nue, mu_nue) + ) / const.hbar, + 0. + ) return C_rho_nue, C_rho_numu, C_n_nue, C_n_numu \ No newline at end of file