diff --git a/README.md b/README.md index 8ea2e25..28d9ef1 100644 --- a/README.md +++ b/README.md @@ -26,9 +26,17 @@ $ conda create -c conda-forge -n diffit python=3.11 numpy numba flake8 pytest ja Data for this project can be found [at this URL](https://portal.nersc.gov/project/hacc/aphearin/diffstar_data/). ## Scripts and demo notebooks + +The `demo_diffstar_sfh.ipynb` notebook in the `docs` folder illustrates how to use the Diffstar model for individual SFH, and how to generate the SFH of a population of galaxies using DiffstarPop. + +The `demo_diffmahpop_diffstarpop_sfh.ipynb` notebook illustrates how to generate a subhalo catalog, and how to generate SFHs for each halo using parameters that reproduce UniverseMachine, IllustrisTNG or Galacticus. + +See `diffstar_fitting_script_umachine_mgash.py` for an example of how to fit the SFHs of a large number of simulated galaxies in parallel with mpi4py. + The `diffstar_fitter_demo.ipynb` notebook demonstrates how to fit the SFH of a simulated galaxy with a diffstar approximation. -See `history_fitting_script.py` for an example of how to fit the SFHs of a large number of simulated galaxies in parallel with mpi4py. + +See `fit_mstar_ssfr_pdfs_mgash.py` for an example of how to use DiffstarPop to fit a set of Mstar and sSFR PDFs, and `measure_smhm_smdpl_script_mpi_mgash.py` for an example of how to generate the target data from a set of Diffstar fits. ## Citing diffstar [The Diffstar paper](https://arxiv.org/abs/2205.04273) has been published in [Monthly Notices of the Royal Astronomical Society](https://academic.oup.com/mnras/article-abstract/518/1/562/6795944?redirectedFrom=fulltext). Citation information for the paper can be found at [this ADS link](https://ui.adsabs.harvard.edu/abs/2023MNRAS.518..562A/abstract), copied below for convenience: diff --git a/diffstar/diffstarpop/kernels/__init__.py b/diffstar/diffstarpop/kernels/__init__.py index 93d911c..b1308d4 100644 --- a/diffstar/diffstarpop/kernels/__init__.py +++ b/diffstar/diffstarpop/kernels/__init__.py @@ -1,4 +1,3 @@ -""" -""" +""" """ # flake8: noqa diff --git a/diffstar/diffstarpop/kernels/params/__init__.py b/diffstar/diffstarpop/kernels/params/__init__.py index cc1cc5b..584d36c 100644 --- a/diffstar/diffstarpop/kernels/params/__init__.py +++ b/diffstar/diffstarpop/kernels/params/__init__.py @@ -7,3 +7,8 @@ DiffstarPop_Params_Diffstarfits_mgash, DiffstarPop_UParams_Diffstarfits_mgash, ) + +from .params_diffstarpopfits_mgash import ( + DiffstarPop_Params_Diffstarpopfits_mgash, + DiffstarPop_UParams_Diffstarpopfits_mgash, +) diff --git a/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_galacticus_in_plus_ex_situ.py b/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_galacticus_in_plus_ex_situ.py index e6c65b6..9176a3a 100644 --- a/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_galacticus_in_plus_ex_situ.py +++ b/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_galacticus_in_plus_ex_situ.py @@ -1,74 +1,82 @@ -"""""" - from collections import OrderedDict, namedtuple -from ..defaults_mgash import DiffstarPopParams, get_unbounded_diffstarpop_params +from ..defaults_mgash import ( + DiffstarPopParams, + get_unbounded_diffstarpop_params, +) + from ..satquenchpop_model import DEFAULT_SATQUENCHPOP_PARAMS SFH_PDF_QUENCH_MU_PDICT = OrderedDict( [ - ("mean_ulgm_mseq_xtp", 12.780), - ("mean_ulgm_mseq_ytp", 11.668), - ("mean_ulgm_mseq_lo", 0.748), + ("mean_ulgm_mseq_xtp", 12.029), + ("mean_ulgm_mseq_ytp", 11.300), + ("mean_ulgm_mseq_lo", 2.812), ("mean_ulgm_mseq_hi", 0.400), - ("mean_ulgy_mseq_int", -9.299), - ("mean_ulgy_mseq_slp", 0.409), + ("mean_ulgy_mseq_xtp", 12.344), + ("mean_ulgy_mseq_ytp", -9.296), + ("mean_ulgy_mseq_lo", 0.902), + ("mean_ulgy_mseq_hi", 0.303), ("mean_ul_mseq_int", -0.639), ("mean_ul_mseq_slp", 0.988), - ("mean_uh_mseq_int", -0.607), - ("mean_uh_mseq_slp", -0.312), - ("mean_ulgm_qseq_xtp", 13.252), - ("mean_ulgm_qseq_ytp", 11.971), - ("mean_ulgm_qseq_lo", 0.639), - ("mean_ulgm_qseq_hi", 0.208), - ("mean_ulgy_qseq_int", -9.243), - ("mean_ulgy_qseq_slp", 0.660), - ("mean_ul_qseq_int", -0.758), - ("mean_ul_qseq_slp", 0.421), - ("mean_uh_qseq_int", -0.661), - ("mean_uh_qseq_slp", -0.334), - ("mean_uqt_int", 1.065), - ("mean_uqt_slp", -0.003), - ("mean_uqs_int", -0.352), - ("mean_uqs_slp", -0.273), - ("mean_udrop_int", -1.938), - ("mean_udrop_slp", 0.551), - ("mean_urej_int", -0.620), - ("mean_urej_slp", -0.206), + ("mean_uh_mseq_int", -0.391), + ("mean_uh_mseq_slp", 0.161), + ("mean_ulgm_qseq_xtp", 13.078), + ("mean_ulgm_qseq_ytp", 11.903), + ("mean_ulgm_qseq_lo", 0.794), + ("mean_ulgm_qseq_hi", 0.230), + ("mean_ulgy_qseq_xtp", 12.112), + ("mean_ulgy_qseq_ytp", -9.421), + ("mean_ulgy_qseq_lo", 0.725), + ("mean_ulgy_qseq_hi", 0.356), + ("mean_ul_qseq_int", -0.714), + ("mean_ul_qseq_slp", 0.616), + ("mean_uh_qseq_int", -0.582), + ("mean_uh_qseq_slp", -0.115), + ("mean_uqt_xtp", 12.922), + ("mean_uqt_ytp", 1.047), + ("mean_uqt_lo", -0.100), + ("mean_uqt_hi", -0.100), + ("mean_uqs_int", -0.298), + ("mean_uqs_slp", -0.411), + ("mean_udrop_int", -1.953), + ("mean_udrop_slp", 0.586), + ("mean_urej_int", -0.853), + ("mean_urej_slp", 0.329), ] ) SFH_PDF_QUENCH_COV_MS_BLOCK_PDICT = OrderedDict( [ - ("std_ulgm_mseq_int", 0.355), - ("std_ulgm_mseq_slp", -0.015), - ("std_ulgy_mseq_int", 0.277), - ("std_ulgy_mseq_slp", 0.103), - ("std_ul_mseq_int", 2.221), - ("std_ul_mseq_slp", 0.900), - ("std_uh_mseq_int", 0.651), - ("std_uh_mseq_slp", -0.155), - ("std_ulgm_qseq_int", 0.313), - ("std_ulgm_qseq_slp", -0.022), - ("std_ulgy_qseq_int", 0.257), - ("std_ulgy_qseq_slp", 0.075), - ("std_ul_qseq_int", 2.241), - ("std_ul_qseq_slp", -0.130), - ("std_uh_qseq_int", 0.482), - ("std_uh_qseq_slp", 0.007), + ("std_ulgm_mseq_int", 0.552), + ("std_ulgm_mseq_slp", -0.433), + ("std_ulgy_mseq_int", 0.298), + ("std_ulgy_mseq_slp", 0.046), + ("std_ul_mseq_int", 2.469), + ("std_ul_mseq_slp", 0.677), + ("std_uh_mseq_int", 0.705), + ("std_uh_mseq_slp", -0.287), + ("std_ulgm_qseq_int", 0.492), + ("std_ulgm_qseq_slp", -0.293), + ("std_ulgy_qseq_int", 0.264), + ("std_ulgy_qseq_slp", 0.063), + ("std_ul_qseq_int", 2.263), + ("std_ul_qseq_slp", -0.165), + ("std_uh_qseq_int", 0.528), + ("std_uh_qseq_slp", -0.067), ] ) SFH_PDF_QUENCH_COV_Q_BLOCK_PDICT = OrderedDict( [ - ("std_uqt_int", 0.067), - ("std_uqt_slp", 0.013), + ("std_uqt_int", 0.075), + ("std_uqt_slp", -0.008), ("std_uqs_int", 0.900), ("std_uqs_slp", 0.900), - ("std_udrop_int", 0.589), - ("std_udrop_slp", 0.049), - ("std_urej_int", 0.867), - ("std_urej_slp", 0.049), + ("std_udrop_int", 0.600), + ("std_udrop_slp", 0.020), + ("std_urej_int", 0.979), + ("std_urej_slp", -0.205), ] ) @@ -76,18 +84,18 @@ [ ("frac_quench_cen_x0_tpeak", 7.000), ("frac_quench_cen_k_tpeak", 2.000), - ("frac_quench_cen_x0_ylotpeak", 11.100), - ("frac_quench_cen_x0_yhitpeak", 13.008), + ("frac_quench_cen_x0_ylotpeak", 11.750), + ("frac_quench_cen_x0_yhitpeak", 12.965), ("frac_quench_cen_ylo_ylotpeak", 0.990), - ("frac_quench_cen_ylo_yhitpeak", 0.446), + ("frac_quench_cen_ylo_yhitpeak", 0.625), ("frac_quench_cen_k", 3.848), ("frac_quench_cen_yhi", 0.971), ("frac_quench_sat_x0_tpeak", 7.000), ("frac_quench_sat_k_tpeak", 2.000), - ("frac_quench_sat_x0_ylotpeak", 11.100), - ("frac_quench_sat_x0_yhitpeak", 13.008), + ("frac_quench_sat_x0_ylotpeak", 11.750), + ("frac_quench_sat_x0_yhitpeak", 12.965), ("frac_quench_sat_ylo_ylotpeak", 0.990), - ("frac_quench_sat_ylo_yhitpeak", 0.446), + ("frac_quench_sat_ylo_yhitpeak", 0.625), ("frac_quench_sat_k", 3.848), ("frac_quench_sat_yhi", 0.971), ] @@ -95,10 +103,10 @@ DELTA_UQT_PDICT = OrderedDict( [ ("delta_uqt_x0", 1.001), - ("delta_uqt_k", 0.836), - ("delta_uqt_ylo", -0.583), - ("delta_uqt_yhi", -0.006), - ("delta_uqt_slope", -0.017), + ("delta_uqt_k", 0.714), + ("delta_uqt_ylo", -0.611), + ("delta_uqt_yhi", 0.002), + ("delta_uqt_slope", -0.024), ] ) SFH_PDF_QUENCH_PDICT = SFH_PDF_FRAC_QUENCH_PDICT.copy() @@ -112,7 +120,6 @@ _UPNAMES = ["u_" + key for key in QseqParams._fields] QseqUParams = namedtuple("QseqUParams", _UPNAMES) - DIFFSTARFITS_GALACTICUS_INPLUSEX_DIFFSTARPOP_PARAMS = DiffstarPopParams( *SFH_PDF_QUENCH_PARAMS, *DEFAULT_SATQUENCHPOP_PARAMS ) diff --git a/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_galacticus_in_situ.py b/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_galacticus_in_situ.py index 5d9348e..708fd1c 100644 --- a/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_galacticus_in_situ.py +++ b/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_galacticus_in_situ.py @@ -1,74 +1,82 @@ -"""""" - from collections import OrderedDict, namedtuple -from ..defaults_mgash import DiffstarPopParams, get_unbounded_diffstarpop_params +from ..defaults_mgash import ( + DiffstarPopParams, + get_unbounded_diffstarpop_params, +) + from ..satquenchpop_model import DEFAULT_SATQUENCHPOP_PARAMS SFH_PDF_QUENCH_MU_PDICT = OrderedDict( [ - ("mean_ulgm_mseq_xtp", 13.561), - ("mean_ulgm_mseq_ytp", 12.181), - ("mean_ulgm_mseq_lo", 0.628), - ("mean_ulgm_mseq_hi", -0.819), - ("mean_ulgy_mseq_int", -9.663), - ("mean_ulgy_mseq_slp", 0.105), + ("mean_ulgm_mseq_xtp", 13.189), + ("mean_ulgm_mseq_ytp", 12.079), + ("mean_ulgm_mseq_lo", 0.783), + ("mean_ulgm_mseq_hi", -0.001), + ("mean_ulgy_mseq_xtp", 11.995), + ("mean_ulgy_mseq_ytp", -9.781), + ("mean_ulgy_mseq_lo", 0.927), + ("mean_ulgy_mseq_hi", 0.091), ("mean_ul_mseq_int", -0.715), ("mean_ul_mseq_slp", 1.717), - ("mean_uh_mseq_int", -0.392), - ("mean_uh_mseq_slp", 0.043), - ("mean_ulgm_qseq_xtp", 12.728), - ("mean_ulgm_qseq_ytp", 11.931), - ("mean_ulgm_qseq_lo", 0.769), - ("mean_ulgm_qseq_hi", 0.123), - ("mean_ulgy_qseq_int", -9.647), - ("mean_ulgy_qseq_slp", 0.382), - ("mean_ul_qseq_int", -0.378), - ("mean_ul_qseq_slp", 1.047), - ("mean_uh_qseq_int", -0.196), - ("mean_uh_qseq_slp", 0.157), - ("mean_uqt_int", 1.034), - ("mean_uqt_slp", -0.096), - ("mean_uqs_int", 0.040), - ("mean_uqs_slp", -0.173), - ("mean_udrop_int", -1.963), - ("mean_udrop_slp", 0.746), - ("mean_urej_int", -0.524), - ("mean_urej_slp", 0.164), + ("mean_uh_mseq_int", -0.253), + ("mean_uh_mseq_slp", 0.364), + ("mean_ulgm_qseq_xtp", 13.473), + ("mean_ulgm_qseq_ytp", 12.089), + ("mean_ulgm_qseq_lo", 0.404), + ("mean_ulgm_qseq_hi", -0.228), + ("mean_ulgy_qseq_xtp", 12.303), + ("mean_ulgy_qseq_ytp", -9.608), + ("mean_ulgy_qseq_lo", 0.446), + ("mean_ulgy_qseq_hi", 0.017), + ("mean_ul_qseq_int", -0.388), + ("mean_ul_qseq_slp", 0.954), + ("mean_uh_qseq_int", -0.153), + ("mean_uh_qseq_slp", 0.284), + ("mean_uqt_xtp", 12.746), + ("mean_uqt_ytp", 1.035), + ("mean_uqt_lo", -0.100), + ("mean_uqt_hi", -0.150), + ("mean_uqs_int", 0.057), + ("mean_uqs_slp", -0.196), + ("mean_udrop_int", -1.917), + ("mean_udrop_slp", 0.684), + ("mean_urej_int", -0.749), + ("mean_urej_slp", 0.472), ] ) SFH_PDF_QUENCH_COV_MS_BLOCK_PDICT = OrderedDict( [ - ("std_ulgm_mseq_int", 0.405), - ("std_ulgm_mseq_slp", -0.094), - ("std_ulgy_mseq_int", 0.229), - ("std_ulgy_mseq_slp", -0.006), - ("std_ul_mseq_int", 2.073), - ("std_ul_mseq_slp", 0.374), - ("std_uh_mseq_int", 1.176), - ("std_uh_mseq_slp", -0.200), - ("std_ulgm_qseq_int", 0.420), - ("std_ulgm_qseq_slp", -0.019), - ("std_ulgy_qseq_int", 0.227), - ("std_ulgy_qseq_slp", -0.014), - ("std_ul_qseq_int", 2.174), - ("std_ul_qseq_slp", -0.137), - ("std_uh_qseq_int", 0.761), - ("std_uh_qseq_slp", -0.141), + ("std_ulgm_mseq_int", 0.527), + ("std_ulgm_mseq_slp", -0.248), + ("std_ulgy_mseq_int", 0.247), + ("std_ulgy_mseq_slp", -0.031), + ("std_ul_mseq_int", 2.156), + ("std_ul_mseq_slp", 0.260), + ("std_uh_mseq_int", 1.112), + ("std_uh_mseq_slp", -0.113), + ("std_ulgm_qseq_int", 0.434), + ("std_ulgm_qseq_slp", -0.097), + ("std_ulgy_qseq_int", 0.229), + ("std_ulgy_qseq_slp", -0.018), + ("std_ul_qseq_int", 2.179), + ("std_ul_qseq_slp", -0.124), + ("std_uh_qseq_int", 0.770), + ("std_uh_qseq_slp", -0.177), ] ) SFH_PDF_QUENCH_COV_Q_BLOCK_PDICT = OrderedDict( [ - ("std_uqt_int", 0.083), - ("std_uqt_slp", 0.070), + ("std_uqt_int", 0.110), + ("std_uqt_slp", 0.032), ("std_uqs_int", 0.900), - ("std_uqs_slp", 0.002), - ("std_udrop_int", 0.607), - ("std_udrop_slp", 0.097), - ("std_urej_int", 0.893), - ("std_urej_slp", 0.260), + ("std_uqs_slp", 0.120), + ("std_udrop_int", 0.639), + ("std_udrop_slp", 0.053), + ("std_urej_int", 1.118), + ("std_urej_slp", -0.044), ] ) @@ -76,29 +84,29 @@ [ ("frac_quench_cen_x0_tpeak", 7.000), ("frac_quench_cen_k_tpeak", 2.000), - ("frac_quench_cen_x0_ylotpeak", 11.708), - ("frac_quench_cen_x0_yhitpeak", 13.900), + ("frac_quench_cen_x0_ylotpeak", 11.692), + ("frac_quench_cen_x0_yhitpeak", 12.963), ("frac_quench_cen_ylo_ylotpeak", 0.990), - ("frac_quench_cen_ylo_yhitpeak", 0.413), + ("frac_quench_cen_ylo_yhitpeak", 0.525), ("frac_quench_cen_k", 3.848), ("frac_quench_cen_yhi", 0.971), ("frac_quench_sat_x0_tpeak", 7.000), ("frac_quench_sat_k_tpeak", 2.000), - ("frac_quench_sat_x0_ylotpeak", 11.708), - ("frac_quench_sat_x0_yhitpeak", 13.900), + ("frac_quench_sat_x0_ylotpeak", 11.692), + ("frac_quench_sat_x0_yhitpeak", 12.963), ("frac_quench_sat_ylo_ylotpeak", 0.990), - ("frac_quench_sat_ylo_yhitpeak", 0.413), + ("frac_quench_sat_ylo_yhitpeak", 0.525), ("frac_quench_sat_k", 3.848), ("frac_quench_sat_yhi", 0.971), ] ) DELTA_UQT_PDICT = OrderedDict( [ - ("delta_uqt_x0", 1.977), - ("delta_uqt_k", 1.334), - ("delta_uqt_ylo", -0.367), - ("delta_uqt_yhi", -0.008), - ("delta_uqt_slope", -0.013), + ("delta_uqt_x0", 1.659), + ("delta_uqt_k", 0.939), + ("delta_uqt_ylo", -0.466), + ("delta_uqt_yhi", 0.000), + ("delta_uqt_slope", -0.022), ] ) SFH_PDF_QUENCH_PDICT = SFH_PDF_FRAC_QUENCH_PDICT.copy() @@ -112,7 +120,6 @@ _UPNAMES = ["u_" + key for key in QseqParams._fields] QseqUParams = namedtuple("QseqUParams", _UPNAMES) - DIFFSTARFITS_GALACTICUS_IN_DIFFSTARPOP_PARAMS = DiffstarPopParams( *SFH_PDF_QUENCH_PARAMS, *DEFAULT_SATQUENCHPOP_PARAMS ) diff --git a/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_smdpl_dr1.py b/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_smdpl_dr1.py index 17debe7..28d4f48 100644 --- a/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_smdpl_dr1.py +++ b/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_smdpl_dr1.py @@ -1,74 +1,82 @@ -"""""" - from collections import OrderedDict, namedtuple -from ..defaults_mgash import DiffstarPopParams, get_unbounded_diffstarpop_params +from ..defaults_mgash import ( + DiffstarPopParams, + get_unbounded_diffstarpop_params, +) + from ..satquenchpop_model import DEFAULT_SATQUENCHPOP_PARAMS SFH_PDF_QUENCH_MU_PDICT = OrderedDict( [ - ("mean_ulgm_mseq_xtp", 12.167), - ("mean_ulgm_mseq_ytp", 12.109), - ("mean_ulgm_mseq_lo", 0.751), - ("mean_ulgm_mseq_hi", 0.173), - ("mean_ulgy_mseq_int", -9.812), - ("mean_ulgy_mseq_slp", 0.625), + ("mean_ulgm_mseq_xtp", 12.097), + ("mean_ulgm_mseq_ytp", 12.082), + ("mean_ulgm_mseq_lo", 0.812), + ("mean_ulgm_mseq_hi", 0.185), + ("mean_ulgy_mseq_xtp", 12.131), + ("mean_ulgy_mseq_ytp", -10.094), + ("mean_ulgy_mseq_lo", 0.929), + ("mean_ulgy_mseq_hi", 0.400), ("mean_ul_mseq_int", -1.815), ("mean_ul_mseq_slp", 1.089), - ("mean_uh_mseq_int", 0.746), - ("mean_uh_mseq_slp", -0.726), - ("mean_ulgm_qseq_xtp", 12.147), - ("mean_ulgm_qseq_ytp", 12.145), - ("mean_ulgm_qseq_lo", 0.940), - ("mean_ulgm_qseq_hi", 0.046), - ("mean_ulgy_qseq_int", -9.829), - ("mean_ulgy_qseq_slp", 0.831), - ("mean_ul_qseq_int", -1.954), - ("mean_ul_qseq_slp", 0.784), - ("mean_uh_qseq_int", 0.986), - ("mean_uh_qseq_slp", -0.186), - ("mean_uqt_int", 1.025), - ("mean_uqt_slp", 0.013), - ("mean_uqs_int", -0.226), - ("mean_uqs_slp", 0.098), - ("mean_udrop_int", -1.823), - ("mean_udrop_slp", 0.071), - ("mean_urej_int", -0.753), - ("mean_urej_slp", -0.236), + ("mean_uh_mseq_int", 0.820), + ("mean_uh_mseq_slp", -0.458), + ("mean_ulgm_qseq_xtp", 12.168), + ("mean_ulgm_qseq_ytp", 12.153), + ("mean_ulgm_qseq_lo", 0.882), + ("mean_ulgm_qseq_hi", 0.040), + ("mean_ulgy_qseq_xtp", 12.803), + ("mean_ulgy_qseq_ytp", -9.594), + ("mean_ulgy_qseq_lo", 0.714), + ("mean_ulgy_qseq_hi", 0.400), + ("mean_ul_qseq_int", -2.700), + ("mean_ul_qseq_slp", -0.220), + ("mean_uh_qseq_int", 1.910), + ("mean_uh_qseq_slp", 1.035), + ("mean_uqt_xtp", 12.005), + ("mean_uqt_ytp", 1.036), + ("mean_uqt_lo", -0.190), + ("mean_uqt_hi", -0.100), + ("mean_uqs_int", -0.230), + ("mean_uqs_slp", 0.104), + ("mean_udrop_int", -1.814), + ("mean_udrop_slp", 0.054), + ("mean_urej_int", -0.739), + ("mean_urej_slp", -0.261), ] ) SFH_PDF_QUENCH_COV_MS_BLOCK_PDICT = OrderedDict( [ ("std_ulgm_mseq_int", 0.254), - ("std_ulgm_mseq_slp", 0.154), - ("std_ulgy_mseq_int", 0.269), - ("std_ulgy_mseq_slp", 0.033), - ("std_ul_mseq_int", 1.925), - ("std_ul_mseq_slp", 0.544), - ("std_uh_mseq_int", 1.463), - ("std_uh_mseq_slp", 0.010), - ("std_ulgm_qseq_int", 0.192), - ("std_ulgm_qseq_slp", 0.127), - ("std_ulgy_qseq_int", 0.287), - ("std_ulgy_qseq_slp", -0.034), - ("std_ul_qseq_int", 2.015), - ("std_ul_qseq_slp", 0.181), - ("std_uh_qseq_int", 1.041), - ("std_uh_qseq_slp", 0.063), + ("std_ulgm_mseq_slp", 0.153), + ("std_ulgy_mseq_int", 0.273), + ("std_ulgy_mseq_slp", 0.027), + ("std_ul_mseq_int", 1.941), + ("std_ul_mseq_slp", 0.516), + ("std_uh_mseq_int", 1.458), + ("std_uh_mseq_slp", 0.020), + ("std_ulgm_qseq_int", 0.194), + ("std_ulgm_qseq_slp", 0.096), + ("std_ulgy_qseq_int", 0.291), + ("std_ulgy_qseq_slp", -0.223), + ("std_ul_qseq_int", 2.091), + ("std_ul_qseq_slp", -0.183), + ("std_uh_qseq_int", 1.017), + ("std_uh_qseq_slp", 0.066), ] ) SFH_PDF_QUENCH_COV_Q_BLOCK_PDICT = OrderedDict( [ - ("std_uqt_int", 0.080), - ("std_uqt_slp", 0.010), - ("std_uqs_int", 0.613), - ("std_uqs_slp", 0.221), - ("std_udrop_int", 0.796), - ("std_udrop_slp", -0.120), - ("std_urej_int", 1.166), - ("std_urej_slp", -0.021), + ("std_uqt_int", 0.079), + ("std_uqt_slp", 0.001), + ("std_uqs_int", 0.615), + ("std_uqs_slp", 0.217), + ("std_udrop_int", 0.798), + ("std_udrop_slp", -0.124), + ("std_urej_int", 1.168), + ("std_urej_slp", -0.025), ] ) @@ -76,29 +84,29 @@ [ ("frac_quench_cen_x0_tpeak", 7.000), ("frac_quench_cen_k_tpeak", 2.000), - ("frac_quench_cen_x0_ylotpeak", 11.238), - ("frac_quench_cen_x0_yhitpeak", 11.780), - ("frac_quench_cen_ylo_ylotpeak", 0.246), - ("frac_quench_cen_ylo_yhitpeak", 0.027), + ("frac_quench_cen_x0_ylotpeak", 11.280), + ("frac_quench_cen_x0_yhitpeak", 11.658), + ("frac_quench_cen_ylo_ylotpeak", 0.010), + ("frac_quench_cen_ylo_yhitpeak", 0.040), ("frac_quench_cen_k", 3.848), ("frac_quench_cen_yhi", 0.971), ("frac_quench_sat_x0_tpeak", 7.000), ("frac_quench_sat_k_tpeak", 2.000), - ("frac_quench_sat_x0_ylotpeak", 11.238), - ("frac_quench_sat_x0_yhitpeak", 11.780), - ("frac_quench_sat_ylo_ylotpeak", 0.246), - ("frac_quench_sat_ylo_yhitpeak", 0.027), + ("frac_quench_sat_x0_ylotpeak", 11.280), + ("frac_quench_sat_x0_yhitpeak", 11.658), + ("frac_quench_sat_ylo_ylotpeak", 0.010), + ("frac_quench_sat_ylo_yhitpeak", 0.040), ("frac_quench_sat_k", 3.848), ("frac_quench_sat_yhi", 0.971), ] ) DELTA_UQT_PDICT = OrderedDict( [ - ("delta_uqt_x0", 2.846), - ("delta_uqt_k", 0.515), - ("delta_uqt_ylo", -0.484), - ("delta_uqt_yhi", 0.036), - ("delta_uqt_slope", -0.072), + ("delta_uqt_x0", 3.260), + ("delta_uqt_k", 0.536), + ("delta_uqt_ylo", -0.372), + ("delta_uqt_yhi", 0.015), + ("delta_uqt_slope", -0.065), ] ) SFH_PDF_QUENCH_PDICT = SFH_PDF_FRAC_QUENCH_PDICT.copy() @@ -112,7 +120,6 @@ _UPNAMES = ["u_" + key for key in QseqParams._fields] QseqUParams = namedtuple("QseqUParams", _UPNAMES) - DIFFSTARFITS_SMDPL_DR1_DIFFSTARPOP_PARAMS = DiffstarPopParams( *SFH_PDF_QUENCH_PARAMS, *DEFAULT_SATQUENCHPOP_PARAMS ) diff --git a/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_smdpl_dr1_nomerging.py b/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_smdpl_dr1_nomerging.py index 5afa440..c11ab6f 100644 --- a/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_smdpl_dr1_nomerging.py +++ b/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_smdpl_dr1_nomerging.py @@ -1,74 +1,82 @@ -"""""" - from collections import OrderedDict, namedtuple -from ..defaults_mgash import DiffstarPopParams, get_unbounded_diffstarpop_params +from ..defaults_mgash import ( + DiffstarPopParams, + get_unbounded_diffstarpop_params, +) + from ..satquenchpop_model import DEFAULT_SATQUENCHPOP_PARAMS SFH_PDF_QUENCH_MU_PDICT = OrderedDict( [ - ("mean_ulgm_mseq_xtp", 12.113), - ("mean_ulgm_mseq_ytp", 12.089), - ("mean_ulgm_mseq_lo", 0.821), - ("mean_ulgm_mseq_hi", 0.108), - ("mean_ulgy_mseq_int", -9.899), - ("mean_ulgy_mseq_slp", 0.389), + ("mean_ulgm_mseq_xtp", 12.133), + ("mean_ulgm_mseq_ytp", 12.097), + ("mean_ulgm_mseq_lo", 0.793), + ("mean_ulgm_mseq_hi", 0.104), + ("mean_ulgy_mseq_xtp", 13.085), + ("mean_ulgy_mseq_ytp", -9.718), + ("mean_ulgy_mseq_lo", 0.638), + ("mean_ulgy_mseq_hi", 0.144), ("mean_ul_mseq_int", -0.804), ("mean_ul_mseq_slp", 0.743), - ("mean_uh_mseq_int", -1.615), - ("mean_uh_mseq_slp", -2.708), - ("mean_ulgm_qseq_xtp", 12.282), - ("mean_ulgm_qseq_ytp", 12.238), - ("mean_ulgm_qseq_lo", 0.781), - ("mean_ulgm_qseq_hi", 0.072), - ("mean_ulgy_qseq_int", -9.977), - ("mean_ulgy_qseq_slp", 0.734), - ("mean_ul_qseq_int", -0.705), - ("mean_ul_qseq_slp", 0.850), - ("mean_uh_qseq_int", -0.350), - ("mean_uh_qseq_slp", -1.369), - ("mean_uqt_int", 0.950), - ("mean_uqt_slp", -0.203), - ("mean_uqs_int", -0.112), - ("mean_uqs_slp", 0.395), - ("mean_udrop_int", -2.143), - ("mean_udrop_slp", 0.232), - ("mean_urej_int", -1.227), - ("mean_urej_slp", -0.010), + ("mean_uh_mseq_int", -1.504), + ("mean_uh_mseq_slp", -2.294), + ("mean_ulgm_qseq_xtp", 12.264), + ("mean_ulgm_qseq_ytp", 12.232), + ("mean_ulgm_qseq_lo", 0.816), + ("mean_ulgm_qseq_hi", 0.075), + ("mean_ulgy_qseq_xtp", 12.567), + ("mean_ulgy_qseq_ytp", -9.840), + ("mean_ulgy_qseq_lo", 0.579), + ("mean_ulgy_qseq_hi", 0.329), + ("mean_ul_qseq_int", -0.847), + ("mean_ul_qseq_slp", 0.753), + ("mean_uh_qseq_int", 1.142), + ("mean_uh_qseq_slp", 0.581), + ("mean_uqt_xtp", 12.862), + ("mean_uqt_ytp", 0.902), + ("mean_uqt_lo", -0.124), + ("mean_uqt_hi", -0.267), + ("mean_uqs_int", -0.109), + ("mean_uqs_slp", 0.391), + ("mean_udrop_int", -2.140), + ("mean_udrop_slp", 0.229), + ("mean_urej_int", -1.217), + ("mean_urej_slp", -0.024), ] ) SFH_PDF_QUENCH_COV_MS_BLOCK_PDICT = OrderedDict( [ - ("std_ulgm_mseq_int", 0.287), - ("std_ulgm_mseq_slp", 0.146), - ("std_ulgy_mseq_int", 0.285), - ("std_ulgy_mseq_slp", 0.039), - ("std_ul_mseq_int", 1.263), - ("std_ul_mseq_slp", 0.445), - ("std_uh_mseq_int", 2.516), + ("std_ulgm_mseq_int", 0.288), + ("std_ulgm_mseq_slp", 0.145), + ("std_ulgy_mseq_int", 0.289), + ("std_ulgy_mseq_slp", 0.034), + ("std_ul_mseq_int", 1.284), + ("std_ul_mseq_slp", 0.415), + ("std_uh_mseq_int", 2.477), ("std_uh_mseq_slp", -0.900), - ("std_ulgm_qseq_int", 0.250), - ("std_ulgm_qseq_slp", 0.190), - ("std_ulgy_qseq_int", 0.307), - ("std_ulgy_qseq_slp", 0.022), - ("std_ul_qseq_int", 1.687), - ("std_ul_qseq_slp", 0.269), - ("std_uh_qseq_int", 1.499), - ("std_uh_qseq_slp", -0.246), + ("std_ulgm_qseq_int", 0.238), + ("std_ulgm_qseq_slp", 0.125), + ("std_ulgy_qseq_int", 0.257), + ("std_ulgy_qseq_slp", -0.220), + ("std_ul_qseq_int", 1.661), + ("std_ul_qseq_slp", -0.455), + ("std_uh_qseq_int", 1.484), + ("std_uh_qseq_slp", 0.218), ] ) SFH_PDF_QUENCH_COV_Q_BLOCK_PDICT = OrderedDict( [ ("std_uqt_int", 0.121), - ("std_uqt_slp", 0.052), - ("std_uqs_int", 0.567), - ("std_uqs_slp", -0.029), - ("std_udrop_int", 0.749), - ("std_udrop_slp", 0.162), - ("std_urej_int", 1.357), - ("std_urej_slp", 0.069), + ("std_uqt_slp", 0.053), + ("std_uqs_int", 0.565), + ("std_uqs_slp", -0.027), + ("std_udrop_int", 0.745), + ("std_udrop_slp", 0.168), + ("std_urej_int", 1.359), + ("std_urej_slp", 0.066), ] ) @@ -76,17 +84,17 @@ [ ("frac_quench_cen_x0_tpeak", 7.000), ("frac_quench_cen_k_tpeak", 2.000), - ("frac_quench_cen_x0_ylotpeak", 11.652), - ("frac_quench_cen_x0_yhitpeak", 11.882), - ("frac_quench_cen_ylo_ylotpeak", 0.555), + ("frac_quench_cen_x0_ylotpeak", 11.349), + ("frac_quench_cen_x0_yhitpeak", 11.862), + ("frac_quench_cen_ylo_ylotpeak", 0.038), ("frac_quench_cen_ylo_yhitpeak", 0.010), ("frac_quench_cen_k", 3.848), ("frac_quench_cen_yhi", 0.971), ("frac_quench_sat_x0_tpeak", 7.000), ("frac_quench_sat_k_tpeak", 2.000), - ("frac_quench_sat_x0_ylotpeak", 11.652), - ("frac_quench_sat_x0_yhitpeak", 11.882), - ("frac_quench_sat_ylo_ylotpeak", 0.555), + ("frac_quench_sat_x0_ylotpeak", 11.349), + ("frac_quench_sat_x0_yhitpeak", 11.862), + ("frac_quench_sat_ylo_ylotpeak", 0.038), ("frac_quench_sat_ylo_yhitpeak", 0.010), ("frac_quench_sat_k", 3.848), ("frac_quench_sat_yhi", 0.971), @@ -94,11 +102,11 @@ ) DELTA_UQT_PDICT = OrderedDict( [ - ("delta_uqt_x0", 1.554), - ("delta_uqt_k", 0.501), - ("delta_uqt_ylo", -0.578), - ("delta_uqt_yhi", 0.037), - ("delta_uqt_slope", -0.050), + ("delta_uqt_x0", 1.548), + ("delta_uqt_k", 0.493), + ("delta_uqt_ylo", -0.494), + ("delta_uqt_yhi", 0.023), + ("delta_uqt_slope", -0.049), ] ) SFH_PDF_QUENCH_PDICT = SFH_PDF_FRAC_QUENCH_PDICT.copy() @@ -112,7 +120,6 @@ _UPNAMES = ["u_" + key for key in QseqParams._fields] QseqUParams = namedtuple("QseqUParams", _UPNAMES) - DIFFSTARFITS_SMDPL_DIFFSTARPOP_PARAMS = DiffstarPopParams( *SFH_PDF_QUENCH_PARAMS, *DEFAULT_SATQUENCHPOP_PARAMS ) diff --git a/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_tng.py b/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_tng.py index 8e158b3..21553f7 100644 --- a/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_tng.py +++ b/diffstar/diffstarpop/kernels/params/params_diffstarfits_mgash_tng.py @@ -1,74 +1,82 @@ -"""""" - from collections import OrderedDict, namedtuple -from ..defaults_mgash import DiffstarPopParams, get_unbounded_diffstarpop_params +from ..defaults_mgash import ( + DiffstarPopParams, + get_unbounded_diffstarpop_params, +) + from ..satquenchpop_model import DEFAULT_SATQUENCHPOP_PARAMS SFH_PDF_QUENCH_MU_PDICT = OrderedDict( [ - ("mean_ulgm_mseq_xtp", 12.315), - ("mean_ulgm_mseq_ytp", 11.579), - ("mean_ulgm_mseq_lo", 0.122), - ("mean_ulgm_mseq_hi", 0.372), - ("mean_ulgy_mseq_int", -9.593), - ("mean_ulgy_mseq_slp", 0.403), + ("mean_ulgm_mseq_xtp", 11.300), + ("mean_ulgm_mseq_ytp", 11.300), + ("mean_ulgm_mseq_lo", 1.296), + ("mean_ulgm_mseq_hi", 0.336), + ("mean_ulgy_mseq_xtp", 13.509), + ("mean_ulgy_mseq_ytp", -9.297), + ("mean_ulgy_mseq_lo", 0.543), + ("mean_ulgy_mseq_hi", 0.247), ("mean_ul_mseq_int", -0.358), ("mean_ul_mseq_slp", 1.187), - ("mean_uh_mseq_int", -0.473), - ("mean_uh_mseq_slp", 1.422), - ("mean_ulgm_qseq_xtp", 11.803), - ("mean_ulgm_qseq_ytp", 11.551), - ("mean_ulgm_qseq_lo", -0.008), - ("mean_ulgm_qseq_hi", 0.282), - ("mean_ulgy_qseq_int", -9.828), - ("mean_ulgy_qseq_slp", 0.562), - ("mean_ul_qseq_int", -0.328), - ("mean_ul_qseq_slp", 1.712), - ("mean_uh_qseq_int", -0.755), - ("mean_uh_qseq_slp", -0.501), - ("mean_uqt_int", 0.990), - ("mean_uqt_slp", -0.071), - ("mean_uqs_int", 0.172), - ("mean_uqs_slp", -0.124), - ("mean_udrop_int", -2.095), - ("mean_udrop_slp", 0.646), - ("mean_urej_int", -0.979), - ("mean_urej_slp", 0.478), + ("mean_uh_mseq_int", -0.807), + ("mean_uh_mseq_slp", 0.066), + ("mean_ulgm_qseq_xtp", 12.622), + ("mean_ulgm_qseq_ytp", 11.767), + ("mean_ulgm_qseq_lo", 0.289), + ("mean_ulgm_qseq_hi", 0.297), + ("mean_ulgy_qseq_xtp", 12.334), + ("mean_ulgy_qseq_ytp", -9.655), + ("mean_ulgy_qseq_lo", 0.572), + ("mean_ulgy_qseq_hi", 0.388), + ("mean_ul_qseq_int", -0.354), + ("mean_ul_qseq_slp", 1.694), + ("mean_uh_qseq_int", -0.924), + ("mean_uh_qseq_slp", -0.869), + ("mean_uqt_xtp", 13.226), + ("mean_uqt_ytp", 0.929), + ("mean_uqt_lo", -0.100), + ("mean_uqt_hi", -0.100), + ("mean_uqs_int", 0.144), + ("mean_uqs_slp", -0.081), + ("mean_udrop_int", -2.026), + ("mean_udrop_slp", 0.543), + ("mean_urej_int", -0.932), + ("mean_urej_slp", 0.409), ] ) SFH_PDF_QUENCH_COV_MS_BLOCK_PDICT = OrderedDict( [ - ("std_ulgm_mseq_int", 0.404), - ("std_ulgm_mseq_slp", -0.013), - ("std_ulgy_mseq_int", 0.281), - ("std_ulgy_mseq_slp", 0.009), - ("std_ul_mseq_int", 1.583), - ("std_ul_mseq_slp", 0.462), - ("std_uh_mseq_int", 1.616), - ("std_uh_mseq_slp", -0.297), - ("std_ulgm_qseq_int", 0.361), - ("std_ulgm_qseq_slp", 0.056), - ("std_ulgy_qseq_int", 0.275), - ("std_ulgy_qseq_slp", -0.011), - ("std_ul_qseq_int", 1.740), - ("std_ul_qseq_slp", -0.135), - ("std_uh_qseq_int", 1.327), - ("std_uh_qseq_slp", -0.477), + ("std_ulgm_mseq_int", 0.379), + ("std_ulgm_mseq_slp", 0.025), + ("std_ulgy_mseq_int", 0.283), + ("std_ulgy_mseq_slp", 0.007), + ("std_ul_mseq_int", 1.665), + ("std_ul_mseq_slp", 0.337), + ("std_uh_mseq_int", 1.588), + ("std_uh_mseq_slp", -0.254), + ("std_ulgm_qseq_int", 0.362), + ("std_ulgm_qseq_slp", 0.001), + ("std_ulgy_qseq_int", 0.274), + ("std_ulgy_qseq_slp", -0.008), + ("std_ul_qseq_int", 1.766), + ("std_ul_qseq_slp", -0.476), + ("std_uh_qseq_int", 1.306), + ("std_uh_qseq_slp", -0.208), ] ) SFH_PDF_QUENCH_COV_Q_BLOCK_PDICT = OrderedDict( [ - ("std_uqt_int", 0.124), - ("std_uqt_slp", 0.082), - ("std_uqs_int", 0.577), - ("std_uqs_slp", 0.204), - ("std_udrop_int", 0.806), - ("std_udrop_slp", -0.100), - ("std_urej_int", 1.197), - ("std_urej_slp", -0.076), + ("std_uqt_int", 0.131), + ("std_uqt_slp", 0.071), + ("std_uqs_int", 0.590), + ("std_uqs_slp", 0.184), + ("std_udrop_int", 0.784), + ("std_udrop_slp", -0.065), + ("std_urej_int", 1.221), + ("std_urej_slp", -0.111), ] ) @@ -76,29 +84,29 @@ [ ("frac_quench_cen_x0_tpeak", 7.000), ("frac_quench_cen_k_tpeak", 2.000), - ("frac_quench_cen_x0_ylotpeak", 11.610), - ("frac_quench_cen_x0_yhitpeak", 12.066), + ("frac_quench_cen_x0_ylotpeak", 11.100), + ("frac_quench_cen_x0_yhitpeak", 12.813), ("frac_quench_cen_ylo_ylotpeak", 0.990), - ("frac_quench_cen_ylo_yhitpeak", 0.099), + ("frac_quench_cen_ylo_yhitpeak", 0.196), ("frac_quench_cen_k", 3.848), ("frac_quench_cen_yhi", 0.971), ("frac_quench_sat_x0_tpeak", 7.000), ("frac_quench_sat_k_tpeak", 2.000), - ("frac_quench_sat_x0_ylotpeak", 11.610), - ("frac_quench_sat_x0_yhitpeak", 12.066), + ("frac_quench_sat_x0_ylotpeak", 11.100), + ("frac_quench_sat_x0_yhitpeak", 12.813), ("frac_quench_sat_ylo_ylotpeak", 0.990), - ("frac_quench_sat_ylo_yhitpeak", 0.099), + ("frac_quench_sat_ylo_yhitpeak", 0.196), ("frac_quench_sat_k", 3.848), ("frac_quench_sat_yhi", 0.971), ] ) DELTA_UQT_PDICT = OrderedDict( [ - ("delta_uqt_x0", 3.294), - ("delta_uqt_k", 4.727), - ("delta_uqt_ylo", -0.326), - ("delta_uqt_yhi", 0.025), - ("delta_uqt_slope", 0.021), + ("delta_uqt_x0", 3.309), + ("delta_uqt_k", 4.719), + ("delta_uqt_ylo", -0.308), + ("delta_uqt_yhi", 0.029), + ("delta_uqt_slope", 0.009), ] ) SFH_PDF_QUENCH_PDICT = SFH_PDF_FRAC_QUENCH_PDICT.copy() @@ -112,7 +120,6 @@ _UPNAMES = ["u_" + key for key in QseqParams._fields] QseqUParams = namedtuple("QseqUParams", _UPNAMES) - DIFFSTARFITS_TNG_DIFFSTARPOP_PARAMS = DiffstarPopParams( *SFH_PDF_QUENCH_PARAMS, *DEFAULT_SATQUENCHPOP_PARAMS ) diff --git a/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_galacticus_in_plus_ex_situ.py b/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_galacticus_in_plus_ex_situ.py index 5efd474..f6eda0d 100644 --- a/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_galacticus_in_plus_ex_situ.py +++ b/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_galacticus_in_plus_ex_situ.py @@ -1,105 +1,113 @@ -"""""" - from collections import OrderedDict, namedtuple -from ..defaults_mgash import DiffstarPopParams, get_unbounded_diffstarpop_params +from ..defaults_mgash import ( + DiffstarPopParams, + get_unbounded_diffstarpop_params, +) + from ..satquenchpop_model import DEFAULT_SATQUENCHPOP_PARAMS SFH_PDF_QUENCH_MU_PDICT = OrderedDict( [ - ("mean_ulgm_mseq_xtp", 11.990), - ("mean_ulgm_mseq_ytp", 11.146), - ("mean_ulgm_mseq_lo", 0.415), - ("mean_ulgm_mseq_hi", -0.261), - ("mean_ulgy_mseq_int", -8.857), - ("mean_ulgy_mseq_slp", 1.175), - ("mean_ul_mseq_int", -2.622), - ("mean_ul_mseq_slp", 0.819), - ("mean_uh_mseq_int", -0.577), - ("mean_uh_mseq_slp", 0.606), - ("mean_ulgm_qseq_xtp", 13.247), - ("mean_ulgm_qseq_ytp", 11.923), - ("mean_ulgm_qseq_lo", 0.290), - ("mean_ulgm_qseq_hi", 0.500), - ("mean_ulgy_qseq_int", -9.306), - ("mean_ulgy_qseq_slp", 0.592), + ("mean_ulgm_mseq_xtp", 11.986), + ("mean_ulgm_mseq_ytp", 11.189), + ("mean_ulgm_mseq_lo", 0.370), + ("mean_ulgm_mseq_hi", 0.147), + ("mean_ulgy_mseq_xtp", 12.677), + ("mean_ulgy_mseq_ytp", -9.236), + ("mean_ulgy_mseq_lo", 0.934), + ("mean_ulgy_mseq_hi", -2.625), + ("mean_ul_mseq_int", -2.997), + ("mean_ul_mseq_slp", 0.001), + ("mean_uh_mseq_int", -0.596), + ("mean_uh_mseq_slp", 0.708), + ("mean_ulgm_qseq_xtp", 13.105), + ("mean_ulgm_qseq_ytp", 11.937), + ("mean_ulgm_qseq_lo", 0.504), + ("mean_ulgm_qseq_hi", 0.616), + ("mean_ulgy_qseq_xtp", 11.979), + ("mean_ulgy_qseq_ytp", -9.432), + ("mean_ulgy_qseq_lo", 0.995), + ("mean_ulgy_qseq_hi", 0.310), ("mean_ul_qseq_int", -2.997), - ("mean_ul_qseq_slp", 3.163), - ("mean_uh_qseq_int", -1.156), - ("mean_uh_qseq_slp", -0.498), - ("mean_uqt_int", 1.109), - ("mean_uqt_slp", 0.024), - ("mean_uqs_int", 1.581), - ("mean_uqs_slp", 0.434), - ("mean_udrop_int", -2.222), - ("mean_udrop_slp", 0.983), - ("mean_urej_int", -1.646), - ("mean_urej_slp", 2.075), + ("mean_ul_qseq_slp", -1.731), + ("mean_uh_qseq_int", -1.134), + ("mean_uh_qseq_slp", -0.500), + ("mean_uqt_xtp", 12.400), + ("mean_uqt_ytp", 1.104), + ("mean_uqt_lo", -0.053), + ("mean_uqt_hi", -0.209), + ("mean_uqs_int", 1.676), + ("mean_uqs_slp", 0.509), + ("mean_udrop_int", -2.189), + ("mean_udrop_slp", 1.064), + ("mean_urej_int", -3.147), + ("mean_urej_slp", 0.202), ] ) SFH_PDF_QUENCH_COV_MS_BLOCK_PDICT = OrderedDict( [ ("std_ulgm_mseq_int", 0.011), - ("std_ulgm_mseq_slp", 0.999), - ("std_ulgy_mseq_int", 0.053), - ("std_ulgy_mseq_slp", 0.999), - ("std_ul_mseq_int", 0.054), - ("std_ul_mseq_slp", 0.058), - ("std_uh_mseq_int", 0.119), - ("std_uh_mseq_slp", 0.999), - ("std_ulgm_qseq_int", 0.011), - ("std_ulgm_qseq_slp", -0.341), - ("std_ulgy_qseq_int", 0.011), - ("std_ulgy_qseq_slp", 0.198), - ("std_ul_qseq_int", 0.825), + ("std_ulgm_mseq_slp", -0.081), + ("std_ulgy_mseq_int", 0.011), + ("std_ulgy_mseq_slp", 0.234), + ("std_ul_mseq_int", 0.013), + ("std_ul_mseq_slp", 0.006), + ("std_uh_mseq_int", 0.012), + ("std_uh_mseq_slp", 0.005), + ("std_ulgm_qseq_int", 0.037), + ("std_ulgm_qseq_slp", 0.030), + ("std_ulgy_qseq_int", 0.025), + ("std_ulgy_qseq_slp", 0.146), + ("std_ul_qseq_int", 0.011), ("std_ul_qseq_slp", -0.999), - ("std_uh_qseq_int", 0.133), - ("std_uh_qseq_slp", 0.082), + ("std_uh_qseq_int", 0.301), + ("std_uh_qseq_slp", -0.083), ] ) SFH_PDF_QUENCH_COV_Q_BLOCK_PDICT = OrderedDict( [ - ("std_uqt_int", 0.122), - ("std_uqt_slp", 0.048), + ("std_uqt_int", 0.114), + ("std_uqt_slp", 0.057), ("std_uqs_int", 0.999), - ("std_uqs_slp", 0.636), - ("std_udrop_int", 0.435), - ("std_udrop_slp", 0.285), - ("std_urej_int", 0.011), - ("std_urej_slp", 0.739), + ("std_uqs_slp", 0.930), + ("std_udrop_int", 0.512), + ("std_udrop_slp", 0.280), + ("std_urej_int", 0.025), + ("std_urej_slp", 0.008), ] ) SFH_PDF_FRAC_QUENCH_PDICT = OrderedDict( [ - ("frac_quench_cen_x0_tpeak", 6.476), - ("frac_quench_cen_k_tpeak", 2.282), - ("frac_quench_cen_x0_ylotpeak", 11.100), - ("frac_quench_cen_x0_yhitpeak", 12.619), + ("frac_quench_cen_x0_tpeak", 6.475), + ("frac_quench_cen_k_tpeak", 2.287), + ("frac_quench_cen_x0_ylotpeak", 11.742), + ("frac_quench_cen_x0_yhitpeak", 12.605), ("frac_quench_cen_ylo_ylotpeak", 0.990), - ("frac_quench_cen_ylo_yhitpeak", 0.028), + ("frac_quench_cen_ylo_yhitpeak", 0.007), ("frac_quench_cen_k", 4.995), - ("frac_quench_cen_yhi", 0.999), - ("frac_quench_sat_x0_tpeak", 5.944), - ("frac_quench_sat_k_tpeak", 9.994), - ("frac_quench_sat_x0_ylotpeak", 11.906), - ("frac_quench_sat_x0_yhitpeak", 12.378), + ("frac_quench_cen_yhi", 0.997), + ("frac_quench_sat_x0_tpeak", 7.617), + ("frac_quench_sat_k_tpeak", 9.998), + ("frac_quench_sat_x0_ylotpeak", 13.994), + ("frac_quench_sat_x0_yhitpeak", 12.583), ("frac_quench_sat_ylo_ylotpeak", 0.999), - ("frac_quench_sat_ylo_yhitpeak", 0.677), + ("frac_quench_sat_ylo_yhitpeak", 0.665), ("frac_quench_sat_k", 4.995), - ("frac_quench_sat_yhi", 0.924), + ("frac_quench_sat_yhi", 0.999), ] ) DELTA_UQT_PDICT = OrderedDict( [ - ("delta_uqt_x0", 2.373), - ("delta_uqt_k", 0.450), - ("delta_uqt_ylo", -0.622), - ("delta_uqt_yhi", 0.075), - ("delta_uqt_slope", -0.017), + ("delta_uqt_x0", 2.423), + ("delta_uqt_k", 0.285), + ("delta_uqt_ylo", -0.758), + ("delta_uqt_yhi", 0.227), + ("delta_uqt_slope", 0.049), ] ) SFH_PDF_QUENCH_PDICT = SFH_PDF_FRAC_QUENCH_PDICT.copy() @@ -114,7 +122,6 @@ _UPNAMES = ["u_" + key for key in QseqParams._fields] QseqUParams = namedtuple("QseqUParams", _UPNAMES) - DIFFSTARPOP_FITS_GALACTICUS_INPLUSEX_DIFFSTARPOP_PARAMS = DiffstarPopParams( *SFH_PDF_QUENCH_PARAMS, *DEFAULT_SATQUENCHPOP_PARAMS ) diff --git a/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_galacticus_in_situ.py b/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_galacticus_in_situ.py index 1f4abff..0c9d504 100644 --- a/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_galacticus_in_situ.py +++ b/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_galacticus_in_situ.py @@ -1,93 +1,101 @@ -"""""" - from collections import OrderedDict, namedtuple -from ..defaults_mgash import DiffstarPopParams, get_unbounded_diffstarpop_params +from ..defaults_mgash import ( + DiffstarPopParams, + get_unbounded_diffstarpop_params, +) + from ..satquenchpop_model import DEFAULT_SATQUENCHPOP_PARAMS SFH_PDF_QUENCH_MU_PDICT = OrderedDict( [ - ("mean_ulgm_mseq_xtp", 12.576), - ("mean_ulgm_mseq_ytp", 11.226), - ("mean_ulgm_mseq_lo", 0.252), - ("mean_ulgm_mseq_hi", -2.900), - ("mean_ulgy_mseq_int", -9.254), - ("mean_ulgy_mseq_slp", 1.003), - ("mean_ul_mseq_int", -2.997), - ("mean_ul_mseq_slp", 12.403), - ("mean_uh_mseq_int", -0.598), - ("mean_uh_mseq_slp", 0.336), - ("mean_ulgm_qseq_xtp", 12.167), - ("mean_ulgm_qseq_ytp", 11.852), - ("mean_ulgm_qseq_lo", 0.840), - ("mean_ulgm_qseq_hi", 0.091), - ("mean_ulgy_qseq_int", -9.671), - ("mean_ulgy_qseq_slp", 0.312), - ("mean_ul_qseq_int", -2.997), - ("mean_ul_qseq_slp", 4.249), - ("mean_uh_qseq_int", -1.484), - ("mean_uh_qseq_slp", -0.251), - ("mean_uqt_int", 0.927), - ("mean_uqt_slp", -0.085), - ("mean_uqs_int", 0.161), - ("mean_uqs_slp", -0.798), - ("mean_udrop_int", -1.809), - ("mean_udrop_slp", 1.253), - ("mean_urej_int", -4.271), - ("mean_urej_slp", -0.664), + ("mean_ulgm_mseq_xtp", 12.077), + ("mean_ulgm_mseq_ytp", 11.198), + ("mean_ulgm_mseq_lo", 0.324), + ("mean_ulgm_mseq_hi", -0.612), + ("mean_ulgy_mseq_xtp", 12.708), + ("mean_ulgy_mseq_ytp", -9.579), + ("mean_ulgy_mseq_lo", 0.753), + ("mean_ulgy_mseq_hi", -3.358), + ("mean_ul_mseq_int", -2.739), + ("mean_ul_mseq_slp", 8.577), + ("mean_uh_mseq_int", 0.164), + ("mean_uh_mseq_slp", 1.255), + ("mean_ulgm_qseq_xtp", 13.525), + ("mean_ulgm_qseq_ytp", 12.100), + ("mean_ulgm_qseq_lo", 0.420), + ("mean_ulgm_qseq_hi", 0.423), + ("mean_ulgy_qseq_xtp", 12.166), + ("mean_ulgy_qseq_ytp", -9.571), + ("mean_ulgy_qseq_lo", 0.702), + ("mean_ulgy_qseq_hi", 0.005), + ("mean_ul_qseq_int", -2.400), + ("mean_ul_qseq_slp", 0.968), + ("mean_uh_qseq_int", -0.915), + ("mean_uh_qseq_slp", -0.459), + ("mean_uqt_xtp", 12.555), + ("mean_uqt_ytp", 1.035), + ("mean_uqt_lo", -0.031), + ("mean_uqt_hi", -0.589), + ("mean_uqs_int", 1.451), + ("mean_uqs_slp", 0.307), + ("mean_udrop_int", -2.247), + ("mean_udrop_slp", 1.486), + ("mean_urej_int", -1.420), + ("mean_urej_slp", 1.622), ] ) SFH_PDF_QUENCH_COV_MS_BLOCK_PDICT = OrderedDict( [ ("std_ulgm_mseq_int", 0.011), - ("std_ulgm_mseq_slp", 0.001), - ("std_ulgy_mseq_int", 0.046), - ("std_ulgy_mseq_slp", 0.132), - ("std_ul_mseq_int", 0.011), - ("std_ul_mseq_slp", 0.005), + ("std_ulgm_mseq_slp", 0.011), + ("std_ulgy_mseq_int", 0.011), + ("std_ulgy_mseq_slp", 0.009), + ("std_ul_mseq_int", 1.725), + ("std_ul_mseq_slp", -0.999), ("std_uh_mseq_int", 0.011), - ("std_uh_mseq_slp", -0.999), - ("std_ulgm_qseq_int", 0.142), - ("std_ulgm_qseq_slp", -0.021), + ("std_uh_mseq_slp", -0.981), + ("std_ulgm_qseq_int", 0.095), + ("std_ulgm_qseq_slp", 0.004), ("std_ulgy_qseq_int", 0.011), - ("std_ulgy_qseq_slp", -0.189), - ("std_ul_qseq_int", 1.080), - ("std_ul_qseq_slp", -0.999), - ("std_uh_qseq_int", 0.011), - ("std_uh_qseq_slp", 0.086), + ("std_ulgy_qseq_slp", -0.122), + ("std_ul_qseq_int", 2.224), + ("std_ul_qseq_slp", -0.412), + ("std_uh_qseq_int", 0.149), + ("std_uh_qseq_slp", -0.996), ] ) SFH_PDF_QUENCH_COV_Q_BLOCK_PDICT = OrderedDict( [ - ("std_uqt_int", 0.024), - ("std_uqt_slp", 0.301), - ("std_uqs_int", 0.711), - ("std_uqs_slp", 0.458), - ("std_udrop_int", 0.716), - ("std_udrop_slp", 0.538), - ("std_urej_int", 0.011), - ("std_urej_slp", 0.002), + ("std_uqt_int", 0.077), + ("std_uqt_slp", 0.045), + ("std_uqs_int", 0.857), + ("std_uqs_slp", 0.999), + ("std_udrop_int", 0.510), + ("std_udrop_slp", 0.334), + ("std_urej_int", 0.020), + ("std_urej_slp", 0.010), ] ) SFH_PDF_FRAC_QUENCH_PDICT = OrderedDict( [ - ("frac_quench_cen_x0_tpeak", 6.501), - ("frac_quench_cen_k_tpeak", 2.281), - ("frac_quench_cen_x0_ylotpeak", 13.193), - ("frac_quench_cen_x0_yhitpeak", 12.520), - ("frac_quench_cen_ylo_ylotpeak", 0.588), + ("frac_quench_cen_x0_tpeak", 6.475), + ("frac_quench_cen_k_tpeak", 2.293), + ("frac_quench_cen_x0_ylotpeak", 11.675), + ("frac_quench_cen_x0_yhitpeak", 12.484), + ("frac_quench_cen_ylo_ylotpeak", 0.990), ("frac_quench_cen_ylo_yhitpeak", 0.001), ("frac_quench_cen_k", 4.995), - ("frac_quench_cen_yhi", 0.902), - ("frac_quench_sat_x0_tpeak", 7.829), - ("frac_quench_sat_k_tpeak", 9.990), - ("frac_quench_sat_x0_ylotpeak", 11.011), - ("frac_quench_sat_x0_yhitpeak", 12.406), + ("frac_quench_cen_yhi", 0.995), + ("frac_quench_sat_x0_tpeak", 8.624), + ("frac_quench_sat_k_tpeak", 2.214), + ("frac_quench_sat_x0_ylotpeak", 12.859), + ("frac_quench_sat_x0_yhitpeak", 12.191), ("frac_quench_sat_ylo_ylotpeak", 0.999), - ("frac_quench_sat_ylo_yhitpeak", 0.510), + ("frac_quench_sat_ylo_yhitpeak", 0.563), ("frac_quench_sat_k", 4.995), ("frac_quench_sat_yhi", 0.999), ] @@ -95,11 +103,11 @@ DELTA_UQT_PDICT = OrderedDict( [ - ("delta_uqt_x0", 1.864), - ("delta_uqt_k", 0.290), - ("delta_uqt_ylo", -0.781), - ("delta_uqt_yhi", 0.295), - ("delta_uqt_slope", 0.011), + ("delta_uqt_x0", 1.273), + ("delta_uqt_k", 0.203), + ("delta_uqt_ylo", -0.983), + ("delta_uqt_yhi", 0.384), + ("delta_uqt_slope", 0.025), ] ) SFH_PDF_QUENCH_PDICT = SFH_PDF_FRAC_QUENCH_PDICT.copy() @@ -114,7 +122,6 @@ _UPNAMES = ["u_" + key for key in QseqParams._fields] QseqUParams = namedtuple("QseqUParams", _UPNAMES) - DIFFSTARPOP_FITS_GALACTICUS_IN_DIFFSTARPOP_PARAMS = DiffstarPopParams( *SFH_PDF_QUENCH_PARAMS, *DEFAULT_SATQUENCHPOP_PARAMS ) diff --git a/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_smdpl_dr1.py b/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_smdpl_dr1.py index 8761427..f497e69 100644 --- a/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_smdpl_dr1.py +++ b/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_smdpl_dr1.py @@ -1,105 +1,113 @@ -"""""" - from collections import OrderedDict, namedtuple -from ..defaults_mgash import DiffstarPopParams, get_unbounded_diffstarpop_params +from ..defaults_mgash import ( + DiffstarPopParams, + get_unbounded_diffstarpop_params, +) + from ..satquenchpop_model import DEFAULT_SATQUENCHPOP_PARAMS SFH_PDF_QUENCH_MU_PDICT = OrderedDict( [ - ("mean_ulgm_mseq_xtp", 12.228), - ("mean_ulgm_mseq_ytp", 11.925), - ("mean_ulgm_mseq_lo", 0.733), - ("mean_ulgm_mseq_hi", -0.212), - ("mean_ulgy_mseq_int", -9.361), - ("mean_ulgy_mseq_slp", 1.696), + ("mean_ulgm_mseq_xtp", 12.842), + ("mean_ulgm_mseq_ytp", 12.447), + ("mean_ulgm_mseq_lo", 0.812), + ("mean_ulgm_mseq_hi", 0.285), + ("mean_ulgy_mseq_xtp", 11.674), + ("mean_ulgy_mseq_ytp", -10.608), + ("mean_ulgy_mseq_lo", 1.895), + ("mean_ulgy_mseq_hi", 0.781), ("mean_ul_mseq_int", -2.997), - ("mean_ul_mseq_slp", 0.021), - ("mean_uh_mseq_int", -3.647), - ("mean_uh_mseq_slp", 1.908), - ("mean_ulgm_qseq_xtp", 11.405), - ("mean_ulgm_qseq_ytp", 11.834), - ("mean_ulgm_qseq_lo", 2.487), - ("mean_ulgm_qseq_hi", 0.225), - ("mean_ulgy_qseq_int", -9.925), - ("mean_ulgy_qseq_slp", 0.626), - ("mean_ul_qseq_int", -2.999), - ("mean_ul_qseq_slp", 0.002), - ("mean_uh_qseq_int", -1.603), - ("mean_uh_qseq_slp", 0.194), - ("mean_uqt_int", 1.245), - ("mean_uqt_slp", -0.030), - ("mean_uqs_int", 0.688), - ("mean_uqs_slp", 1.099), - ("mean_udrop_int", -2.064), - ("mean_udrop_slp", -0.114), - ("mean_urej_int", -9.313), - ("mean_urej_slp", -11.160), + ("mean_ul_mseq_slp", 0.001), + ("mean_uh_mseq_int", -4.995), + ("mean_uh_mseq_slp", -2.820), + ("mean_ulgm_qseq_xtp", 11.011), + ("mean_ulgm_qseq_ytp", 12.142), + ("mean_ulgm_qseq_lo", 4.995), + ("mean_ulgm_qseq_hi", 0.010), + ("mean_ulgy_qseq_xtp", 12.836), + ("mean_ulgy_qseq_ytp", -9.690), + ("mean_ulgy_qseq_lo", 0.602), + ("mean_ulgy_qseq_hi", 0.842), + ("mean_ul_qseq_int", -2.997), + ("mean_ul_qseq_slp", 0.001), + ("mean_uh_qseq_int", -2.241), + ("mean_uh_qseq_slp", 0.606), + ("mean_uqt_xtp", 13.533), + ("mean_uqt_ytp", 1.127), + ("mean_uqt_lo", -0.011), + ("mean_uqt_hi", -0.238), + ("mean_uqs_int", 1.028), + ("mean_uqs_slp", 2.132), + ("mean_udrop_int", -2.177), + ("mean_udrop_slp", 0.308), + ("mean_urej_int", -4.516), + ("mean_urej_slp", -2.300), ] ) SFH_PDF_QUENCH_COV_MS_BLOCK_PDICT = OrderedDict( [ - ("std_ulgm_mseq_int", 0.047), - ("std_ulgm_mseq_slp", -0.062), - ("std_ulgy_mseq_int", 0.213), - ("std_ulgy_mseq_slp", 0.154), - ("std_ul_mseq_int", 0.025), - ("std_ul_mseq_slp", 0.023), - ("std_uh_mseq_int", 0.498), - ("std_uh_mseq_slp", -0.443), - ("std_ulgm_qseq_int", 0.300), - ("std_ulgm_qseq_slp", -0.252), - ("std_ulgy_qseq_int", 0.016), - ("std_ulgy_qseq_slp", -0.186), - ("std_ul_qseq_int", 0.018), - ("std_ul_qseq_slp", -0.003), - ("std_uh_qseq_int", 0.744), - ("std_uh_qseq_slp", -0.314), + ("std_ulgm_mseq_int", 0.074), + ("std_ulgm_mseq_slp", -0.106), + ("std_ulgy_mseq_int", 0.011), + ("std_ulgy_mseq_slp", 0.001), + ("std_ul_mseq_int", 0.014), + ("std_ul_mseq_slp", 0.009), + ("std_uh_mseq_int", 0.011), + ("std_uh_mseq_slp", -0.999), + ("std_ulgm_qseq_int", 0.243), + ("std_ulgm_qseq_slp", -0.162), + ("std_ulgy_qseq_int", 0.091), + ("std_ulgy_qseq_slp", -0.229), + ("std_ul_qseq_int", 0.013), + ("std_ul_qseq_slp", 0.007), + ("std_uh_qseq_int", 0.702), + ("std_uh_qseq_slp", -0.274), ] ) SFH_PDF_QUENCH_COV_Q_BLOCK_PDICT = OrderedDict( [ - ("std_uqt_int", 0.033), - ("std_uqt_slp", 0.063), - ("std_uqs_int", 0.013), - ("std_uqs_slp", -0.154), - ("std_udrop_int", 0.652), - ("std_udrop_slp", -0.998), - ("std_urej_int", 0.203), - ("std_urej_slp", -0.633), + ("std_uqt_int", 0.029), + ("std_uqt_slp", 0.041), + ("std_uqs_int", 0.011), + ("std_uqs_slp", 0.001), + ("std_udrop_int", 0.599), + ("std_udrop_slp", -0.868), + ("std_urej_int", 0.974), + ("std_urej_slp", -0.050), ] ) SFH_PDF_FRAC_QUENCH_PDICT = OrderedDict( [ - ("frac_quench_cen_x0_tpeak", 10.626), - ("frac_quench_cen_k_tpeak", 9.976), - ("frac_quench_cen_x0_ylotpeak", 11.006), - ("frac_quench_cen_x0_yhitpeak", 12.101), + ("frac_quench_cen_x0_tpeak", 11.207), + ("frac_quench_cen_k_tpeak", 9.998), + ("frac_quench_cen_x0_ylotpeak", 11.001), + ("frac_quench_cen_x0_yhitpeak", 12.606), ("frac_quench_cen_ylo_ylotpeak", 0.999), ("frac_quench_cen_ylo_yhitpeak", 0.001), - ("frac_quench_cen_k", 1.896), + ("frac_quench_cen_k", 1.954), ("frac_quench_cen_yhi", 0.999), - ("frac_quench_sat_x0_tpeak", 5.829), - ("frac_quench_sat_k_tpeak", 9.814), - ("frac_quench_sat_x0_ylotpeak", 11.418), - ("frac_quench_sat_x0_yhitpeak", 11.258), - ("frac_quench_sat_ylo_ylotpeak", 0.998), + ("frac_quench_sat_x0_tpeak", 9.372), + ("frac_quench_sat_k_tpeak", 9.995), + ("frac_quench_sat_x0_ylotpeak", 11.810), + ("frac_quench_sat_x0_yhitpeak", 11.547), + ("frac_quench_sat_ylo_ylotpeak", 0.524), ("frac_quench_sat_ylo_yhitpeak", 0.001), - ("frac_quench_sat_k", 4.997), - ("frac_quench_sat_yhi", 0.922), + ("frac_quench_sat_k", 4.995), + ("frac_quench_sat_yhi", 0.797), ] ) DELTA_UQT_PDICT = OrderedDict( [ - ("delta_uqt_x0", 10.958), - ("delta_uqt_k", 0.092), - ("delta_uqt_ylo", -0.594), - ("delta_uqt_yhi", 0.285), - ("delta_uqt_slope", -0.072), + ("delta_uqt_x0", 2.173), + ("delta_uqt_k", 0.110), + ("delta_uqt_ylo", -0.713), + ("delta_uqt_yhi", 0.217), + ("delta_uqt_slope", -0.023), ] ) SFH_PDF_QUENCH_PDICT = SFH_PDF_FRAC_QUENCH_PDICT.copy() @@ -114,7 +122,6 @@ _UPNAMES = ["u_" + key for key in QseqParams._fields] QseqUParams = namedtuple("QseqUParams", _UPNAMES) - DIFFSTARPOP_FITS_SMDPL_DR1_DIFFSTARPOP_PARAMS = DiffstarPopParams( *SFH_PDF_QUENCH_PARAMS, *DEFAULT_SATQUENCHPOP_PARAMS ) diff --git a/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_smdpl_dr1_nomerging.py b/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_smdpl_dr1_nomerging.py index e5aefe1..bfd532a 100644 --- a/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_smdpl_dr1_nomerging.py +++ b/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_smdpl_dr1_nomerging.py @@ -1,105 +1,113 @@ -"""""" - from collections import OrderedDict, namedtuple -from ..defaults_mgash import DiffstarPopParams, get_unbounded_diffstarpop_params +from ..defaults_mgash import ( + DiffstarPopParams, + get_unbounded_diffstarpop_params, +) + from ..satquenchpop_model import DEFAULT_SATQUENCHPOP_PARAMS SFH_PDF_QUENCH_MU_PDICT = OrderedDict( [ - ("mean_ulgm_mseq_xtp", 12.126), - ("mean_ulgm_mseq_ytp", 11.925), - ("mean_ulgm_mseq_lo", 0.809), - ("mean_ulgm_mseq_hi", -0.045), - ("mean_ulgy_mseq_int", -9.342), - ("mean_ulgy_mseq_slp", 1.631), - ("mean_ul_mseq_int", 3.450), - ("mean_ul_mseq_slp", 12.619), + ("mean_ulgm_mseq_xtp", 12.900), + ("mean_ulgm_mseq_ytp", 12.526), + ("mean_ulgm_mseq_lo", 0.745), + ("mean_ulgm_mseq_hi", -0.061), + ("mean_ulgy_mseq_xtp", 11.885), + ("mean_ulgy_mseq_ytp", -10.536), + ("mean_ulgy_mseq_lo", 1.644), + ("mean_ulgy_mseq_hi", 0.461), + ("mean_ul_mseq_int", -2.806), + ("mean_ul_mseq_slp", -0.321), ("mean_uh_mseq_int", -4.995), - ("mean_uh_mseq_slp", 0.424), - ("mean_ulgm_qseq_xtp", 12.547), - ("mean_ulgm_qseq_ytp", 12.283), - ("mean_ulgm_qseq_lo", 0.612), - ("mean_ulgm_qseq_hi", -0.188), - ("mean_ulgy_qseq_int", -9.907), - ("mean_ulgy_qseq_slp", 0.596), - ("mean_ul_qseq_int", -2.790), - ("mean_ul_qseq_slp", 1.725), - ("mean_uh_qseq_int", -1.461), - ("mean_uh_qseq_slp", -0.218), - ("mean_uqt_int", 0.870), - ("mean_uqt_slp", -0.319), - ("mean_uqs_int", 0.602), - ("mean_uqs_slp", -0.250), + ("mean_uh_mseq_slp", 1.479), + ("mean_ulgm_qseq_xtp", 13.349), + ("mean_ulgm_qseq_ytp", 12.240), + ("mean_ulgm_qseq_lo", 0.345), + ("mean_ulgm_qseq_hi", -0.122), + ("mean_ulgy_qseq_xtp", 11.938), + ("mean_ulgy_qseq_ytp", -9.896), + ("mean_ulgy_qseq_lo", 0.922), + ("mean_ulgy_qseq_hi", 0.450), + ("mean_ul_qseq_int", -0.597), + ("mean_ul_qseq_slp", 0.255), + ("mean_uh_qseq_int", -1.440), + ("mean_uh_qseq_slp", -0.799), + ("mean_uqt_xtp", 11.305), + ("mean_uqt_ytp", 1.181), + ("mean_uqt_lo", -0.001), + ("mean_uqt_hi", -0.445), + ("mean_uqs_int", 0.665), + ("mean_uqs_slp", 0.073), ("mean_udrop_int", -2.997), - ("mean_udrop_slp", 0.431), - ("mean_urej_int", -3.370), - ("mean_urej_slp", 1.119), + ("mean_udrop_slp", 0.960), + ("mean_urej_int", -4.668), + ("mean_urej_slp", 2.971), ] ) SFH_PDF_QUENCH_COV_MS_BLOCK_PDICT = OrderedDict( [ - ("std_ulgm_mseq_int", 0.011), - ("std_ulgm_mseq_slp", -0.127), - ("std_ulgy_mseq_int", 0.078), - ("std_ulgy_mseq_slp", 0.050), - ("std_ul_mseq_int", 1.309), - ("std_ul_mseq_slp", -0.908), - ("std_uh_mseq_int", 1.005), - ("std_uh_mseq_slp", 0.984), - ("std_ulgm_qseq_int", 0.362), - ("std_ulgm_qseq_slp", -0.065), - ("std_ulgy_qseq_int", 0.055), - ("std_ulgy_qseq_slp", -0.162), - ("std_ul_qseq_int", 1.129), - ("std_ul_qseq_slp", -0.999), - ("std_uh_qseq_int", 0.050), - ("std_uh_qseq_slp", -0.033), + ("std_ulgm_mseq_int", 0.063), + ("std_ulgm_mseq_slp", -0.058), + ("std_ulgy_mseq_int", 0.057), + ("std_ulgy_mseq_slp", 0.246), + ("std_ul_mseq_int", 0.871), + ("std_ul_mseq_slp", -0.999), + ("std_uh_mseq_int", 0.011), + ("std_uh_mseq_slp", 0.014), + ("std_ulgm_qseq_int", 0.440), + ("std_ulgm_qseq_slp", -0.309), + ("std_ulgy_qseq_int", 0.045), + ("std_ulgy_qseq_slp", 0.119), + ("std_ul_qseq_int", 0.762), + ("std_ul_qseq_slp", 0.636), + ("std_uh_qseq_int", 0.011), + ("std_uh_qseq_slp", 0.001), ] ) SFH_PDF_QUENCH_COV_Q_BLOCK_PDICT = OrderedDict( [ - ("std_uqt_int", 0.071), - ("std_uqt_slp", 0.009), - ("std_uqs_int", 0.032), - ("std_uqs_slp", 0.046), - ("std_udrop_int", 0.252), - ("std_udrop_slp", -0.405), - ("std_urej_int", 1.104), - ("std_urej_slp", -0.999), + ("std_uqt_int", 0.054), + ("std_uqt_slp", 0.021), + ("std_uqs_int", 0.077), + ("std_uqs_slp", 0.074), + ("std_udrop_int", 0.288), + ("std_udrop_slp", -0.999), + ("std_urej_int", 0.965), + ("std_urej_slp", -0.998), ] ) SFH_PDF_FRAC_QUENCH_PDICT = OrderedDict( [ - ("frac_quench_cen_x0_tpeak", 10.225), + ("frac_quench_cen_x0_tpeak", 9.731), ("frac_quench_cen_k_tpeak", 9.989), - ("frac_quench_cen_x0_ylotpeak", 11.011), - ("frac_quench_cen_x0_yhitpeak", 12.579), - ("frac_quench_cen_ylo_ylotpeak", 0.222), - ("frac_quench_cen_ylo_yhitpeak", 0.223), - ("frac_quench_cen_k", 4.995), - ("frac_quench_cen_yhi", 0.979), - ("frac_quench_sat_x0_tpeak", 3.998), - ("frac_quench_sat_k_tpeak", 4.400), - ("frac_quench_sat_x0_ylotpeak", 11.870), - ("frac_quench_sat_x0_yhitpeak", 11.456), - ("frac_quench_sat_ylo_ylotpeak", 0.999), - ("frac_quench_sat_ylo_yhitpeak", 0.203), + ("frac_quench_cen_x0_ylotpeak", 13.998), + ("frac_quench_cen_x0_yhitpeak", 13.654), + ("frac_quench_cen_ylo_ylotpeak", 0.999), + ("frac_quench_cen_ylo_yhitpeak", 0.302), + ("frac_quench_cen_k", 2.169), + ("frac_quench_cen_yhi", 0.999), + ("frac_quench_sat_x0_tpeak", 9.464), + ("frac_quench_sat_k_tpeak", 0.587), + ("frac_quench_sat_x0_ylotpeak", 11.097), + ("frac_quench_sat_x0_yhitpeak", 13.986), + ("frac_quench_sat_ylo_ylotpeak", 0.501), + ("frac_quench_sat_ylo_yhitpeak", 0.001), ("frac_quench_sat_k", 4.995), - ("frac_quench_sat_yhi", 0.887), + ("frac_quench_sat_yhi", 0.999), ] ) DELTA_UQT_PDICT = OrderedDict( [ ("delta_uqt_x0", 1.001), - ("delta_uqt_k", 0.673), - ("delta_uqt_ylo", -0.857), - ("delta_uqt_yhi", 0.131), - ("delta_uqt_slope", -0.041), + ("delta_uqt_k", 0.578), + ("delta_uqt_ylo", -0.569), + ("delta_uqt_yhi", 0.309), + ("delta_uqt_slope", 0.098), ] ) SFH_PDF_QUENCH_PDICT = SFH_PDF_FRAC_QUENCH_PDICT.copy() @@ -114,7 +122,6 @@ _UPNAMES = ["u_" + key for key in QseqParams._fields] QseqUParams = namedtuple("QseqUParams", _UPNAMES) - DIFFSTARPOP_FITS_SMDPL_DIFFSTARPOP_PARAMS = DiffstarPopParams( *SFH_PDF_QUENCH_PARAMS, *DEFAULT_SATQUENCHPOP_PARAMS ) diff --git a/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_tng.py b/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_tng.py index 022ff64..468dc41 100644 --- a/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_tng.py +++ b/diffstar/diffstarpop/kernels/params/params_diffstarpopfits_mgash_tng.py @@ -1,103 +1,113 @@ from collections import OrderedDict, namedtuple -from ..defaults_mgash import DiffstarPopParams, get_unbounded_diffstarpop_params +from ..defaults_mgash import ( + DiffstarPopParams, + get_unbounded_diffstarpop_params, +) + from ..satquenchpop_model import DEFAULT_SATQUENCHPOP_PARAMS SFH_PDF_QUENCH_MU_PDICT = OrderedDict( [ - ("mean_ulgm_mseq_xtp", 12.142), - ("mean_ulgm_mseq_ytp", 11.376), - ("mean_ulgm_mseq_lo", 0.489), - ("mean_ulgm_mseq_hi", 0.279), - ("mean_ulgy_mseq_int", -9.181), - ("mean_ulgy_mseq_slp", 1.319), - ("mean_ul_mseq_int", 0.442), - ("mean_ul_mseq_slp", 1.391), - ("mean_uh_mseq_int", -2.218), - ("mean_uh_mseq_slp", -1.589), - ("mean_ulgm_qseq_xtp", 13.299), - ("mean_ulgm_qseq_ytp", 11.894), - ("mean_ulgm_qseq_lo", 0.036), - ("mean_ulgm_qseq_hi", 0.361), - ("mean_ulgy_qseq_int", -9.558), - ("mean_ulgy_qseq_slp", 0.568), - ("mean_ul_qseq_int", -2.997), - ("mean_ul_qseq_slp", -1.791), - ("mean_uh_qseq_int", -1.426), - ("mean_uh_qseq_slp", -0.129), - ("mean_uqt_int", 0.983), - ("mean_uqt_slp", -0.358), - ("mean_uqs_int", 1.406), - ("mean_uqs_slp", -0.414), + ("mean_ulgm_mseq_xtp", 11.934), + ("mean_ulgm_mseq_ytp", 11.392), + ("mean_ulgm_mseq_lo", 0.654), + ("mean_ulgm_mseq_hi", 0.409), + ("mean_ulgy_mseq_xtp", 12.716), + ("mean_ulgy_mseq_ytp", -9.372), + ("mean_ulgy_mseq_lo", 0.691), + ("mean_ulgy_mseq_hi", 0.999), + ("mean_ul_mseq_int", 0.350), + ("mean_ul_mseq_slp", 2.462), + ("mean_uh_mseq_int", -2.144), + ("mean_uh_mseq_slp", -0.700), + ("mean_ulgm_qseq_xtp", 13.590), + ("mean_ulgm_qseq_ytp", 11.911), + ("mean_ulgm_qseq_lo", 0.218), + ("mean_ulgm_qseq_hi", 0.294), + ("mean_ulgy_qseq_xtp", 12.038), + ("mean_ulgy_qseq_ytp", -9.768), + ("mean_ulgy_qseq_lo", 1.350), + ("mean_ulgy_qseq_hi", 0.597), + ("mean_ul_qseq_int", -1.536), + ("mean_ul_qseq_slp", -0.132), + ("mean_uh_qseq_int", -1.129), + ("mean_uh_qseq_slp", -0.216), + ("mean_uqt_xtp", 13.551), + ("mean_uqt_ytp", 0.698), + ("mean_uqt_lo", -0.332), + ("mean_uqt_hi", -0.018), + ("mean_uqs_int", 1.324), + ("mean_uqs_slp", 0.301), ("mean_udrop_int", -2.997), - ("mean_udrop_slp", 1.363), - ("mean_urej_int", -9.691), - ("mean_urej_slp", -0.406), + ("mean_udrop_slp", 1.093), + ("mean_urej_int", -9.199), + ("mean_urej_slp", -0.854), ] ) SFH_PDF_QUENCH_COV_MS_BLOCK_PDICT = OrderedDict( [ ("std_ulgm_mseq_int", 0.011), - ("std_ulgm_mseq_slp", 0.037), + ("std_ulgm_mseq_slp", 0.001), ("std_ulgy_mseq_int", 0.011), - ("std_ulgy_mseq_slp", 0.001), - ("std_ul_mseq_int", 0.123), - ("std_ul_mseq_slp", 0.346), - ("std_uh_mseq_int", 0.054), + ("std_ulgy_mseq_slp", -0.184), + ("std_ul_mseq_int", 0.085), + ("std_ul_mseq_slp", 0.994), + ("std_uh_mseq_int", 0.103), ("std_uh_mseq_slp", -0.999), - ("std_ulgm_qseq_int", 0.015), - ("std_ulgm_qseq_slp", -0.435), - ("std_ulgy_qseq_int", 0.036), - ("std_ulgy_qseq_slp", -0.103), - ("std_ul_qseq_int", 0.161), - ("std_ul_qseq_slp", -0.466), - ("std_uh_qseq_int", 0.577), - ("std_uh_qseq_slp", -0.268), + ("std_ulgm_qseq_int", 0.011), + ("std_ulgm_qseq_slp", -0.175), + ("std_ulgy_qseq_int", 0.011), + ("std_ulgy_qseq_slp", -0.122), + ("std_ul_qseq_int", 0.014), + ("std_ul_qseq_slp", -0.002), + ("std_uh_qseq_int", 0.575), + ("std_uh_qseq_slp", -0.469), ] ) SFH_PDF_QUENCH_COV_Q_BLOCK_PDICT = OrderedDict( [ - ("std_uqt_int", 0.066), - ("std_uqt_slp", -0.002), - ("std_uqs_int", 0.353), - ("std_uqs_slp", -0.253), + ("std_uqt_int", 0.073), + ("std_uqt_slp", -0.009), + ("std_uqs_int", 0.042), + ("std_uqs_slp", -0.035), ("std_udrop_int", 0.011), - ("std_udrop_slp", 0.662), - ("std_urej_int", 0.100), - ("std_urej_slp", -0.074), + ("std_udrop_slp", 0.753), + ("std_urej_int", 0.127), + ("std_urej_slp", -0.063), ] ) SFH_PDF_FRAC_QUENCH_PDICT = OrderedDict( [ - ("frac_quench_cen_x0_tpeak", 10.224), - ("frac_quench_cen_k_tpeak", 2.036), - ("frac_quench_cen_x0_ylotpeak", 13.985), - ("frac_quench_cen_x0_yhitpeak", 12.398), - ("frac_quench_cen_ylo_ylotpeak", 0.999), - ("frac_quench_cen_ylo_yhitpeak", 0.234), - ("frac_quench_cen_k", 4.291), + ("frac_quench_cen_x0_tpeak", 13.615), + ("frac_quench_cen_k_tpeak", 9.240), + ("frac_quench_cen_x0_ylotpeak", 13.966), + ("frac_quench_cen_x0_yhitpeak", 11.541), + ("frac_quench_cen_ylo_ylotpeak", 0.041), + ("frac_quench_cen_ylo_yhitpeak", 0.737), + ("frac_quench_cen_k", 4.995), ("frac_quench_cen_yhi", 0.999), - ("frac_quench_sat_x0_tpeak", 10.161), - ("frac_quench_sat_k_tpeak", 9.993), - ("frac_quench_sat_x0_ylotpeak", 13.082), - ("frac_quench_sat_x0_yhitpeak", 12.443), + ("frac_quench_sat_x0_tpeak", 11.905), + ("frac_quench_sat_k_tpeak", 4.158), + ("frac_quench_sat_x0_ylotpeak", 12.469), + ("frac_quench_sat_x0_yhitpeak", 12.456), ("frac_quench_sat_ylo_ylotpeak", 0.999), - ("frac_quench_sat_ylo_yhitpeak", 0.002), + ("frac_quench_sat_ylo_yhitpeak", 0.001), ("frac_quench_sat_k", 4.995), - ("frac_quench_sat_yhi", 0.843), + ("frac_quench_sat_yhi", 0.999), ] ) DELTA_UQT_PDICT = OrderedDict( [ - ("delta_uqt_x0", 3.678), - ("delta_uqt_k", 0.554), - ("delta_uqt_ylo", -0.605), - ("delta_uqt_yhi", 0.050), - ("delta_uqt_slope", 0.020), + ("delta_uqt_x0", 2.532), + ("delta_uqt_k", 0.454), + ("delta_uqt_ylo", -0.977), + ("delta_uqt_yhi", -0.002), + ("delta_uqt_slope", -0.030), ] ) SFH_PDF_QUENCH_PDICT = SFH_PDF_FRAC_QUENCH_PDICT.copy() @@ -112,7 +122,6 @@ _UPNAMES = ["u_" + key for key in QseqParams._fields] QseqUParams = namedtuple("QseqUParams", _UPNAMES) - DIFFSTARPOP_FITS_TNG_DIFFSTARPOP_PARAMS = DiffstarPopParams( *SFH_PDF_QUENCH_PARAMS, *DEFAULT_SATQUENCHPOP_PARAMS ) diff --git a/diffstar/diffstarpop/kernels/sfh_pdf_mgash.py b/diffstar/diffstarpop/kernels/sfh_pdf_mgash.py index 73adb2b..fdfe8a7 100644 --- a/diffstar/diffstarpop/kernels/sfh_pdf_mgash.py +++ b/diffstar/diffstarpop/kernels/sfh_pdf_mgash.py @@ -14,6 +14,7 @@ smoothly_clipped_line, ) + TODAY = 13.8 LGT0 = jnp.log10(TODAY) @@ -24,34 +25,40 @@ SFH_PDF_QUENCH_MU_PDICT = OrderedDict( [ - ("mean_ulgm_mseq_xtp", 12.126), - ("mean_ulgm_mseq_ytp", 11.925), - ("mean_ulgm_mseq_lo", 0.809), - ("mean_ulgm_mseq_hi", -0.045), - ("mean_ulgy_mseq_int", -9.342), - ("mean_ulgy_mseq_slp", 1.631), - ("mean_ul_mseq_int", 3.450), - ("mean_ul_mseq_slp", 12.619), + ("mean_ulgm_mseq_xtp", 12.900), + ("mean_ulgm_mseq_ytp", 12.526), + ("mean_ulgm_mseq_lo", 0.745), + ("mean_ulgm_mseq_hi", -0.061), + ("mean_ulgy_mseq_xtp", 11.885), + ("mean_ulgy_mseq_ytp", -10.536), + ("mean_ulgy_mseq_lo", 1.644), + ("mean_ulgy_mseq_hi", 0.461), + ("mean_ul_mseq_int", -2.806), + ("mean_ul_mseq_slp", -0.321), ("mean_uh_mseq_int", -4.995), - ("mean_uh_mseq_slp", 0.424), - ("mean_ulgm_qseq_xtp", 12.547), - ("mean_ulgm_qseq_ytp", 12.283), - ("mean_ulgm_qseq_lo", 0.612), - ("mean_ulgm_qseq_hi", -0.188), - ("mean_ulgy_qseq_int", -9.907), - ("mean_ulgy_qseq_slp", 0.596), - ("mean_ul_qseq_int", -2.790), - ("mean_ul_qseq_slp", 1.725), - ("mean_uh_qseq_int", -1.461), - ("mean_uh_qseq_slp", -0.218), - ("mean_uqt_int", 0.870), - ("mean_uqt_slp", -0.319), - ("mean_uqs_int", 0.602), - ("mean_uqs_slp", -0.250), + ("mean_uh_mseq_slp", 1.479), + ("mean_ulgm_qseq_xtp", 13.349), + ("mean_ulgm_qseq_ytp", 12.240), + ("mean_ulgm_qseq_lo", 0.345), + ("mean_ulgm_qseq_hi", -0.122), + ("mean_ulgy_qseq_xtp", 11.938), + ("mean_ulgy_qseq_ytp", -9.896), + ("mean_ulgy_qseq_lo", 0.922), + ("mean_ulgy_qseq_hi", 0.450), + ("mean_ul_qseq_int", -0.597), + ("mean_ul_qseq_slp", 0.255), + ("mean_uh_qseq_int", -1.440), + ("mean_uh_qseq_slp", -0.799), + ("mean_uqt_xtp", 11.305), + ("mean_uqt_ytp", 1.181), + ("mean_uqt_lo", -0.001), + ("mean_uqt_hi", -0.445), + ("mean_uqs_int", 0.665), + ("mean_uqs_slp", 0.073), ("mean_udrop_int", -2.997), - ("mean_udrop_slp", 0.431), - ("mean_urej_int", -3.370), - ("mean_urej_slp", 1.119), + ("mean_udrop_slp", 0.960), + ("mean_urej_int", -4.668), + ("mean_urej_slp", 2.971), ] ) SFH_PDF_QUENCH_MU_BOUNDS_PDICT = OrderedDict( @@ -59,8 +66,10 @@ mean_ulgm_mseq_ytp=(11.0, 14.0), mean_ulgm_mseq_lo=(-1.0, 5.0), mean_ulgm_mseq_hi=(-5.0, 1.0), - mean_ulgy_mseq_int=(-13.0, -7.0), - mean_ulgy_mseq_slp=(-20.0, 20.0), + mean_ulgy_mseq_xtp=(11.0, 14.0), + mean_ulgy_mseq_ytp=(-13.0, -8.0), + mean_ulgy_mseq_lo=(-1.0, 5.0), + mean_ulgy_mseq_hi=(-5.0, 1.0), mean_ul_mseq_int=(-3.0, 5.0), mean_ul_mseq_slp=(-20.0, 20.0), mean_uh_mseq_int=(-5.0, 3.0), @@ -69,14 +78,18 @@ mean_ulgm_qseq_ytp=(11.0, 14.0), mean_ulgm_qseq_lo=(-1.0, 5.0), mean_ulgm_qseq_hi=(-5.0, 1.0), - mean_ulgy_qseq_int=(-13.0, -7.0), - mean_ulgy_qseq_slp=(-20.0, 20.0), + mean_ulgy_qseq_xtp=(11.0, 14.0), + mean_ulgy_qseq_ytp=(-13.0, -8.0), + mean_ulgy_qseq_lo=(-1.0, 5.0), + mean_ulgy_qseq_hi=(-5.0, 1.0), mean_ul_qseq_int=(-3.0, 5.0), mean_ul_qseq_slp=(-20.0, 20.0), mean_uh_qseq_int=(-5.0, 3.0), mean_uh_qseq_slp=(-20.0, 20.0), - mean_uqt_int=(0.0, 2.0), - mean_uqt_slp=(-20.0, 20.0), + mean_uqt_xtp=(11.0, 14.0), + mean_uqt_ytp=(0.0, 2.0), + mean_uqt_lo=(-1.0, 0.0), + mean_uqt_hi=(-1.0, 0.0), mean_uqs_int=(-5.0, 2.0), mean_uqs_slp=(-20.0, 20.0), mean_udrop_int=(-3.0, 2.0), @@ -87,22 +100,22 @@ SFH_PDF_QUENCH_COV_MS_BLOCK_PDICT = OrderedDict( [ - ("std_ulgm_mseq_int", 0.011), - ("std_ulgm_mseq_slp", -0.127), - ("std_ulgy_mseq_int", 0.078), - ("std_ulgy_mseq_slp", 0.050), - ("std_ul_mseq_int", 1.309), - ("std_ul_mseq_slp", -0.908), - ("std_uh_mseq_int", 1.005), - ("std_uh_mseq_slp", 0.984), - ("std_ulgm_qseq_int", 0.362), - ("std_ulgm_qseq_slp", -0.065), - ("std_ulgy_qseq_int", 0.055), - ("std_ulgy_qseq_slp", -0.162), - ("std_ul_qseq_int", 1.129), - ("std_ul_qseq_slp", -0.999), - ("std_uh_qseq_int", 0.050), - ("std_uh_qseq_slp", -0.033), + ("std_ulgm_mseq_int", 0.063), + ("std_ulgm_mseq_slp", -0.058), + ("std_ulgy_mseq_int", 0.057), + ("std_ulgy_mseq_slp", 0.246), + ("std_ul_mseq_int", 0.871), + ("std_ul_mseq_slp", -0.999), + ("std_uh_mseq_int", 0.011), + ("std_uh_mseq_slp", 0.014), + ("std_ulgm_qseq_int", 0.440), + ("std_ulgm_qseq_slp", -0.309), + ("std_ulgy_qseq_int", 0.045), + ("std_ulgy_qseq_slp", 0.119), + ("std_ul_qseq_int", 0.762), + ("std_ul_qseq_slp", 0.636), + ("std_uh_qseq_int", 0.011), + ("std_uh_qseq_slp", 0.001), ] ) SFH_PDF_QUENCH_COV_MS_BLOCK_BOUNDS_PDICT = OrderedDict( @@ -126,14 +139,14 @@ SFH_PDF_QUENCH_COV_Q_BLOCK_PDICT = OrderedDict( [ - ("std_uqt_int", 0.071), - ("std_uqt_slp", 0.009), - ("std_uqs_int", 0.032), - ("std_uqs_slp", 0.046), - ("std_udrop_int", 0.252), - ("std_udrop_slp", -0.405), - ("std_urej_int", 1.104), - ("std_urej_slp", -0.999), + ("std_uqt_int", 0.054), + ("std_uqt_slp", 0.021), + ("std_uqs_int", 0.077), + ("std_uqs_slp", 0.074), + ("std_udrop_int", 0.288), + ("std_udrop_slp", -0.999), + ("std_urej_int", 0.965), + ("std_urej_slp", -0.998), ] ) SFH_PDF_QUENCH_COV_Q_BLOCK_BOUNDS_PDICT = OrderedDict( @@ -149,22 +162,22 @@ SFH_PDF_FRAC_QUENCH_PDICT = OrderedDict( [ - ("frac_quench_cen_x0_tpeak", 10.225), + ("frac_quench_cen_x0_tpeak", 9.731), ("frac_quench_cen_k_tpeak", 9.989), - ("frac_quench_cen_x0_ylotpeak", 11.011), - ("frac_quench_cen_x0_yhitpeak", 12.579), - ("frac_quench_cen_ylo_ylotpeak", 0.222), - ("frac_quench_cen_ylo_yhitpeak", 0.223), - ("frac_quench_cen_k", 4.995), - ("frac_quench_cen_yhi", 0.979), - ("frac_quench_sat_x0_tpeak", 3.998), - ("frac_quench_sat_k_tpeak", 4.400), - ("frac_quench_sat_x0_ylotpeak", 11.870), - ("frac_quench_sat_x0_yhitpeak", 11.456), - ("frac_quench_sat_ylo_ylotpeak", 0.999), - ("frac_quench_sat_ylo_yhitpeak", 0.203), + ("frac_quench_cen_x0_ylotpeak", 13.998), + ("frac_quench_cen_x0_yhitpeak", 13.654), + ("frac_quench_cen_ylo_ylotpeak", 0.999), + ("frac_quench_cen_ylo_yhitpeak", 0.302), + ("frac_quench_cen_k", 2.169), + ("frac_quench_cen_yhi", 0.999), + ("frac_quench_sat_x0_tpeak", 9.464), + ("frac_quench_sat_k_tpeak", 0.587), + ("frac_quench_sat_x0_ylotpeak", 11.097), + ("frac_quench_sat_x0_yhitpeak", 13.986), + ("frac_quench_sat_ylo_ylotpeak", 0.501), + ("frac_quench_sat_ylo_yhitpeak", 0.001), ("frac_quench_sat_k", 4.995), - ("frac_quench_sat_yhi", 0.887), + ("frac_quench_sat_yhi", 0.999), ] ) SFH_PDF_FRAC_QUENCH_BOUNDS_PDICT = OrderedDict( @@ -210,10 +223,10 @@ DELTA_UQT_PDICT = OrderedDict( [ ("delta_uqt_x0", 1.001), - ("delta_uqt_k", 0.673), - ("delta_uqt_ylo", -0.857), - ("delta_uqt_yhi", 0.131), - ("delta_uqt_slope", -0.041), + ("delta_uqt_k", 0.578), + ("delta_uqt_ylo", -0.569), + ("delta_uqt_yhi", 0.309), + ("delta_uqt_slope", 0.098), ] ) DELTA_UQT_BOUNDS_PDICT = OrderedDict( @@ -316,11 +329,12 @@ def _get_mean_u_params_mseq(params, logmp0): params.mean_ulgm_mseq_hi, ) - ulgy = line_model( + ulgy = Mcrit_model( logmp0, - params.mean_ulgy_mseq_int, - params.mean_ulgy_mseq_slp, - *BOUNDING_VALS.mean_ulgy, + params.mean_ulgy_mseq_xtp, + params.mean_ulgy_mseq_ytp, + params.mean_ulgy_mseq_lo, + params.mean_ulgy_mseq_hi, ) ul = line_model( @@ -351,11 +365,12 @@ def _get_mean_u_params_qseq(params, logmp0, tpeak): params.mean_ulgm_qseq_hi, ) - ulgy = line_model( + ulgy = Mcrit_model( logmp0, - params.mean_ulgy_qseq_int, - params.mean_ulgy_qseq_slp, - *BOUNDING_VALS.mean_ulgy, + params.mean_ulgy_qseq_xtp, + params.mean_ulgy_qseq_ytp, + params.mean_ulgy_qseq_lo, + params.mean_ulgy_qseq_hi, ) ul = line_model( @@ -372,11 +387,12 @@ def _get_mean_u_params_qseq(params, logmp0, tpeak): *BOUNDING_VALS.mean_uh, ) - _uqt = line_model( + _uqt = Mcrit_model( logmp0, - params.mean_uqt_int, - params.mean_uqt_slp, - *BOUNDING_VALS.mean_uqt, + params.mean_uqt_xtp, + params.mean_uqt_ytp, + params.mean_uqt_lo, + params.mean_uqt_hi, ) delta_uqt = _delta_uqt(params, logmp0, tpeak) uqt = _uqt + delta_uqt diff --git a/diffstar/diffstarpop/kernels/tests/test_satquench_model.py b/diffstar/diffstarpop/kernels/tests/test_satquench_model.py index d0f949e..cbdd1d0 100644 --- a/diffstar/diffstarpop/kernels/tests/test_satquench_model.py +++ b/diffstar/diffstarpop/kernels/tests/test_satquench_model.py @@ -1,5 +1,4 @@ -""" -""" +""" """ import numpy as np from jax import random as jran diff --git a/diffstar/diffstarpop/sumstats/smdpl_smhm_targets.py b/diffstar/diffstarpop/sumstats/smdpl_smhm_targets.py index 992c014..239ab86 100644 --- a/diffstar/diffstarpop/sumstats/smdpl_smhm_targets.py +++ b/diffstar/diffstarpop/sumstats/smdpl_smhm_targets.py @@ -1,5 +1,4 @@ -""" -""" +""" """ import os diff --git a/diffstar/diffstarpop/sumstats/tests/test_smhm.py b/diffstar/diffstarpop/sumstats/tests/test_smhm.py index 277a30f..c3f4bd8 100644 --- a/diffstar/diffstarpop/sumstats/tests/test_smhm.py +++ b/diffstar/diffstarpop/sumstats/tests/test_smhm.py @@ -1,5 +1,4 @@ -""" -""" +""" """ import numpy as np diff --git a/docs/source/demo_diffmahpop_diffstarpop_sfh.ipynb b/docs/source/demo_diffmahpop_diffstarpop_sfh.ipynb new file mode 100644 index 0000000..0b95dbc --- /dev/null +++ b/docs/source/demo_diffmahpop_diffstarpop_sfh.ipynb @@ -0,0 +1,450 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "2f5e4d41", + "metadata": {}, + "source": [ + "# Generating halo and galaxy populations with DiffmahPop and DiffstarPop\n", + "\n", + "This notebook gives a basic illustrations of how to use `DiffmahPop` and `DiffstarPop` withing the `diffsky` pipeline to generate a catalog of simulated mass accretion histories and star formation histories." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "dfda3a65", + "metadata": {}, + "outputs": [], + "source": [ + "%matplotlib inline" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "740ca990", + "metadata": {}, + "outputs": [], + "source": [ + "import numpy as np\n", + "from matplotlib import pyplot as plt\n", + "import matplotlib.gridspec as gridspec\n", + "from matplotlib.patches import Patch\n", + "from matplotlib.lines import Line2D\n", + "from scipy.optimize import curve_fit\n", + "mred = u\"#d62728\"\n", + "morange = u\"#ff7f0e\"\n", + "mgreen = u\"#2ca02c\"\n", + "mblue = u\"#1f77b4\"\n", + "mpurple = u\"#9467bd\"\n", + "\n", + "colors = [mpurple,mblue,mgreen,morange, mred]" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "49e88ec7", + "metadata": {}, + "outputs": [], + "source": [ + "import subprocess\n", + "\n", + "\n", + "def try_enable_latex():\n", + " \"\"\"Try enabling LaTeX text rendering in matplotlib,\n", + " fallback if not available.\"\"\"\n", + " try:\n", + " # Quick check: can we run latex?\n", + " subprocess.check_call(\n", + " [\"latex\", \"--version\"],\n", + " stdout=subprocess.DEVNULL,\n", + " stderr=subprocess.DEVNULL,\n", + " )\n", + " plt.rc(\"text\", usetex=True)\n", + " plt.rc(\"font\", family=\"serif\", size=22)\n", + " plt.rc('figure', figsize=(6,4)) \n", + " print(\"LaTeX rendering enabled.\")\n", + " except (subprocess.CalledProcessError, FileNotFoundError):\n", + " # LaTeX not installed or failed\n", + " plt.rc(\"text\", usetex=False)\n", + " plt.rc(\"font\", family=\"serif\", size=22)\n", + " plt.rc('figure', figsize=(6,4)) \n", + " print(\"LaTeX not available, falling back to default mathtext.\")\n", + "\n", + "try_enable_latex()" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "7de791ef", + "metadata": {}, + "outputs": [], + "source": [ + "\n", + "from jax import jit as jjit\n", + "from jax import numpy as jnp\n", + "from jax import random as jran\n", + "from jax import value_and_grad\n", + "\n", + "# Some constants\n", + "from dsps.constants import T_TABLE_MIN\n", + "from diffstar.defaults import LGT0, TODAY\n" + ] + }, + { + "cell_type": "markdown", + "id": "af7e46f5", + "metadata": {}, + "source": [ + "### Generate a simulated halo catalog with `diffmahpop`\n", + "\n", + "Here we generate a simulated Monte Carlo realization of a subhalo catalog at a single redshift using the diffmahpop wrapper in the diffsky pipeline. \n", + "\n", + "The `DiffmahPop` model is a generator of mass accretion histories for a population of haloes, using the diffmah model.\n", + "\n", + "Please note that the diffmahpop wrapper within the `diffsky` repo is likely to change in the future, since it is in a development stage.\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "27d5a751", + "metadata": {}, + "outputs": [], + "source": [ + "from diffsky.mass_functions.mc_diffmah_tpeak import mc_subhalos\n", + "\n", + "ran_key = jran.PRNGKey(0)\n", + "\n", + "# Generate a random subhalo catalog\n", + "subcat_key, ran_key = jran.split(ran_key, 2)\n", + "lgmp_min = 11.25\n", + "z_obs = 0.01\n", + "Lbox = 75.0\n", + "volume_com = Lbox**3\n", + "subcat = mc_subhalos(subcat_key, z_obs, lgmp_min=lgmp_min, volume_com=volume_com)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "ce913bc0", + "metadata": {}, + "outputs": [], + "source": [ + "plt.hist(subcat.logmp0, np.linspace(lgmp_min, 14.5, 100))\n", + "plt.xlabel(r\"$\\log M_{p,0}\\, [M_{\\odot}]$\")\n", + "plt.show()" + ] + }, + { + "cell_type": "markdown", + "id": "150fcab3", + "metadata": {}, + "source": [ + "### Load `DiffstarPop` parameters fitted to different simulations\n", + "\n", + "DiffstarPop has a set of default parameters, but it also has a set of parameters that best-fit three different types of simulations:\n", + "\n", + "- The hydrodynamical simulation IllustrisTNG.\n", + "- The semi-analyitical model (SAM) Galaticus.\n", + "- The empirical model UniverseMachine.\n", + "\n", + "These can be accessed through the `diffstar.diffstarpop.kernels.params` module.\n", + "\n", + "They can be used to place informed priors from galaxy formation models into the SFHs of galaxies." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "ed459a74", + "metadata": {}, + "outputs": [], + "source": [ + "from diffstar.diffstarpop import DEFAULT_DIFFSTARPOP_PARAMS\n", + "from diffstar.diffstarpop.kernels.params import (\n", + " DiffstarPop_Params_Diffstarpopfits_mgash, \n", + ")\n", + "print(DiffstarPop_Params_Diffstarpopfits_mgash.keys())\n", + "\n", + "# These are fit for in-situ only SFHs.\n", + "DIFFSTARPOP_UM = DiffstarPop_Params_Diffstarpopfits_mgash[\"smdpl_dr1_nomerging\"]\n", + "DIFFSTARPOP_TNG = DiffstarPop_Params_Diffstarpopfits_mgash[\"tng\"]\n", + "DIFFSTARPOP_GALCUS = DiffstarPop_Params_Diffstarpopfits_mgash[\"galacticus_in_situ\"]\n", + "\n", + "# These are fit for in-plus-ex-situ SFHs.\n", + "DIFFSTARPOP_UM_plus_exsitu = DiffstarPop_Params_Diffstarpopfits_mgash[\"smdpl_dr1\"]\n", + "DIFFSTARPOP_GALCUS_plus_exsitu = DiffstarPop_Params_Diffstarpopfits_mgash[\"galacticus_in_plus_ex_situ\"]\n" + ] + }, + { + "cell_type": "markdown", + "id": "a24d8558", + "metadata": {}, + "source": [ + "### Calculate the SFHs using the UniverseMachine params\n", + "\n", + "Here we calculate the MAH from the halo parameters that we previously generated. \n", + "\n", + "We generate star formation histories for each halo. A galaxy in a given halo can be in the main sequence or quenched, with a certain probability given by the quenching fraction. We generate both, along with the `frac_q` that each halo might be quenched. We also generate a monte-carlo realization for the halo catalog, assigning whereas each halo is hosting a main sequence or quenched galaxy, via `mc_is_q`." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "174f463a", + "metadata": {}, + "outputs": [], + "source": [ + "from diffmah.diffmah_kernels import mah_halopop\n", + "from diffstar.diffstarpop import mc_diffstar_sfh_galpop\n", + "\n", + "# Create a table of times where to calculate the MAH and SFH\n", + "ntimes = 50\n", + "tarr = np.linspace(T_TABLE_MIN, TODAY, ntimes)\n", + "\n", + "# Calculate the mass accreation history of every halo\n", + "dmhdt_fit, log_mah_fit = mah_halopop(subcat.mah_params, tarr, LGT0)\n", + "\n", + "# Manually set the infall data of each halo to no infall,\n", + "# since this part of the Diffstarpop model has not been calibrated\n", + "\n", + "n_halos = subcat.logmhost_ult_inf.shape[0]\n", + "\n", + "lgmu_infall = -1.0 * np.ones(n_halos)\n", + "logmhost_infall = 13.0 * np.ones(n_halos)\n", + "gyr_since_infall = -99.0 * np.ones(n_halos)\n", + "\n", + "# compute SFHs for the default galaxy population\n", + "args = (\n", + " DIFFSTARPOP_UM,\n", + " subcat.mah_params,\n", + " subcat.logmp0,\n", + " subcat.upids,\n", + " lgmu_infall,\n", + " logmhost_infall,\n", + " gyr_since_infall,\n", + " ran_key,\n", + " tarr,\n", + ")\n", + "\n", + "(\n", + " diffstar_params_ms,\n", + " diffstar_params_q,\n", + " default_sfh_ms,\n", + " default_sfh_q,\n", + " frac_q,\n", + " mc_is_q,\n", + ") = mc_diffstar_sfh_galpop(*args)\n", + "\n", + "# select at random if a galaxy is MS or Q based on frac_q.\n", + "default_sfh = np.zeros_like(default_sfh_ms) \n", + "default_sfh[mc_is_q] = default_sfh_q[mc_is_q]\n", + "default_sfh[~mc_is_q] = default_sfh_ms[~mc_is_q]" + ] + }, + { + "cell_type": "markdown", + "id": "8008bb83", + "metadata": {}, + "source": [ + "### Make some plots\n", + "\n", + "Here we plot the average MAH and SFH histories we have generated, for halos of different present-day halo mass.\n", + "\n", + "We also plot a few individual MAH and SFH histories." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "5d28d160", + "metadata": {}, + "outputs": [], + "source": [ + "fig, ax = plt.subplots(1,2, figsize=(16,6))\n", + "mpeak_vals = np.arange(11.5, 14, 0.5)\n", + "for i, mpeak in enumerate(mpeak_vals):\n", + " sel = (subcat.logmp0 > mpeak - 0.2) & (subcat.logmp0 < mpeak + 0.2)\n", + " mean_log_mah_fit = np.mean(log_mah_fit[sel], axis=0)\n", + " mean_default_sfh = np.mean(default_sfh[sel], axis=0)\n", + " range_log_mah_fit = np.percentile(log_mah_fit[sel], [15.865, 84.135], axis=0)\n", + " range_mean_default_sfh = np.percentile(default_sfh[sel], [15.865, 84.135], axis=0)\n", + " \n", + " ax[0].plot(tarr, 10**mean_log_mah_fit, color=colors[i])\n", + " ax[0].fill_between(tarr, 10**range_log_mah_fit[0], 10**range_log_mah_fit[1], color=colors[i], alpha=0.1)\n", + " ax[1].plot(tarr, mean_default_sfh, color=colors[i])\n", + " ax[1].fill_between(tarr, range_mean_default_sfh[0], range_mean_default_sfh[1], color=colors[i], alpha=0.1)\n", + "\n", + "\n", + "ax[0].set_ylim(1e9, 1e14)\n", + "ax[0].set_xticks(np.arange(1.0, 14.0, 2.0))\n", + "ax[0].set_xlabel(\"Cosmic time [Gyr]\")\n", + "ax[1].set_xlabel(\"Cosmic time [Gyr]\")\n", + "ax[0].set_yscale('log')\n", + "\n", + "ax[0].set_xlim(0.0, 13.7)\n", + "ax[1].set_xlim(0.0, 13.7)\n", + "ax[1].set_yscale('log')\n", + "ax[1].set_ylim(5e-3, 5e2)\n", + "ax[1].set_ylabel(r\"$\\langle \\dot{M}_\\star | M_{p,0} \\rangle [M_{\\odot}/{\\rm yr}]$\")\n", + "ax[0].set_ylabel(r\"$\\langle M_{\\rm halo} | M_{p,0} \\rangle [M_{\\odot}]$\")\n", + "\n", + "fig.subplots_adjust(wspace=0.25)\n", + "fig.suptitle(\"Average histories\")\n", + "\n", + "plt.show()\n", + "\n", + "fig, ax = plt.subplots(1,2, figsize=(16,6))\n", + "for i, mpeak in enumerate(mpeak_vals):\n", + " sel = (subcat.logmp0 > mpeak - 0.2) & (subcat.logmp0 < mpeak + 0.2)\n", + "\n", + " ax[0].plot(tarr, 10**log_mah_fit[sel][np.random.choice(int(sel.sum()), 5)].T, color=colors[i])\n", + " ax[1].plot(tarr, default_sfh[sel][np.random.choice(int(sel.sum()), 5)].T, color=colors[i])\n", + "\n", + "\n", + "\n", + "ax[0].set_yscale('log')\n", + "ax[1].set_yscale('log')\n", + "ax[1].set_ylim(5e-3, 5e2)\n", + "\n", + "ax[0].set_xlim(0.0, 13.7)\n", + "ax[1].set_xlim(0.0, 13.7)\n", + "ax[0].set_ylim(1e9, 1e14)\n", + "ax[0].set_xlabel(\"Cosmic time [Gyr]\")\n", + "ax[1].set_xlabel(\"Cosmic time [Gyr]\")\n", + "ax[0].set_ylabel(r\"$ M_{\\rm halo} | M_{p,0} \\,[M_{\\odot}]$\")\n", + "ax[1].set_ylabel(r\"$\\dot{M}_\\star | M_{p,0} \\,[M_{\\odot}/{\\rm yr}]$\")\n", + "fig.subplots_adjust(wspace=0.25)\n", + "fig.suptitle(\"Individual samples\")\n", + "plt.show()" + ] + }, + { + "cell_type": "markdown", + "id": "3fefff10", + "metadata": {}, + "source": [ + "### Compare SFHs for `DiffstarPop` params fitted to different simulations.\n", + "\n", + "Here we plot the average SFHs using the best-fit diffstarpop values obtained from different simulations." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "00d9abad", + "metadata": {}, + "outputs": [], + "source": [ + "fig, ax = plt.subplots(1,4, figsize=(30,6), sharex=True)\n", + "mpeak_vals = np.arange(11.5, 14, 0.5)\n", + "\n", + "for i, mpeak in enumerate(mpeak_vals):\n", + " sel = (subcat.logmp0 > mpeak - 0.2) & (subcat.logmp0 < mpeak + 0.2)\n", + " mean_log_mah_fit = np.mean(log_mah_fit[sel], axis=0)\n", + " range_log_mah_fit = np.percentile(log_mah_fit[sel], [15.865, 84.135], axis=0)\n", + " \n", + " ax[0].plot(tarr, mean_log_mah_fit, color=colors[i])\n", + " ax[0].fill_between(tarr, range_log_mah_fit[0], range_log_mah_fit[1], color=colors[i], alpha=0.1)\n", + "\n", + "for k, params in enumerate([DIFFSTARPOP_UM, DIFFSTARPOP_TNG, DIFFSTARPOP_GALCUS]):\n", + " # compute SFHs for the default galaxy population\n", + " args = (\n", + " params,\n", + " subcat.mah_params,\n", + " subcat.logmp0,\n", + " subcat.upids,\n", + " lgmu_infall,\n", + " logmhost_infall,\n", + " gyr_since_infall,\n", + " ran_key,\n", + " tarr,\n", + " )\n", + "\n", + " (\n", + " diffstar_params_ms,\n", + " diffstar_params_q,\n", + " default_sfh_ms,\n", + " default_sfh_q,\n", + " frac_q,\n", + " mc_is_q,\n", + " ) = mc_diffstar_sfh_galpop(*args)\n", + "\n", + " default_sfh = np.zeros_like(default_sfh_ms) \n", + " default_sfh[mc_is_q] = default_sfh_q[mc_is_q]\n", + " default_sfh[~mc_is_q] = default_sfh_ms[~mc_is_q]\n", + "\n", + " for i, mpeak in enumerate(mpeak_vals):\n", + " sel = (subcat.logmp0 > mpeak - 0.2) & (subcat.logmp0 < mpeak + 0.2)\n", + " mean_default_sfh = np.mean(default_sfh[sel], axis=0)\n", + " range_mean_default_sfh = np.percentile(default_sfh[sel], [15.865, 84.135], axis=0)\n", + " \n", + " ax[k+1].plot(tarr, mean_default_sfh, color=colors[i])\n", + " ax[k+1].fill_between(tarr, range_mean_default_sfh[0], range_mean_default_sfh[1], color=colors[i], alpha=0.1)\n", + "\n", + "\n", + "ax[0].set_ylim(9, 14)\n", + "ax[0].set_xticks(np.arange(1.0, 14.0, 2.0))\n", + "ax[0].set_xlabel(\"Cosmic time [Gyr]\")\n", + "ax[1].set_xlabel(\"Cosmic time [Gyr]\")\n", + "ax[2].set_xlabel(\"Cosmic time [Gyr]\")\n", + "ax[3].set_xlabel(\"Cosmic time [Gyr]\")\n", + "ax[0].set_xlim(0.0, 13.7)\n", + "\n", + "ax[1].set_yscale('log')\n", + "ax[2].set_yscale('log')\n", + "ax[3].set_yscale('log')\n", + "ax[1].set_ylim(5e-3, 5e2)\n", + "ax[2].set_ylim(5e-3, 5e2)\n", + "ax[3].set_ylim(5e-3, 5e2)\n", + "ax[0].set_ylabel(r\"$\\langle M_{\\rm halo} | M_{p,0} \\rangle [M_{\\odot}]$\")\n", + "ax[1].set_ylabel(r\"$\\langle \\dot{M}_\\star | M_{p,0} \\rangle [M_{\\odot}/{\\rm yr}]$\")\n", + "ax[2].set_ylabel(r\"$\\langle \\dot{M}_\\star | M_{p,0} \\rangle [M_{\\odot}/{\\rm yr}]$\")\n", + "ax[3].set_ylabel(r\"$\\langle \\dot{M}_\\star | M_{p,0} \\rangle [M_{\\odot}/{\\rm yr}]$\")\n", + "\n", + "ax[1].set_title(\"UniverseMachine\")\n", + "ax[2].set_title(\"IllustrisTNG\")\n", + "ax[3].set_title(\"Galacticus\")\n", + "\n", + "fig.subplots_adjust(wspace=0.25)\n", + "\n", + "\n", + "plt.show()" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "471a75b8", + "metadata": {}, + "outputs": [], + "source": [] + } + ], + "metadata": { + "kernelspec": { + "display_name": "diffstuff", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.11.9" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/docs/source/demo_diffstar_sfh.ipynb b/docs/source/demo_diffstar_sfh.ipynb index a900723..416b8f9 100644 --- a/docs/source/demo_diffstar_sfh.ipynb +++ b/docs/source/demo_diffstar_sfh.ipynb @@ -14,6 +14,41 @@ "In the cell below, we'll use the default diffmah and diffstar parameters, and then use the `sfh_singlegal` function to calculate the SFH." ] }, + { + "cell_type": "code", + "execution_count": null, + "id": "e4643163", + "metadata": {}, + "outputs": [], + "source": [ + "import subprocess\n", + "from matplotlib import pyplot as plt\n", + "\n", + "\n", + "def try_enable_latex():\n", + " \"\"\"Try enabling LaTeX text rendering in matplotlib,\n", + " fallback if not available.\"\"\"\n", + " try:\n", + " # Quick check: can we run latex?\n", + " subprocess.check_call(\n", + " [\"latex\", \"--version\"],\n", + " stdout=subprocess.DEVNULL,\n", + " stderr=subprocess.DEVNULL,\n", + " )\n", + " plt.rc(\"text\", usetex=True)\n", + " plt.rc(\"font\", family=\"serif\", size=22)\n", + " plt.rc('figure', figsize=(6,4)) \n", + " print(\"LaTeX rendering enabled.\")\n", + " except (subprocess.CalledProcessError, FileNotFoundError):\n", + " # LaTeX not installed or failed\n", + " plt.rc(\"text\", usetex=False)\n", + " plt.rc(\"font\", family=\"serif\", size=22)\n", + " plt.rc('figure', figsize=(6,4)) \n", + " print(\"LaTeX not available, falling back to default mathtext.\")\n", + "\n", + "try_enable_latex()" + ] + }, { "cell_type": "code", "execution_count": null, @@ -49,7 +84,6 @@ "metadata": {}, "outputs": [], "source": [ - "from matplotlib import pyplot as plt\n", "\n", "fig, ax = plt.subplots(1, 1)\n", "ylim = ax.set_ylim(2e-3, 50)\n", @@ -58,7 +92,8 @@ "__=ax.plot(tarr, sfh_gal, color='k')\n", "\n", "xlabel = ax.set_xlabel(r'${\\rm cosmic\\ time\\ [Gyr]}$')\n", - "ylabel = ax.set_ylabel(r'${\\rm SFR\\ [M_{\\odot}/yr]}$')" + "ylabel = ax.set_ylabel(r'${\\rm SFR\\ [M_{\\odot}/yr]}$')\n", + "ax.set_xticks(np.arange(1.0, 14.0, 2.0))" ] }, { @@ -167,21 +202,14 @@ "\n", "\n", "xlabel = ax.set_xlabel(r'${\\rm cosmic\\ time\\ [Gyr]}$')\n", - "ylabel = ax.set_ylabel(r'${\\rm SFR\\ [M_{\\odot}/yr]}$')" + "ylabel = ax.set_ylabel(r'${\\rm SFR\\ [M_{\\odot}/yr]}$')\n", + "ax.set_xticks(np.arange(1.0, 14.0, 2.0))\n" ] - }, - { - "cell_type": "code", - "execution_count": null, - "id": "1b780056", - "metadata": {}, - "outputs": [], - "source": [] } ], "metadata": { "kernelspec": { - "display_name": "Python 3 (ipykernel)", + "display_name": "diffstuff", "language": "python", "name": "python3" }, @@ -195,7 +223,7 @@ "name": "python", "nbconvert_exporter": "python", "pygments_lexer": "ipython3", - "version": "3.12.9" + "version": "3.11.9" } }, "nbformat": 4, diff --git a/docs/source/index.rst b/docs/source/index.rst index 28e1eb7..e14cb5e 100644 --- a/docs/source/index.rst +++ b/docs/source/index.rst @@ -20,6 +20,7 @@ User Guide installation.rst demo_diffstar_sfh.ipynb + demo_diffmahpop_diffstarpop_sfh.ipynb reference.rst See :ref:`Citation Information ` for how to acknowledge Diffstar. diff --git a/docs/source/rtd_environment.yaml b/docs/source/rtd_environment.yaml index d9b877c..4c56d42 100644 --- a/docs/source/rtd_environment.yaml +++ b/docs/source/rtd_environment.yaml @@ -15,4 +15,6 @@ dependencies: - jax - jaxlib - diffmah>=0.7.0 - - h5py \ No newline at end of file + - h5py + - dsps + - diffsky \ No newline at end of file diff --git a/scripts/diffstarpop_scripts/fit_get_loss_helpers_mgash.py b/scripts/diffstarpop_scripts/fit_get_loss_helpers_mgash.py new file mode 100644 index 0000000..d8c8bb8 --- /dev/null +++ b/scripts/diffstarpop_scripts/fit_get_loss_helpers_mgash.py @@ -0,0 +1,1374 @@ +""" """ + +import os +import h5py +import numpy as np +from diffmah.diffmah_kernels import DiffmahParams, mah_halopop +from diffstar.defaults import LGT0 +from jax import random as jran +from jax import numpy as jnp + + +def get_loss_data_smhm(indir, nhalos): + # Load SMHM data --------------------------------------------- + print("Loading SMHM data...") + + with h5py.File(indir + "smdpl_smhm.h5", "r") as hdf: + redshift_targets = hdf["redshift_targets"][:] + # smhm_diff = hdf["smhm_diff"][:] + smhm = hdf["smhm"][:] + logmh_bins = hdf["logmh_bins"][:] + age_targets = hdf["age_targets"][:] + """ + hdfout["counts_diff"] = wcounts + hdfout["hist_diff"] = whist + hdfout["counts"] = counts + hdfout["hist"] = hist + hdfout["smhm_diff"] = whist / wcounts + hdfout["smhm"] = hist / counts + hdfout["logmh_bins"] = smhm_utils.LOGMH_BINS + hdfout["subvol_used"] = subvol_used + """ + + logmh_binsc = 0.5 * (logmh_bins[1:] + logmh_bins[:-1]) + + with h5py.File(indir + "smdpl_smhm_samples_haloes.h5", "r") as hdf: + logmh_id = hdf["logmh_id"][:] + # logmh_val = hdf["logmh_id"][:] + mah_params_samp = hdf["mah_params_samp"][:] + upid_samp = hdf["upid_samp"][:] + # ms_params_samp = hdf["ms_params_samp"][:] + # q_params_samp = hdf["q_params_samp"][:] + tobs_id = hdf["tobs_id"][:] + # tobs_val = hdf["tobs_val"][:] + # redshift_val = hdf["redshift_val"][:] + + # Create loss_data --------------------------------------------- + print("Creating loss data...") + + ran_key = jran.PRNGKey(np.random.randint(2**32)) + + lgmu_infall = -1.0 + logmhost_infall = 13.0 + gyr_since_infall = -99.0 # 2.0 + + mah_params_data = [] + logmp0_data = [] + upid_data = [] + lgmu_infall_data = [] + logmhost_infall_data = [] + gyr_since_infall_data = [] + t_obs_targets = [] + smhm_targets = [] + + tarr_logm0 = np.logspace(-1, LGT0, 50) + + for i in range(len(age_targets)): + t_target = age_targets[i] + + for j in range(len(logmh_binsc)): + sel = (tobs_id == i) & (logmh_id == j) + + if sel.sum() < nhalos: + continue + arange_sel = np.arange(len(tobs_id))[sel] + arange_sel = np.random.choice(arange_sel, nhalos, replace=False) + mah_params_data.append(mah_params_samp[:, arange_sel]) + upid_data.append(upid_samp[arange_sel]) + lgmu_infall_data.append(np.ones(len(arange_sel)) * lgmu_infall) + logmhost_infall_data.append(np.ones(len(arange_sel)) * logmhost_infall) + gyr_since_infall_data.append(np.ones(len(arange_sel)) * gyr_since_infall) + t_obs_targets.append(t_target) + smhm_targets.append(smhm[i, j]) + mah_pars_ntuple = DiffmahParams(*mah_params_samp[:, arange_sel]) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, LGT0) + logmp0_data.append(log_mah_fit[:, -1]) + + mah_params_data = np.array(mah_params_data) + logmp0_data = np.array(logmp0_data) + upid_data = np.array(upid_data) + lgmu_infall_data = np.array(lgmu_infall_data) + logmhost_infall_data = np.array(logmhost_infall_data) + gyr_since_infall_data = np.array(gyr_since_infall_data) + t_obs_targets = np.array(t_obs_targets) + smhm_targets = np.array(smhm_targets) + + ran_key_data = jran.split(ran_key, len(smhm_targets)) + loss_data = ( + mah_params_data, + logmp0_data, + lgmu_infall_data, + logmhost_infall_data, + gyr_since_infall_data, + ran_key_data, + t_obs_targets, + smhm_targets, + ) + + plot_data = ( + age_targets, + logmh_binsc, + tobs_id, + logmh_id, + tarr_logm0, + lgmu_infall, + logmhost_infall, + gyr_since_infall, + ran_key, + redshift_targets, + smhm, + mah_params_samp, + upid_samp, + ) + + return loss_data, plot_data + + +def get_loss_data_pdfs_mstar(indir, nhalos): + # Load PDF data --------------------------------------------- + print("Loading PDF Mstar data...") + + fname = os.path.join(indir, "smdpl_smhm.h5") + with h5py.File(fname, "r") as hdf: + redshift_targets = hdf["redshift_targets"][:] + smhm_diff = hdf["smhm_diff"][:] + smhm = hdf["smhm"][:] + logmh_bins = hdf["logmh_bins"][:] + age_targets = hdf["age_targets"][:] + hist = hdf["hist"][:] + counts = hdf["counts"][:] + counts_cen = hdf["counts_cen"][:] + counts_sat = hdf["counts_sat"][:] + + logmh_binsc = 0.5 * (logmh_bins[1:] + logmh_bins[:-1]) + + fname = os.path.join(indir, "smdpl_smhm_samples_haloes.h5") + with h5py.File(fname, "r") as hdf: + logmh_id = hdf["logmh_id"][:] + logmh_val = hdf["logmh_id"][:] + mah_params_samp = hdf["mah_params_samp"][:] + ms_params_samp = hdf["ms_params_samp"][:] + q_params_samp = hdf["q_params_samp"][:] + upid_samp = hdf["upid_samp"][:] + tobs_id = hdf["tobs_id"][:] + tobs_val = hdf["tobs_val"][:] + redshift_val = hdf["redshift_val"][:] + + fname = os.path.join(indir, "smdpl_mstar_ssfr.h5") + with h5py.File(fname, "r") as hdf: + mstar_wcounts = hdf["mstar_wcounts"][:] + mstar_counts = hdf["mstar_counts"][:] + mstar_ssfr_wcounts_cent = hdf["mstar_ssfr_wcounts_cent"][:] + mstar_ssfr_wcounts_sat = hdf["mstar_ssfr_wcounts_sat"][:] + logssfr_bins_pdf = hdf["logssfr_bins_pdf"][:] + logmstar_bins_pdf = hdf["logmstar_bins_pdf"][:] + """ + hdfout["mstar_wcounts"] = mstar_wcounts + hdfout["mstar_counts"] = mstar_counts + hdfout["mstar_ssfr_wcounts_cent"] = mstar_ssfr_wcounts_cent + hdfout["mstar_ssfr_wcounts_sat"] = mstar_ssfr_wcounts_sat + hdfout["logmh_bins"] = smhm_utils.LOGMH_BINS + hdfout["logmstar_bins_pdf"] = smhm_utils.LOGMSTAR_BINS_PDF + hdfout["logssfr_bins_pdf"] = smhm_utils.LOGSSFR_BINS_PDF + hdfout["redshift_targets"] = redshift_targets + hdfout["age_targets"] = age_targets + """ + + logssfr_binsc_pdf = 0.5 * (logssfr_bins_pdf[1:] + logssfr_bins_pdf[:-1]) + logmstar_binsc_pdf = 0.5 * (logmstar_bins_pdf[1:] + logmstar_bins_pdf[:-1]) + + # Create loss_data --------------------------------------------- + print("Creating loss data...") + + ran_key = jran.PRNGKey(np.random.randint(2**32)) + + lgmu_infall = -1.0 + logmhost_infall = 13.0 + gyr_since_infall = -99.0 # 2.0 + + mah_params_data = [] + logmp0_data = [] + upid_data = [] + lgmu_infall_data = [] + logmhost_infall_data = [] + gyr_since_infall_data = [] + t_obs_targets = [] + mstar_counts_target = [] + + tarr_logm0 = np.logspace(-1, LGT0, 50) + + for i in range(len(age_targets)): + t_target = age_targets[i] + + for j in range(len(logmh_binsc)): + sel = (tobs_id == i) & (logmh_id == j) + + if sel.sum() < 50: + continue + replace = True if sel.sum() < nhalos else False + arange_sel = np.arange(len(tobs_id))[sel] + arange_sel = np.random.choice(arange_sel, nhalos, replace=replace) + mah_params_data.append(mah_params_samp[:, arange_sel]) + upid_data.append(upid_samp[arange_sel]) + lgmu_infall_data.append(np.ones(len(arange_sel)) * lgmu_infall) + logmhost_infall_data.append(np.ones(len(arange_sel)) * logmhost_infall) + gyr_since_infall_data.append(np.ones(len(arange_sel)) * gyr_since_infall) + t_obs_targets.append(t_target) + mstar_counts_target.append(mstar_wcounts[i, j] / mstar_wcounts[i, j].sum()) + mah_pars_ntuple = DiffmahParams(*mah_params_samp[:, arange_sel]) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, LGT0) + logmp0_data.append(log_mah_fit[:, -1]) + # break + + mah_params_data = np.array(mah_params_data) + logmp0_data = np.array(logmp0_data) + upid_data = np.array(upid_data) + lgmu_infall_data = np.array(lgmu_infall_data) + logmhost_infall_data = np.array(logmhost_infall_data) + gyr_since_infall_data = np.array(gyr_since_infall_data) + t_obs_targets = np.array(t_obs_targets) + mstar_counts_target = np.array(mstar_counts_target) + + ran_key_data = jran.split(ran_key, len(mstar_counts_target)) + loss_data_mstar = ( + mah_params_data, + logmp0_data, + upid_data, + lgmu_infall_data, + logmhost_infall_data, + gyr_since_infall_data, + ran_key_data, + t_obs_targets, + logmstar_bins_pdf, + mstar_counts_target, + ) + + # Create loss_data for plot --------------------------------------------- + print("Creating loss data for plot...") + nhalos_plot = 10000 + ran_key = jran.PRNGKey(np.random.randint(2**32)) + + lgmu_infall = -1.0 + logmhost_infall = 13.0 + gyr_since_infall = -99.0 # 2.0 + + mah_params_data = [] + logmp0_data = [] + upid_data = [] + lgmu_infall_data = [] + logmhost_infall_data = [] + gyr_since_infall_data = [] + t_obs_targets = [] + mstar_counts_target = [] + + tarr_logm0 = np.logspace(-1, LGT0, 50) + + for i in range(len(age_targets)): + t_target = age_targets[i] + + for j in range(len(logmh_binsc)): + sel = (tobs_id == i) & (logmh_id == j) + + if sel.sum() < 50: + continue + replace = True if sel.sum() < nhalos_plot else False + arange_sel = np.arange(len(tobs_id))[sel] + arange_sel = np.random.choice(arange_sel, nhalos_plot, replace=replace) + mah_params_data.append(mah_params_samp[:, arange_sel]) + upid_data.append(upid_samp[arange_sel]) + lgmu_infall_data.append(np.ones(len(arange_sel)) * lgmu_infall) + logmhost_infall_data.append(np.ones(len(arange_sel)) * logmhost_infall) + gyr_since_infall_data.append(np.ones(len(arange_sel)) * gyr_since_infall) + t_obs_targets.append(t_target) + mstar_counts_target.append(mstar_wcounts[i, j] / mstar_wcounts[i, j].sum()) + mah_pars_ntuple = DiffmahParams(*mah_params_samp[:, arange_sel]) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, LGT0) + logmp0_data.append(log_mah_fit[:, -1]) + # break + + mah_params_data = np.array(mah_params_data) + logmp0_data = np.array(logmp0_data) + upid_data = np.array(upid_data) + lgmu_infall_data = np.array(lgmu_infall_data) + logmhost_infall_data = np.array(logmhost_infall_data) + gyr_since_infall_data = np.array(gyr_since_infall_data) + t_obs_targets = np.array(t_obs_targets) + mstar_counts_target = np.array(mstar_counts_target) + + ran_key_data = jran.split(ran_key, len(mstar_counts_target)) + loss_data_mstar_pred = ( + mah_params_data, + logmp0_data, + upid_data, + lgmu_infall_data, + logmhost_infall_data, + gyr_since_infall_data, + ran_key_data, + t_obs_targets, + logmstar_bins_pdf, + mstar_counts_target, + ) + plot_data = ( + logmstar_bins_pdf, + mstar_wcounts, + age_targets, + redshift_targets, + tobs_id, + logmh_id, + logmh_binsc, + loss_data_mstar_pred, + ) + + return loss_data_mstar, plot_data + + +def prepare_ragged(indx_pdf, nmhalo_pdf, index_mhalo): + """Run this outside jit once per dataset to create dense, shape-stable arrays.""" + nz = len(indx_pdf) + Mmax = max(len(ix) for ix in indx_pdf) + + # Build dense (nz, Mmax) arrays for indices, weights, and mask + idx_np = np.zeros((nz, Mmax), dtype=jnp.int32) + w_np = np.zeros((nz, Mmax), dtype=nmhalo_pdf.dtype) + # msk_np = np.zeros((nz, Mmax), dtype=bool) + + for z in range(nz): + m = len(indx_pdf[z]) + idx_np[z, :m] = jnp.asarray(indx_pdf[z]) + w_np[z, :m] = jnp.asarray(nmhalo_pdf[z, index_mhalo[z]]) + # msk_np[z, :m] = True + idx = jnp.asarray(idx_np) + w = jnp.asarray(w_np) + # msk = jnp.asarray(msk_np) + # return idx, w, msk # shapes: (nz, Mmax), (nz, Mmax), (nz, Mmax) + return idx, w # shapes: (nz, Mmax), (nz, Mmax) + + +def get_loss_data_pdfs_ssfr_central(indir, nhalos): + + print("Loading PDF Mstar data...") + + fname = os.path.join(indir, "smdpl_smhm.h5") + with h5py.File(fname, "r") as hdf: + redshift_targets = hdf["redshift_targets"][:] + smhm_diff = hdf["smhm_diff"][:] + smhm = hdf["smhm"][:] + logmh_bins = hdf["logmh_bins"][:] + age_targets = hdf["age_targets"][:] + hist = hdf["hist"][:] + counts = hdf["counts"][:] + counts_cen = hdf["counts_cen"][:] + counts_sat = hdf["counts_sat"][:] + + logmh_binsc = 0.5 * (logmh_bins[1:] + logmh_bins[:-1]) + + fname = os.path.join(indir, "smdpl_smhm_samples_haloes.h5") + with h5py.File(fname, "r") as hdf: + logmh_id = hdf["logmh_id"][:] + logmh_val = hdf["logmh_id"][:] + mah_params_samp = hdf["mah_params_samp"][:] + ms_params_samp = hdf["ms_params_samp"][:] + q_params_samp = hdf["q_params_samp"][:] + upid_samp = hdf["upid_samp"][:] + tobs_id = hdf["tobs_id"][:] + tobs_val = hdf["tobs_val"][:] + redshift_val = hdf["redshift_val"][:] + + fname = os.path.join(indir, "smdpl_mstar_ssfr.h5") + with h5py.File(fname, "r") as hdf: + mstar_wcounts = hdf["mstar_wcounts"][:] + mstar_counts = hdf["mstar_counts"][:] + mstar_ssfr_wcounts_cent = hdf["mstar_ssfr_wcounts_cent"][:] + mstar_ssfr_wcounts_sat = hdf["mstar_ssfr_wcounts_sat"][:] + logssfr_bins_pdf = hdf["logssfr_bins_pdf"][:] + logmstar_bins_pdf = hdf["logmstar_bins_pdf"][:] + """ + hdfout["mstar_wcounts"] = mstar_wcounts + hdfout["mstar_counts"] = mstar_counts + hdfout["mstar_ssfr_wcounts_cent"] = mstar_ssfr_wcounts_cent + hdfout["mstar_ssfr_wcounts_sat"] = mstar_ssfr_wcounts_sat + hdfout["logmh_bins"] = smhm_utils.LOGMH_BINS + hdfout["logmstar_bins_pdf"] = smhm_utils.LOGMSTAR_BINS_PDF + hdfout["logssfr_bins_pdf"] = smhm_utils.LOGSSFR_BINS_PDF + hdfout["redshift_targets"] = redshift_targets + hdfout["age_targets"] = age_targets + """ + + logssfr_binsc_pdf = 0.5 * (logssfr_bins_pdf[1:] + logssfr_bins_pdf[:-1]) + logmstar_binsc_pdf = 0.5 * (logmstar_bins_pdf[1:] + logmstar_bins_pdf[:-1]) + + mhalo_pdf_hist = hist / np.sum(hist, axis=1)[:, None] + mhalo_pdf = counts / np.sum(counts, axis=1)[:, None] + mhalo_pdf_cen = counts_cen / np.sum(counts_cen, axis=1)[:, None] + mhalo_pdf_sat = counts_sat / np.sum(counts_sat, axis=1)[:, None] + + # Create loss_data --------------------------------------------- + print("Creating loss data...") + + ran_key = jran.PRNGKey(np.random.randint(2**32)) + + lgmu_infall = -1.0 + logmhost_infall = 13.0 + gyr_since_infall = -99.0 # 2.0 + + mah_params_data = [] + logmp0_data = [] + upid_data = [] + lgmu_infall_data = [] + logmhost_infall_data = [] + gyr_since_infall_data = [] + t_obs_targets = [] + + mstar_ssfr_pdfs_cent = np.clip(mstar_ssfr_wcounts_cent, 0.0, None) + mstar_ssfr_pdfs_cent = ( + mstar_ssfr_pdfs_cent + / np.sum(mstar_ssfr_pdfs_cent, axis=(2, 3))[:, :, None, None] + ) + mstar_ssfr_pdfs_cent = np.where( + np.isnan(mstar_ssfr_pdfs_cent), 0.0, mstar_ssfr_pdfs_cent + ) + mstar_ssfr_pdfs_cent = np.einsum( + "zmab,zm->zab", mstar_ssfr_pdfs_cent, mhalo_pdf_cen + ) + + tarr_logm0 = np.logspace(-1, LGT0, 50) + + index_mhalo = [] + indx_pdf = [] + _run_indx = 0 + for i in range(len(age_targets)): + t_target = age_targets[i] + index_mhalo_atz = [] + indx_pdf_atz = [] + for j in range(len(logmh_binsc)): + sel = (tobs_id == i) & (logmh_id == j) & (upid_samp == -1) + + if sel.sum() < 50: + print(i, j) + continue + arange_sel = np.arange(len(tobs_id))[sel] + replace = True if sel.sum() < nhalos else False + arange_sel = np.random.choice(arange_sel, nhalos, replace=replace) + mah_params_data.append(mah_params_samp[:, arange_sel]) + upid_data.append(upid_samp[arange_sel]) + lgmu_infall_data.append(np.ones(len(arange_sel)) * lgmu_infall) + logmhost_infall_data.append(np.ones(len(arange_sel)) * logmhost_infall) + gyr_since_infall_data.append(np.ones(len(arange_sel)) * gyr_since_infall) + t_obs_targets.append(t_target) + mah_pars_ntuple = DiffmahParams(*mah_params_samp[:, arange_sel]) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, LGT0) + logmp0_data.append(log_mah_fit[:, -1]) + + index_mhalo_atz.append(j) + indx_pdf_atz.append(_run_indx) + _run_indx += 1 + + index_mhalo.append(np.array(index_mhalo_atz)) + indx_pdf.append(np.array(indx_pdf_atz)) + # break + # target_mstar_ids = np.array([4, 9, 14, 17, 19, 22]) + # target_mstar_ids = np.array([9, 14, 17, 19, 22]) + target_mstar_ids = np.array([8, 10, 13, 16, 19]) + # target_mstar_ids = np.array([9, 12, 15, 17, 19]) + # target_mstar_ids = np.array([14, 17, 19, 22]) + print(logmstar_binsc_pdf[target_mstar_ids]) + target_data = np.zeros( + (len(age_targets), len(target_mstar_ids), len(logssfr_binsc_pdf)) + ) + for i in range(len(age_targets)): + for j, jval in enumerate(target_mstar_ids): + if mstar_ssfr_pdfs_cent[i, jval].sum() > 0.0: + target_data[i, j] = ( + mstar_ssfr_pdfs_cent[i, jval] / mstar_ssfr_pdfs_cent[i, jval].sum() + ) + + mah_params_data = np.array(mah_params_data) + logmp0_data = np.array(logmp0_data) + upid_data = np.array(upid_data) + lgmu_infall_data = np.array(lgmu_infall_data) + logmhost_infall_data = np.array(logmhost_infall_data) + gyr_since_infall_data = np.array(gyr_since_infall_data) + t_obs_targets = np.array(t_obs_targets) + + ran_key_data = jran.split(ran_key, len(t_obs_targets)) + + ndbins_lo = [] + ndbins_hi = [] + for i in range(len(logmstar_bins_pdf) - 1): + for j in range(len(logssfr_bins_pdf) - 1): + ndbins_lo.append([logmstar_bins_pdf[i], logssfr_bins_pdf[j]]) + ndbins_hi.append([logmstar_bins_pdf[i + 1], logssfr_bins_pdf[j + 1]]) + ndbins_lo = np.array(ndbins_lo) + ndbins_hi = np.array(ndbins_hi) + + indx_pdf, mhalo_pdf_cen_ragged = prepare_ragged( + indx_pdf, mhalo_pdf_cen, index_mhalo + ) + + loss_data_ssfr = ( + mah_params_data, + logmp0_data, + upid_data, + lgmu_infall_data, + logmhost_infall_data, + gyr_since_infall_data, + ran_key_data, + t_obs_targets, + ndbins_lo, + ndbins_hi, + logmstar_bins_pdf, + logssfr_bins_pdf, + mhalo_pdf_cen_ragged, + indx_pdf, + jnp.asarray(target_mstar_ids), + target_data, + ) + + # Create loss_data for plot --------------------------------------------- + print("Creating loss data for plot...") + nhalos_plot = 10000 + + ran_key = jran.PRNGKey(np.random.randint(2**32)) + + lgmu_infall = -1.0 + logmhost_infall = 13.0 + gyr_since_infall = -99.0 # 2.0 + + mah_params_data = [] + logmp0_data = [] + upid_data = [] + lgmu_infall_data = [] + logmhost_infall_data = [] + gyr_since_infall_data = [] + t_obs_targets = [] + + mstar_ssfr_pdfs_cent = np.clip(mstar_ssfr_wcounts_cent, 0.0, None) + mstar_ssfr_pdfs_cent = ( + mstar_ssfr_pdfs_cent + / np.sum(mstar_ssfr_pdfs_cent, axis=(2, 3))[:, :, None, None] + ) + mstar_ssfr_pdfs_cent = np.where( + np.isnan(mstar_ssfr_pdfs_cent), 0.0, mstar_ssfr_pdfs_cent + ) + mstar_ssfr_pdfs_cent = np.einsum( + "zmab,zm->zab", mstar_ssfr_pdfs_cent, mhalo_pdf_cen + ) + + index_mhalo = [] + indx_pdf = [] + _run_indx = 0 + for i in range(len(age_targets)): + t_target = age_targets[i] + index_mhalo_atz = [] + indx_pdf_atz = [] + for j in range(len(logmh_binsc)): + sel = (tobs_id == i) & (logmh_id == j) & (upid_samp == -1) + + if sel.sum() < 50: + print(i, j, t_target, logmh_binsc[j], sel.sum()) + continue + arange_sel = np.arange(len(tobs_id))[sel] + replace = True if sel.sum() < nhalos_plot else False + arange_sel = np.random.choice(arange_sel, nhalos_plot, replace=replace) + mah_params_data.append(mah_params_samp[:, arange_sel]) + upid_data.append(upid_samp[arange_sel]) + lgmu_infall_data.append(np.ones(len(arange_sel)) * lgmu_infall) + logmhost_infall_data.append(np.ones(len(arange_sel)) * logmhost_infall) + gyr_since_infall_data.append(np.ones(len(arange_sel)) * gyr_since_infall) + t_obs_targets.append(t_target) + mah_pars_ntuple = DiffmahParams(*mah_params_samp[:, arange_sel]) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, LGT0) + logmp0_data.append(log_mah_fit[:, -1]) + + index_mhalo_atz.append(j) + indx_pdf_atz.append(_run_indx) + _run_indx += 1 + + index_mhalo.append(np.array(index_mhalo_atz)) + indx_pdf.append(np.array(indx_pdf_atz)) + # break + # target_mstar_ids = np.array([4, 9, 14, 17, 19, 22]) + # target_mstar_ids = np.array([9, 14, 17, 19, 22]) + target_mstar_ids = np.array([8, 10, 13, 16, 19]) + # target_mstar_ids = np.array([9, 12, 15, 17, 19]) + # target_mstar_ids = np.array([14, 17, 19, 22]) + print(logmstar_binsc_pdf[target_mstar_ids]) + target_data = np.zeros( + (len(age_targets), len(target_mstar_ids), len(logssfr_binsc_pdf)) + ) + for i in range(len(age_targets)): + for j, jval in enumerate(target_mstar_ids): + if mstar_ssfr_pdfs_cent[i, jval].sum() > 0.0: + target_data[i, j] = ( + mstar_ssfr_pdfs_cent[i, jval] / mstar_ssfr_pdfs_cent[i, jval].sum() + ) + + mah_params_data = np.array(mah_params_data) + logmp0_data = np.array(logmp0_data) + upid_data = np.array(upid_data) + lgmu_infall_data = np.array(lgmu_infall_data) + logmhost_infall_data = np.array(logmhost_infall_data) + gyr_since_infall_data = np.array(gyr_since_infall_data) + t_obs_targets = np.array(t_obs_targets) + + ran_key_data = jran.split(ran_key, len(t_obs_targets)) + + ndbins_lo = [] + ndbins_hi = [] + for i in range(len(logmstar_bins_pdf) - 1): + for j in range(len(logssfr_bins_pdf) - 1): + ndbins_lo.append([logmstar_bins_pdf[i], logssfr_bins_pdf[j]]) + ndbins_hi.append([logmstar_bins_pdf[i + 1], logssfr_bins_pdf[j + 1]]) + ndbins_lo = np.array(ndbins_lo) + ndbins_hi = np.array(ndbins_hi) + + indx_pdf, mhalo_pdf_cen_ragged = prepare_ragged( + indx_pdf, mhalo_pdf_cen, index_mhalo + ) + + loss_data_ssfr_pred = ( + mah_params_data, + logmp0_data, + upid_data, + lgmu_infall_data, + logmhost_infall_data, + gyr_since_infall_data, + ran_key_data, + t_obs_targets, + ndbins_lo, + ndbins_hi, + logmstar_bins_pdf, + logssfr_bins_pdf, + mhalo_pdf_cen_ragged, + indx_pdf, + jnp.asarray(target_mstar_ids), + target_data, + ) + + plot_data = ( + jnp.asarray(target_mstar_ids), + logssfr_binsc_pdf, + target_data, + loss_data_ssfr_pred, + ) + + return loss_data_ssfr, plot_data + + +def get_loss_data_pdfs_ssfr_satellite(indir, nhalos): + + print("Loading PDF Mstar data...") + + fname = os.path.join(indir, "smdpl_smhm.h5") + with h5py.File(fname, "r") as hdf: + redshift_targets = hdf["redshift_targets"][:] + smhm_diff = hdf["smhm_diff"][:] + smhm = hdf["smhm"][:] + logmh_bins = hdf["logmh_bins"][:] + age_targets = hdf["age_targets"][:] + hist = hdf["hist"][:] + counts = hdf["counts"][:] + counts_cen = hdf["counts_cen"][:] + counts_sat = hdf["counts_sat"][:] + + logmh_binsc = 0.5 * (logmh_bins[1:] + logmh_bins[:-1]) + + fname = os.path.join(indir, "smdpl_smhm_samples_haloes.h5") + with h5py.File(fname, "r") as hdf: + logmh_id = hdf["logmh_id"][:] + logmh_val = hdf["logmh_id"][:] + mah_params_samp = hdf["mah_params_samp"][:] + ms_params_samp = hdf["ms_params_samp"][:] + q_params_samp = hdf["q_params_samp"][:] + upid_samp = hdf["upid_samp"][:] + tobs_id = hdf["tobs_id"][:] + tobs_val = hdf["tobs_val"][:] + redshift_val = hdf["redshift_val"][:] + + fname = os.path.join(indir, "smdpl_mstar_ssfr.h5") + with h5py.File(fname, "r") as hdf: + mstar_wcounts = hdf["mstar_wcounts"][:] + mstar_counts = hdf["mstar_counts"][:] + mstar_ssfr_wcounts_cent = hdf["mstar_ssfr_wcounts_cent"][:] + mstar_ssfr_wcounts_sat = hdf["mstar_ssfr_wcounts_sat"][:] + logssfr_bins_pdf = hdf["logssfr_bins_pdf"][:] + logmstar_bins_pdf = hdf["logmstar_bins_pdf"][:] + """ + hdfout["mstar_wcounts"] = mstar_wcounts + hdfout["mstar_counts"] = mstar_counts + hdfout["mstar_ssfr_wcounts_cent"] = mstar_ssfr_wcounts_cent + hdfout["mstar_ssfr_wcounts_sat"] = mstar_ssfr_wcounts_sat + hdfout["logmh_bins"] = smhm_utils.LOGMH_BINS + hdfout["logmstar_bins_pdf"] = smhm_utils.LOGMSTAR_BINS_PDF + hdfout["logssfr_bins_pdf"] = smhm_utils.LOGSSFR_BINS_PDF + hdfout["redshift_targets"] = redshift_targets + hdfout["age_targets"] = age_targets + """ + + logssfr_binsc_pdf = 0.5 * (logssfr_bins_pdf[1:] + logssfr_bins_pdf[:-1]) + logmstar_binsc_pdf = 0.5 * (logmstar_bins_pdf[1:] + logmstar_bins_pdf[:-1]) + + mhalo_pdf_hist = hist / np.sum(hist, axis=1)[:, None] + mhalo_pdf = counts / np.sum(counts, axis=1)[:, None] + mhalo_pdf_cen = counts_cen / np.sum(counts_cen, axis=1)[:, None] + mhalo_pdf_sat = counts_sat / np.sum(counts_sat, axis=1)[:, None] + + # Create loss_data --------------------------------------------- + print("Creating loss data...") + + ran_key = jran.PRNGKey(np.random.randint(2**32)) + + lgmu_infall = -1.0 + logmhost_infall = 13.0 + gyr_since_infall = -99.0 # 2.0 + + mah_params_data = [] + logmp0_data = [] + upid_data = [] + lgmu_infall_data = [] + logmhost_infall_data = [] + gyr_since_infall_data = [] + t_obs_targets = [] + + mstar_ssfr_pdfs_sat = np.clip(mstar_ssfr_wcounts_sat, 0.0, None) + mstar_ssfr_pdfs_sat = ( + mstar_ssfr_pdfs_sat / np.sum(mstar_ssfr_pdfs_sat, axis=(2, 3))[:, :, None, None] + ) + mstar_ssfr_pdfs_sat = np.where( + np.isnan(mstar_ssfr_pdfs_sat), 0.0, mstar_ssfr_pdfs_sat + ) + mstar_ssfr_pdfs_sat = np.einsum("zmab,zm->zab", mstar_ssfr_pdfs_sat, mhalo_pdf_sat) + + index_mhalo = [] + indx_pdf = [] + _run_indx = 0 + + tarr_logm0 = np.logspace(-1, LGT0, 50) + for i in range(len(age_targets)): + t_target = age_targets[i] + index_mhalo_atz = [] + indx_pdf_atz = [] + for j in range(len(logmh_binsc)): + sel = (tobs_id == i) & (logmh_id == j) & (upid_samp != -1) + + if sel.sum() < 50: + print(i, j) + continue + arange_sel = np.arange(len(tobs_id))[sel] + replace = True if sel.sum() < nhalos else False + arange_sel = np.random.choice(arange_sel, nhalos, replace=replace) + mah_params_data.append(mah_params_samp[:, arange_sel]) + upid_data.append(upid_samp[arange_sel]) + lgmu_infall_data.append(np.ones(len(arange_sel)) * lgmu_infall) + logmhost_infall_data.append(np.ones(len(arange_sel)) * logmhost_infall) + gyr_since_infall_data.append(np.ones(len(arange_sel)) * gyr_since_infall) + t_obs_targets.append(t_target) + mah_pars_ntuple = DiffmahParams(*mah_params_samp[:, arange_sel]) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, LGT0) + logmp0_data.append(log_mah_fit[:, -1]) + + index_mhalo_atz.append(j) + indx_pdf_atz.append(_run_indx) + _run_indx += 1 + + index_mhalo.append(np.array(index_mhalo_atz)) + indx_pdf.append(np.array(indx_pdf_atz)) + + # break + # target_mstar_ids = np.array([4, 9, 14, 17, 19, 22]) + # target_mstar_ids = np.array([9, 14, 17, 19, 22]) + target_mstar_ids = np.array([8, 10, 13, 16, 19]) + # target_mstar_ids = np.array([9, 12, 15, 17]) + # target_mstar_ids = np.array([14, 17, 19, 22]) + print(logmstar_binsc_pdf[target_mstar_ids]) + target_data_sat = np.zeros( + (len(age_targets), len(target_mstar_ids), len(logssfr_binsc_pdf)) + ) + for i in range(len(age_targets)): + for j, jval in enumerate(target_mstar_ids): + if mstar_ssfr_pdfs_sat[i, jval].sum() > 0.0: + target_data_sat[i, j] = ( + mstar_ssfr_pdfs_sat[i, jval] / mstar_ssfr_pdfs_sat[i, jval].sum() + ) + + mah_params_data = np.array(mah_params_data) + logmp0_data = np.array(logmp0_data) + upid_data = np.array(upid_data) + lgmu_infall_data = np.array(lgmu_infall_data) + logmhost_infall_data = np.array(logmhost_infall_data) + gyr_since_infall_data = np.array(gyr_since_infall_data) + t_obs_targets = np.array(t_obs_targets) + + ran_key_data = jran.split(ran_key, len(t_obs_targets)) + + ndbins_lo = [] + ndbins_hi = [] + for i in range(len(logmstar_bins_pdf) - 1): + for j in range(len(logssfr_bins_pdf) - 1): + ndbins_lo.append([logmstar_bins_pdf[i], logssfr_bins_pdf[j]]) + ndbins_hi.append([logmstar_bins_pdf[i + 1], logssfr_bins_pdf[j + 1]]) + ndbins_lo = np.array(ndbins_lo) + ndbins_hi = np.array(ndbins_hi) + + indx_pdf, mhalo_pdf_sat_ragged = prepare_ragged( + indx_pdf, mhalo_pdf_sat, index_mhalo + ) + + loss_data_ssfr_sat = ( + mah_params_data, + logmp0_data, + upid_data, + lgmu_infall_data, + logmhost_infall_data, + gyr_since_infall_data, + ran_key_data, + t_obs_targets, + ndbins_lo, + ndbins_hi, + logmstar_bins_pdf, + logssfr_bins_pdf, + mhalo_pdf_sat_ragged, + indx_pdf, + jnp.asarray(target_mstar_ids), + target_data_sat, + ) + + # Create loss_data for plot --------------------------------------------- + print("Creating loss data for plot...") + nhalos_plot = 10000 + + ran_key = jran.PRNGKey(np.random.randint(2**32)) + + lgmu_infall = -1.0 + logmhost_infall = 13.0 + gyr_since_infall = -99.0 # 2.0 + + mah_params_data = [] + logmp0_data = [] + upid_data = [] + lgmu_infall_data = [] + logmhost_infall_data = [] + gyr_since_infall_data = [] + t_obs_targets = [] + + mstar_ssfr_pdfs_sat = np.clip(mstar_ssfr_wcounts_sat, 0.0, None) + mstar_ssfr_pdfs_sat = ( + mstar_ssfr_pdfs_sat / np.sum(mstar_ssfr_pdfs_sat, axis=(2, 3))[:, :, None, None] + ) + mstar_ssfr_pdfs_sat = np.where( + np.isnan(mstar_ssfr_pdfs_sat), 0.0, mstar_ssfr_pdfs_sat + ) + mstar_ssfr_pdfs_sat = np.einsum("zmab,zm->zab", mstar_ssfr_pdfs_sat, mhalo_pdf_sat) + + index_mhalo = [] + indx_pdf = [] + _run_indx = 0 + + for i in range(len(age_targets)): + t_target = age_targets[i] + index_mhalo_atz = [] + indx_pdf_atz = [] + for j in range(len(logmh_binsc)): + sel = (tobs_id == i) & (logmh_id == j) & (upid_samp != -1) + + if sel.sum() < 50: + print(i, j) + continue + arange_sel = np.arange(len(tobs_id))[sel] + replace = True if sel.sum() < nhalos_plot else False + arange_sel = np.random.choice(arange_sel, nhalos_plot, replace=replace) + mah_params_data.append(mah_params_samp[:, arange_sel]) + upid_data.append(upid_samp[arange_sel]) + lgmu_infall_data.append(np.ones(len(arange_sel)) * lgmu_infall) + logmhost_infall_data.append(np.ones(len(arange_sel)) * logmhost_infall) + gyr_since_infall_data.append(np.ones(len(arange_sel)) * gyr_since_infall) + t_obs_targets.append(t_target) + mah_pars_ntuple = DiffmahParams(*mah_params_samp[:, arange_sel]) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, LGT0) + logmp0_data.append(log_mah_fit[:, -1]) + + index_mhalo_atz.append(j) + indx_pdf_atz.append(_run_indx) + _run_indx += 1 + + index_mhalo.append(np.array(index_mhalo_atz)) + indx_pdf.append(np.array(indx_pdf_atz)) + # break + # target_mstar_ids = np.array([4, 9, 14, 17, 19, 22]) + # target_mstar_ids = np.array([9, 14, 17, 19, 22]) + target_mstar_ids = np.array([8, 10, 13, 16, 19]) + # target_mstar_ids = np.array([9, 12, 15, 17]) + # target_mstar_ids = np.array([14, 17, 19, 22]) + print(logmstar_binsc_pdf[target_mstar_ids]) + target_data_sat = np.zeros( + (len(age_targets), len(target_mstar_ids), len(logssfr_binsc_pdf)) + ) + for i in range(len(age_targets)): + for j, jval in enumerate(target_mstar_ids): + if mstar_ssfr_pdfs_sat[i, jval].sum() > 0.0: + target_data_sat[i, j] = ( + mstar_ssfr_pdfs_sat[i, jval] / mstar_ssfr_pdfs_sat[i, jval].sum() + ) + + mah_params_data = np.array(mah_params_data) + logmp0_data = np.array(logmp0_data) + upid_data = np.array(upid_data) + lgmu_infall_data = np.array(lgmu_infall_data) + logmhost_infall_data = np.array(logmhost_infall_data) + gyr_since_infall_data = np.array(gyr_since_infall_data) + t_obs_targets = np.array(t_obs_targets) + + ran_key_data = jran.split(ran_key, len(t_obs_targets)) + + ndbins_lo = [] + ndbins_hi = [] + for i in range(len(logmstar_bins_pdf) - 1): + for j in range(len(logssfr_bins_pdf) - 1): + ndbins_lo.append([logmstar_bins_pdf[i], logssfr_bins_pdf[j]]) + ndbins_hi.append([logmstar_bins_pdf[i + 1], logssfr_bins_pdf[j + 1]]) + ndbins_lo = np.array(ndbins_lo) + ndbins_hi = np.array(ndbins_hi) + + indx_pdf, mhalo_pdf_sat_ragged = prepare_ragged( + indx_pdf, mhalo_pdf_sat, index_mhalo + ) + + loss_data_ssfr_sat_pred = ( + mah_params_data, + logmp0_data, + upid_data, + lgmu_infall_data, + logmhost_infall_data, + gyr_since_infall_data, + ran_key_data, + t_obs_targets, + ndbins_lo, + ndbins_hi, + logmstar_bins_pdf, + logssfr_bins_pdf, + mhalo_pdf_sat_ragged, + indx_pdf, + jnp.asarray(target_mstar_ids), + target_data_sat, + ) + + plot_data = ( + jnp.asarray(target_mstar_ids), + logssfr_binsc_pdf, + target_data_sat, + loss_data_ssfr_sat_pred, + ) + + return loss_data_ssfr_sat, plot_data + + +def get_loss_data_pdfs_mstar_cen(indir, nhalos): + # Load PDF data --------------------------------------------- + print("Loading PDF Mstar for centrals data...") + + fname = os.path.join(indir, "smdpl_smhm.h5") + with h5py.File(fname, "r") as hdf: + redshift_targets = hdf["redshift_targets"][:] + smhm_diff = hdf["smhm_diff"][:] + smhm = hdf["smhm"][:] + logmh_bins = hdf["logmh_bins"][:] + age_targets = hdf["age_targets"][:] + hist = hdf["hist"][:] + counts = hdf["counts"][:] + counts_cen = hdf["counts_cen"][:] + counts_sat = hdf["counts_sat"][:] + + logmh_binsc = 0.5 * (logmh_bins[1:] + logmh_bins[:-1]) + + fname = os.path.join(indir, "smdpl_smhm_samples_haloes.h5") + with h5py.File(fname, "r") as hdf: + logmh_id = hdf["logmh_id"][:] + logmh_val = hdf["logmh_id"][:] + mah_params_samp = hdf["mah_params_samp"][:] + ms_params_samp = hdf["ms_params_samp"][:] + q_params_samp = hdf["q_params_samp"][:] + upid_samp = hdf["upid_samp"][:] + tobs_id = hdf["tobs_id"][:] + tobs_val = hdf["tobs_val"][:] + redshift_val = hdf["redshift_val"][:] + + fname = os.path.join(indir, "smdpl_mstar_ssfr.h5") + with h5py.File(fname, "r") as hdf: + mstar_wcounts = hdf["mstar_wcounts"][:] + mstar_counts = hdf["mstar_counts"][:] + mstar_ssfr_wcounts_cent = hdf["mstar_ssfr_wcounts_cent"][:] + mstar_ssfr_wcounts_sat = hdf["mstar_ssfr_wcounts_sat"][:] + logssfr_bins_pdf = hdf["logssfr_bins_pdf"][:] + logmstar_bins_pdf = hdf["logmstar_bins_pdf"][:] + """ + hdfout["mstar_wcounts"] = mstar_wcounts + hdfout["mstar_counts"] = mstar_counts + hdfout["mstar_ssfr_wcounts_cent"] = mstar_ssfr_wcounts_cent + hdfout["mstar_ssfr_wcounts_sat"] = mstar_ssfr_wcounts_sat + hdfout["logmh_bins"] = smhm_utils.LOGMH_BINS + hdfout["logmstar_bins_pdf"] = smhm_utils.LOGMSTAR_BINS_PDF + hdfout["logssfr_bins_pdf"] = smhm_utils.LOGSSFR_BINS_PDF + hdfout["redshift_targets"] = redshift_targets + hdfout["age_targets"] = age_targets + """ + + mstar_wcounts_cen = np.sum(mstar_ssfr_wcounts_cent, axis=3) + mstar_wcounts_sat = np.sum(mstar_ssfr_wcounts_sat, axis=3) + + logssfr_binsc_pdf = 0.5 * (logssfr_bins_pdf[1:] + logssfr_bins_pdf[:-1]) + logmstar_binsc_pdf = 0.5 * (logmstar_bins_pdf[1:] + logmstar_bins_pdf[:-1]) + + # Create loss_data --------------------------------------------- + print("Creating loss data...") + + ran_key = jran.PRNGKey(np.random.randint(2**32)) + + lgmu_infall = -1.0 + logmhost_infall = 13.0 + gyr_since_infall = -99.0 # 2.0 + + # Centrals + mah_params_data = [] + logmp0_data = [] + upid_data = [] + lgmu_infall_data = [] + logmhost_infall_data = [] + gyr_since_infall_data = [] + t_obs_targets = [] + mstar_counts_target = [] + + tarr_logm0 = np.logspace(-1, LGT0, 50) + + for i in range(len(age_targets)): + t_target = age_targets[i] + + for j in range(len(logmh_binsc)): + sel = (tobs_id == i) & (logmh_id == j) & (upid_samp == -1) + + if sel.sum() < 50: + continue + replace = True if sel.sum() < nhalos else False + arange_sel = np.arange(len(tobs_id))[sel] + arange_sel = np.random.choice(arange_sel, nhalos, replace=replace) + mah_params_data.append(mah_params_samp[:, arange_sel]) + upid_data.append(upid_samp[arange_sel]) + lgmu_infall_data.append(np.ones(len(arange_sel)) * lgmu_infall) + logmhost_infall_data.append(np.ones(len(arange_sel)) * logmhost_infall) + gyr_since_infall_data.append(np.ones(len(arange_sel)) * gyr_since_infall) + t_obs_targets.append(t_target) + mstar_counts_target.append( + mstar_wcounts_cen[i, j] / mstar_wcounts_cen[i, j].sum() + ) + mah_pars_ntuple = DiffmahParams(*mah_params_samp[:, arange_sel]) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, LGT0) + logmp0_data.append(log_mah_fit[:, -1]) + # break + + mah_params_data = np.array(mah_params_data) + logmp0_data = np.array(logmp0_data) + upid_data = np.array(upid_data) + lgmu_infall_data = np.array(lgmu_infall_data) + logmhost_infall_data = np.array(logmhost_infall_data) + gyr_since_infall_data = np.array(gyr_since_infall_data) + t_obs_targets = np.array(t_obs_targets) + mstar_counts_target = np.array(mstar_counts_target) + + ran_key_data = jran.split(ran_key, len(mstar_counts_target)) + loss_data_mstar = ( + mah_params_data, + logmp0_data, + upid_data, + lgmu_infall_data, + logmhost_infall_data, + gyr_since_infall_data, + ran_key_data, + t_obs_targets, + logmstar_bins_pdf, + mstar_counts_target, + ) + + # Create loss_data for plot --------------------------------------------- + print("Creating loss data for plot...") + nhalos_plot = 10000 + ran_key = jran.PRNGKey(np.random.randint(2**32)) + + lgmu_infall = -1.0 + logmhost_infall = 13.0 + gyr_since_infall = -99.0 # 2.0 + + mah_params_data = [] + logmp0_data = [] + upid_data = [] + lgmu_infall_data = [] + logmhost_infall_data = [] + gyr_since_infall_data = [] + t_obs_targets = [] + mstar_counts_target = [] + + tarr_logm0 = np.logspace(-1, LGT0, 50) + + for i in range(len(age_targets)): + t_target = age_targets[i] + + for j in range(len(logmh_binsc)): + sel = (tobs_id == i) & (logmh_id == j) & (upid_samp == -1) + + if sel.sum() < 50: + continue + replace = True if sel.sum() < nhalos_plot else False + arange_sel = np.arange(len(tobs_id))[sel] + arange_sel = np.random.choice(arange_sel, nhalos_plot, replace=replace) + mah_params_data.append(mah_params_samp[:, arange_sel]) + upid_data.append(upid_samp[arange_sel]) + lgmu_infall_data.append(np.ones(len(arange_sel)) * lgmu_infall) + logmhost_infall_data.append(np.ones(len(arange_sel)) * logmhost_infall) + gyr_since_infall_data.append(np.ones(len(arange_sel)) * gyr_since_infall) + t_obs_targets.append(t_target) + mstar_counts_target.append( + mstar_wcounts_cen[i, j] / mstar_wcounts_cen[i, j].sum() + ) + mah_pars_ntuple = DiffmahParams(*mah_params_samp[:, arange_sel]) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, LGT0) + logmp0_data.append(log_mah_fit[:, -1]) + # break + + mah_params_data = np.array(mah_params_data) + logmp0_data = np.array(logmp0_data) + upid_data = np.array(upid_data) + lgmu_infall_data = np.array(lgmu_infall_data) + logmhost_infall_data = np.array(logmhost_infall_data) + gyr_since_infall_data = np.array(gyr_since_infall_data) + t_obs_targets = np.array(t_obs_targets) + mstar_counts_target = np.array(mstar_counts_target) + + ran_key_data = jran.split(ran_key, len(mstar_counts_target)) + loss_data_mstar_pred = ( + mah_params_data, + logmp0_data, + upid_data, + lgmu_infall_data, + logmhost_infall_data, + gyr_since_infall_data, + ran_key_data, + t_obs_targets, + logmstar_bins_pdf, + mstar_counts_target, + ) + plot_data = ( + logmstar_bins_pdf, + mstar_wcounts, + age_targets, + redshift_targets, + tobs_id, + logmh_id, + logmh_binsc, + loss_data_mstar_pred, + ) + + return loss_data_mstar, plot_data + + +def get_loss_data_pdfs_mstar_sat(indir, nhalos): + # Load PDF data --------------------------------------------- + print("Loading PDF Mstar for satellites data...") + + fname = os.path.join(indir, "smdpl_smhm.h5") + with h5py.File(fname, "r") as hdf: + redshift_targets = hdf["redshift_targets"][:] + smhm_diff = hdf["smhm_diff"][:] + smhm = hdf["smhm"][:] + logmh_bins = hdf["logmh_bins"][:] + age_targets = hdf["age_targets"][:] + hist = hdf["hist"][:] + counts = hdf["counts"][:] + counts_cen = hdf["counts_cen"][:] + counts_sat = hdf["counts_sat"][:] + + logmh_binsc = 0.5 * (logmh_bins[1:] + logmh_bins[:-1]) + + fname = os.path.join(indir, "smdpl_smhm_samples_haloes.h5") + with h5py.File(fname, "r") as hdf: + logmh_id = hdf["logmh_id"][:] + logmh_val = hdf["logmh_id"][:] + mah_params_samp = hdf["mah_params_samp"][:] + ms_params_samp = hdf["ms_params_samp"][:] + q_params_samp = hdf["q_params_samp"][:] + upid_samp = hdf["upid_samp"][:] + tobs_id = hdf["tobs_id"][:] + tobs_val = hdf["tobs_val"][:] + redshift_val = hdf["redshift_val"][:] + + fname = os.path.join(indir, "smdpl_mstar_ssfr.h5") + with h5py.File(fname, "r") as hdf: + mstar_wcounts = hdf["mstar_wcounts"][:] + mstar_counts = hdf["mstar_counts"][:] + mstar_ssfr_wcounts_cent = hdf["mstar_ssfr_wcounts_cent"][:] + mstar_ssfr_wcounts_sat = hdf["mstar_ssfr_wcounts_sat"][:] + logssfr_bins_pdf = hdf["logssfr_bins_pdf"][:] + logmstar_bins_pdf = hdf["logmstar_bins_pdf"][:] + """ + hdfout["mstar_wcounts"] = mstar_wcounts + hdfout["mstar_counts"] = mstar_counts + hdfout["mstar_ssfr_wcounts_cent"] = mstar_ssfr_wcounts_cent + hdfout["mstar_ssfr_wcounts_sat"] = mstar_ssfr_wcounts_sat + hdfout["logmh_bins"] = smhm_utils.LOGMH_BINS + hdfout["logmstar_bins_pdf"] = smhm_utils.LOGMSTAR_BINS_PDF + hdfout["logssfr_bins_pdf"] = smhm_utils.LOGSSFR_BINS_PDF + hdfout["redshift_targets"] = redshift_targets + hdfout["age_targets"] = age_targets + """ + + mstar_wcounts_cen = np.sum(mstar_ssfr_wcounts_cent, axis=3) + mstar_wcounts_sat = np.sum(mstar_ssfr_wcounts_sat, axis=3) + + logssfr_binsc_pdf = 0.5 * (logssfr_bins_pdf[1:] + logssfr_bins_pdf[:-1]) + logmstar_binsc_pdf = 0.5 * (logmstar_bins_pdf[1:] + logmstar_bins_pdf[:-1]) + + # Create loss_data --------------------------------------------- + print("Creating loss data...") + + ran_key = jran.PRNGKey(np.random.randint(2**32)) + + lgmu_infall = -1.0 + logmhost_infall = 13.0 + gyr_since_infall = -99.0 # 2.0 + + # Centrals + mah_params_data = [] + logmp0_data = [] + upid_data = [] + lgmu_infall_data = [] + logmhost_infall_data = [] + gyr_since_infall_data = [] + t_obs_targets = [] + mstar_counts_target = [] + + tarr_logm0 = np.logspace(-1, LGT0, 50) + + for i in range(len(age_targets)): + t_target = age_targets[i] + + for j in range(len(logmh_binsc)): + sel = (tobs_id == i) & (logmh_id == j) & (upid_samp != -1) + + if sel.sum() < 50: + continue + replace = True if sel.sum() < nhalos else False + arange_sel = np.arange(len(tobs_id))[sel] + arange_sel = np.random.choice(arange_sel, nhalos, replace=replace) + mah_params_data.append(mah_params_samp[:, arange_sel]) + upid_data.append(upid_samp[arange_sel]) + lgmu_infall_data.append(np.ones(len(arange_sel)) * lgmu_infall) + logmhost_infall_data.append(np.ones(len(arange_sel)) * logmhost_infall) + gyr_since_infall_data.append(np.ones(len(arange_sel)) * gyr_since_infall) + t_obs_targets.append(t_target) + mstar_counts_target.append( + mstar_wcounts_sat[i, j] / mstar_wcounts_sat[i, j].sum() + ) + mah_pars_ntuple = DiffmahParams(*mah_params_samp[:, arange_sel]) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, LGT0) + logmp0_data.append(log_mah_fit[:, -1]) + # break + + mah_params_data = np.array(mah_params_data) + logmp0_data = np.array(logmp0_data) + upid_data = np.array(upid_data) + lgmu_infall_data = np.array(lgmu_infall_data) + logmhost_infall_data = np.array(logmhost_infall_data) + gyr_since_infall_data = np.array(gyr_since_infall_data) + t_obs_targets = np.array(t_obs_targets) + mstar_counts_target = np.array(mstar_counts_target) + + ran_key_data = jran.split(ran_key, len(mstar_counts_target)) + loss_data_mstar = ( + mah_params_data, + logmp0_data, + upid_data, + lgmu_infall_data, + logmhost_infall_data, + gyr_since_infall_data, + ran_key_data, + t_obs_targets, + logmstar_bins_pdf, + mstar_counts_target, + ) + + # Create loss_data for plot --------------------------------------------- + print("Creating loss data for plot...") + nhalos_plot = 10000 + ran_key = jran.PRNGKey(np.random.randint(2**32)) + + lgmu_infall = -1.0 + logmhost_infall = 13.0 + gyr_since_infall = -99.0 # 2.0 + + mah_params_data = [] + logmp0_data = [] + upid_data = [] + lgmu_infall_data = [] + logmhost_infall_data = [] + gyr_since_infall_data = [] + t_obs_targets = [] + mstar_counts_target = [] + + tarr_logm0 = np.logspace(-1, LGT0, 50) + + for i in range(len(age_targets)): + t_target = age_targets[i] + + for j in range(len(logmh_binsc)): + sel = (tobs_id == i) & (logmh_id == j) & (upid_samp != -1) + + if sel.sum() < 50: + continue + replace = True if sel.sum() < nhalos_plot else False + arange_sel = np.arange(len(tobs_id))[sel] + arange_sel = np.random.choice(arange_sel, nhalos_plot, replace=replace) + mah_params_data.append(mah_params_samp[:, arange_sel]) + upid_data.append(upid_samp[arange_sel]) + lgmu_infall_data.append(np.ones(len(arange_sel)) * lgmu_infall) + logmhost_infall_data.append(np.ones(len(arange_sel)) * logmhost_infall) + gyr_since_infall_data.append(np.ones(len(arange_sel)) * gyr_since_infall) + t_obs_targets.append(t_target) + mstar_counts_target.append( + mstar_wcounts_sat[i, j] / mstar_wcounts_sat[i, j].sum() + ) + mah_pars_ntuple = DiffmahParams(*mah_params_samp[:, arange_sel]) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, LGT0) + logmp0_data.append(log_mah_fit[:, -1]) + # break + + mah_params_data = np.array(mah_params_data) + logmp0_data = np.array(logmp0_data) + upid_data = np.array(upid_data) + lgmu_infall_data = np.array(lgmu_infall_data) + logmhost_infall_data = np.array(logmhost_infall_data) + gyr_since_infall_data = np.array(gyr_since_infall_data) + t_obs_targets = np.array(t_obs_targets) + mstar_counts_target = np.array(mstar_counts_target) + + ran_key_data = jran.split(ran_key, len(mstar_counts_target)) + loss_data_mstar_pred = ( + mah_params_data, + logmp0_data, + upid_data, + lgmu_infall_data, + logmhost_infall_data, + gyr_since_infall_data, + ran_key_data, + t_obs_targets, + logmstar_bins_pdf, + mstar_counts_target, + ) + plot_data = ( + logmstar_bins_pdf, + mstar_wcounts, + age_targets, + redshift_targets, + tobs_id, + logmh_id, + logmh_binsc, + loss_data_mstar_pred, + ) + + return loss_data_mstar, plot_data diff --git a/scripts/diffstarpop_scripts/fit_mstar_ssfr_pdfs_mgash.py b/scripts/diffstarpop_scripts/fit_mstar_ssfr_pdfs_mgash.py new file mode 100644 index 0000000..1cb4d43 --- /dev/null +++ b/scripts/diffstarpop_scripts/fit_mstar_ssfr_pdfs_mgash.py @@ -0,0 +1,228 @@ +import os +import h5py +import numpy as np +from jax import ( + numpy as jnp, + jit as jjit, + random as jran, + grad, + vmap, +) +import argparse +from time import time +from scipy.optimize import minimize +from jax.example_libraries import optimizers as jax_opt + +from collections import OrderedDict, namedtuple + +from diffstar.defaults import TODAY, LGT0 +from diffmah.diffmah_kernels import mah_halopop + +from diffstar.diffstarpop.loss_kernels.mstar_ssfr_loss_mgash import ( + loss_mstar_kern_tobs_grad_wrapper, + loss_mstar_ssfr_kern_tobs_grad_wrapper, + loss_combined_wrapper, + loss_combined_3loss_wrapper, +) + +from diffstar.diffstarpop.kernels.defaults_mgash import ( + DEFAULT_DIFFSTARPOP_U_PARAMS, + DEFAULT_DIFFSTARPOP_PARAMS, + get_bounded_diffstarpop_params, +) + +from fit_get_loss_helpers_mgash import ( + get_loss_data_pdfs_mstar, + get_loss_data_pdfs_ssfr_central, + get_loss_data_pdfs_ssfr_satellite, +) +from diffstar.diffstarpop.kernels.params import ( + DiffstarPop_UParams_Diffstarfits_mgash, +) + +BEBOP_SMHM_MEAN_DATA = "/lcrc/project/halotools/alarcon/results/" + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument( + "-indir", help="input drn", type=str, default=BEBOP_SMHM_MEAN_DATA + ) + parser.add_argument( + "-outdir", help="output drn", type=str, default=BEBOP_SMHM_MEAN_DATA + ) + parser.add_argument( + "-nhalos", help="Number of halos for fitting", type=int, default=100 + ) + parser.add_argument( + "-nstep", help="Number of steps for fitting", type=int, default=1000 + ) + parser.add_argument( + "-outname", + help="output fname for best params", + type=str, + default="bestfit_diffstarpop_params", + ) + parser.add_argument( + "-loss_type", + help="Which data to target", + type=str, + choices=["mstar", "mstar_ssfr_cen", "mstar_ssfr_cen_sat"], + default="mstar", + ) + parser.add_argument( + "--params_path", + type=str, + default=None, + help="Path were diffstarpop params are stored", + ) + parser.add_argument( + "--print_loss", + type=int, + default=100, + help="How many steps before printing current loss", + ) + + args = parser.parse_args() + indir = args.indir + outdir = args.outdir + nhalos = args.nhalos + n_step = args.nstep + outname = args.outname + params_path = args.params_path + loss_type = args.loss_type + + # Load MStar pdf data --------------------------------------------- + + if loss_type == "mstar": + loss_data_mstar, plot_data_mstar = get_loss_data_pdfs_mstar(indir, nhalos) + loss_data = (loss_data_mstar,) + elif loss_type == "mstar_ssfr_cen": + loss_data_mstar, plot_data_mstar = get_loss_data_pdfs_mstar(indir, nhalos) + loss_data_ssfr_cen, plot_data_ssfr_cen = get_loss_data_pdfs_ssfr_central( + indir, nhalos + ) + loss_data = (loss_data_mstar, loss_data_ssfr_cen) + elif loss_type == "mstar_ssfr_cen_sat": + loss_data_mstar, plot_data_mstar = get_loss_data_pdfs_mstar(indir, nhalos) + loss_data_ssfr_cen, plot_data_ssfr_cen = get_loss_data_pdfs_ssfr_central( + indir, nhalos + ) + loss_data_ssfr_sat, plot_data_ssfr_sat = get_loss_data_pdfs_ssfr_satellite( + indir, nhalos + ) + loss_data = (loss_data_mstar, loss_data_ssfr_cen, loss_data_ssfr_sat) + + # Define loss kernel --------------------------------------------- + if loss_type == "mstar": + loss_kernel = loss_mstar_kern_tobs_grad_wrapper + elif loss_type == "mstar_ssfr_cen": + loss_kernel = loss_combined_wrapper + elif loss_type == "mstar_ssfr_cen_sat": + loss_kernel = loss_combined_3loss_wrapper + + # Register params --------------------------------------------- + + if params_path is None: + all_u_params = jnp.asarray(DEFAULT_DIFFSTARPOP_U_PARAMS) + elif params_path.startswith("diffstarfits"): + sim_name = params_path.split("_")[1:] + sim_name = ("_").join(sim_name) + params_tuple = DiffstarPop_UParams_Diffstarfits_mgash[sim_name] + all_u_params = jnp.asarray(params_tuple) + else: + params = np.load(params_path) + all_u_params = params["diffstarpop_u_params"] + + # Run fitter --------------------------------------------- + print("Running fitter...") + + params_init = jnp.asarray(all_u_params) + loss_kernel(params_init, *loss_data) + + start = time() + + step_size = 0.01 + + loss_arr = np.zeros(n_step).astype("f4") + np.inf + + opt_init, opt_update, get_params = jax_opt.adam(step_size) + opt_state = opt_init(params_init) + + n_params = len(params_init) + params_arr = np.zeros((n_step, n_params)).astype("f4") + + n_mah = 100 + + ran_key = jran.PRNGKey(np.random.randint(2**32)) + + no_nan_grads_arr = np.zeros(n_step) + for istep in range(n_step): + start = time() + ran_key, subkey = jran.split(ran_key, 2) + + p = np.array(get_params(opt_state)) + + loss, grads = loss_kernel(p, *loss_data) + + no_nan_params = np.all(np.isfinite(p)) + no_nan_loss = np.isfinite(loss) + no_nan_grads = np.all(np.isfinite(grads)) + if ~no_nan_loss: + print("NaN in loss, trying to take extra gradient step") + opt_state2 = opt_update(istep, current_grads, opt_state) + p2 = np.array(get_params(opt_state2)) + loss2, grads2 = loss_kernel(p2, *loss_data) + no_nan_params2 = np.all(np.isfinite(p2)) + no_nan_loss2 = np.isfinite(loss2) + no_nan_grads2 = np.all(np.isfinite(grads2)) + if ~no_nan_params2 | ~no_nan_loss2 | ~no_nan_grads2: + print("Extra step failed") + continue + else: + p = p2.copy() + loss = loss2.copy() + grads = grads2.copy() + no_nan_params = np.all(np.isfinite(p)) + no_nan_loss = np.isfinite(loss) + no_nan_grads = np.all(np.isfinite(grads)) + + if ~no_nan_params | ~no_nan_loss | ~no_nan_grads: + # break + if istep > 0: + indx_best = np.nanargmin(loss_arr[:istep]) + best_fit_params = params_arr[indx_best] + best_fit_loss = loss_arr[indx_best] + else: + best_fit_params = np.copy(p) + best_fit_loss = 999.99 + else: + params_arr[istep, :] = p + loss_arr[istep] = loss + opt_state = opt_update(istep, grads, opt_state) + + current_grads = grads.copy() + + no_nan_grads_arr[istep] = ~no_nan_grads + end = time() + if istep % args.print_loss == 0: + print(istep, loss, end - start, no_nan_grads) + if ~no_nan_grads: + break + + argmin_best = np.argmin(loss_arr) + best_fit_u_params = params_arr[argmin_best] + + def return_params_from_result(best_fit_u_params): + bestfit_u_tuple = DEFAULT_DIFFSTARPOP_U_PARAMS._make(best_fit_u_params) + diffstarpop_params = get_bounded_diffstarpop_params(bestfit_u_tuple) + return diffstarpop_params + + best_fit_params = return_params_from_result(best_fit_u_params) + best_fit_params = jnp.asarray(best_fit_params) + + np.savez( + os.path.join(outdir, outname) + ".npz", + diffstarpop_params=best_fit_params, + diffstarpop_u_params=best_fit_u_params, + ) diff --git a/scripts/diffstarpop_scripts/measure_smhm_galacticus_script_mpi_mgash.py b/scripts/diffstarpop_scripts/measure_smhm_galacticus_script_mpi_mgash.py new file mode 100644 index 0000000..e359d0c --- /dev/null +++ b/scripts/diffstarpop_scripts/measure_smhm_galacticus_script_mpi_mgash.py @@ -0,0 +1,232 @@ +"""This module tabulates for SMDPL""" + +import argparse +import os +from time import time +import subprocess + +import h5py +import numpy as np +import gc + +import smhm_utils_galacticus_mgash as smhm_utils + +from mpi4py import MPI + +TMP_OUTPATH = "_tmp_subvol_{0}_galacticus_smhm.h5" + +if __name__ == "__main__": + comm = MPI.COMM_WORLD + rank, nranks = comm.Get_rank(), comm.Get_size() + + parser = argparse.ArgumentParser() + parser.add_argument( + "-diffmah_drn", help="input drn", type=str, default=smhm_utils.BEBOP_GALAC + ) + parser.add_argument( + "-diffstar_drn", + help="input drn", + type=str, + default=smhm_utils.BEBOP_GALAC, + ) + + parser.add_argument("-outdrn", help="output directory", type=str) + parser.add_argument( + "-sfh_type", + help="Type of star formation histories", + choices=["in_situ", "in_plus_ex_situ"], + default="in_situ", + ) + parser.add_argument( + "-z_bins", + nargs="+", + type=float, + help="List of redshift bins", + default=smhm_utils.Z_BINS, + ) + + args = parser.parse_args() + + outdrn = args.outdrn + diffmah_drn = args.diffmah_drn + diffstar_drn = args.diffstar_drn + sfh_type = args.sfh_type + + redshift_targets = args.z_bins + nz, nm = len(redshift_targets), smhm_utils.LOGMH_BINS.size - 1 + nmstar = len(smhm_utils.LOGMSTAR_BINS_PDF) - 1 + nssfr = len(smhm_utils.LOGSSFR_BINS_PDF) - 1 + + haloes_data = [] + print("Beginning loop over subvolumes...\n") + + start = time() + _res = smhm_utils.create_target_data( + sfh_type, redshift_targets, diffmah_drn=diffmah_drn, diffstar_drn=diffstar_drn + ) + ( + wcounts_i, + whist_i, + counts_i, + hist_i, + age_targets, + haloes, + counts_cen_i, + counts_sat_i, + ) = _res + + ( + logmh_id, + logmh_val, + mah_params_samp, + ms_params_samp, + q_params_samp, + upid_samp, + tobs_id, + tobs_val, + redshift_val, + ) = haloes + + _res = smhm_utils.create_pdf_target_data( + sfh_type, redshift_targets, diffmah_drn=diffmah_drn, diffstar_drn=diffstar_drn + ) + + fnout = os.path.join(outdrn, TMP_OUTPATH.format(rank)) + with h5py.File(fnout, "w") as hdfout: + hdfout["wcounts_i"] = wcounts_i + hdfout["whist_i"] = whist_i + hdfout["counts_i"] = counts_i + hdfout["hist_i"] = hist_i + hdfout["age_targets"] = age_targets + hdfout["counts_cen_i"] = counts_cen_i + hdfout["counts_sat_i"] = counts_sat_i + hdfout["mstar_wcounts_i"] = _res[0] + hdfout["mstar_counts_i"] = _res[1] + hdfout["mstar_ssfr_wcounts_cent_i"] = _res[2] + hdfout["mstar_ssfr_wcounts_sat_i"] = _res[3] + + hdfout["logmh_id"] = logmh_id + hdfout["logmh_val"] = logmh_val + hdfout["mah_params_samp"] = mah_params_samp + hdfout["ms_params_samp"] = ms_params_samp + hdfout["q_params_samp"] = q_params_samp + hdfout["upid_samp"] = upid_samp + hdfout["tobs_id"] = tobs_id + hdfout["tobs_val"] = tobs_val + hdfout["redshift_val"] = redshift_val + + end = time() + runtime = end - start + print( + f"...computed sumstat counts for subvolume {rank}", + "Time: %.2f seconds." % (end - start), + ) + + comm.Barrier() + + if rank == 0: + print("Collecting all data in rank 0.") + start = time() + + wcounts = np.zeros((nz, nm)) + whist = np.zeros_like(wcounts) + counts = np.zeros_like(wcounts) + hist = np.zeros_like(wcounts) + counts_cen = np.zeros_like(wcounts) + counts_sat = np.zeros_like(wcounts) + + mstar_wcounts = np.zeros((nz, nm, nmstar)) + mstar_counts = np.zeros((nz, nm, nmstar)) + + mstar_ssfr_wcounts_cent = np.zeros((nz, nm, nmstar, nssfr)) + mstar_ssfr_wcounts_sat = np.zeros((nz, nm, nmstar, nssfr)) + + for i in range(1): + fnout = os.path.join(outdrn, TMP_OUTPATH.format(i)) + with h5py.File(fnout, "r") as hdfout: + wcounts = wcounts + hdfout["wcounts_i"][:] + whist = whist + hdfout["whist_i"][:] + counts = counts + hdfout["counts_i"][:] + hist = hist + hdfout["hist_i"][:] + counts_cen = counts_cen + hdfout["counts_cen_i"][:] + counts_sat = counts_sat + hdfout["counts_sat_i"][:] + mstar_wcounts += hdfout["mstar_wcounts_i"][:] + mstar_counts += hdfout["mstar_counts_i"][:] + mstar_ssfr_wcounts_cent += hdfout["mstar_ssfr_wcounts_cent_i"][:] + mstar_ssfr_wcounts_sat += hdfout["mstar_ssfr_wcounts_sat_i"][:] + + haloes_data.append( + ( + hdfout["logmh_id"][:], + hdfout["logmh_val"][:], + hdfout["mah_params_samp"][:], + hdfout["ms_params_samp"][:], + hdfout["q_params_samp"][:], + hdfout["upid_samp"][:], + hdfout["tobs_id"][:], + hdfout["tobs_val"][:], + hdfout["redshift_val"][:], + ) + ) + + sampled_haloes = smhm_utils.concatenate_samples_haloes(haloes_data) + end = time() + runtime = end - start + print("Ended collecting data. Time: %.2f seconds." % (end - start)) + + print("Saving final target data...") + + fnout = os.path.join(outdrn, "smdpl_smhm.h5") + with h5py.File(fnout, "w") as hdfout: + hdfout["counts_diff"] = wcounts + hdfout["hist_diff"] = whist + hdfout["counts"] = counts + hdfout["hist"] = hist + hdfout["counts_cen"] = counts_cen + hdfout["counts_sat"] = counts_sat + hdfout["smhm_diff"] = whist / wcounts + hdfout["smhm"] = hist / counts + hdfout["logmh_bins"] = smhm_utils.LOGMH_BINS + hdfout["redshift_targets"] = redshift_targets + hdfout["age_targets"] = age_targets + + ( + logmh_id, + logmh_val, + mah_params_samp, + ms_params_samp, + q_params_samp, + upid_samp, + tobs_id, + tobs_val, + redshift_val, + ) = sampled_haloes + + fnout = os.path.join(outdrn, "smdpl_smhm_samples_haloes.h5") + with h5py.File(fnout, "w") as hdfout: + hdfout["logmh_id"] = logmh_id + hdfout["logmh_val"] = logmh_val + hdfout["mah_params_samp"] = mah_params_samp + hdfout["ms_params_samp"] = ms_params_samp + hdfout["q_params_samp"] = q_params_samp + hdfout["upid_samp"] = upid_samp + hdfout["tobs_id"] = tobs_id + hdfout["tobs_val"] = tobs_val + hdfout["redshift_val"] = redshift_val + + fnout = os.path.join(outdrn, "smdpl_mstar_ssfr.h5") + with h5py.File(fnout, "w") as hdfout: + hdfout["mstar_wcounts"] = mstar_wcounts + hdfout["mstar_counts"] = mstar_counts + hdfout["mstar_ssfr_wcounts_cent"] = mstar_ssfr_wcounts_cent + hdfout["mstar_ssfr_wcounts_sat"] = mstar_ssfr_wcounts_sat + hdfout["logmh_bins"] = smhm_utils.LOGMH_BINS + hdfout["logmstar_bins_pdf"] = smhm_utils.LOGMSTAR_BINS_PDF + hdfout["logssfr_bins_pdf"] = smhm_utils.LOGSSFR_BINS_PDF + hdfout["redshift_targets"] = redshift_targets + hdfout["age_targets"] = age_targets + + # clean up all temporary data + bnpat = os.path.join(outdrn, "_tmp_subvol_*") + command = "rm " + bnpat + subprocess.os.system(command) diff --git a/scripts/diffstarpop_scripts/measure_smhm_smdpl_script_mpi_mgash.py b/scripts/diffstarpop_scripts/measure_smhm_smdpl_script_mpi_mgash.py new file mode 100644 index 0000000..34a6247 --- /dev/null +++ b/scripts/diffstarpop_scripts/measure_smhm_smdpl_script_mpi_mgash.py @@ -0,0 +1,291 @@ +"""This module tabulates for SMDPL""" + +import argparse +import os +from time import time +import subprocess + +import h5py +import numpy as np +import gc +import re + +import smhm_utils_smdpl_mgash as smhm_utils + +from mpi4py import MPI + + +if __name__ == "__main__": + comm = MPI.COMM_WORLD + rank, nranks = comm.Get_rank(), comm.Get_size() + + parser = argparse.ArgumentParser() + parser.add_argument( + "-n_subvol_max", + help="Last subvolume", + type=int, + default=smhm_utils.N_SUBVOL_SMDPL, + ) + parser.add_argument( + "-diffmah_drn", + help="input drn", + type=str, + default=smhm_utils.LCRC_NOMERGING_DIFFMAH_DRN, + ) + parser.add_argument( + "-diffstar_drn", + help="input drn", + type=str, + default=smhm_utils.LCRC_NOMERGING_DIFFSTAR_DRN, + ) + parser.add_argument( + "-sim_name", + help="Simulation name", + choices=["DR1", "DR1_nomerging", "other"], + default="DR1_nomerging", + ) + parser.add_argument( + "-z_bins", + nargs="+", + type=float, + help="List of redshift bins", + default=smhm_utils.Z_BINS, + ) + + parser.add_argument("-outdrn", help="output directory", type=str, default="") + args = parser.parse_args() + n_subvol_max = args.n_subvol_max + outdrn = args.outdrn + sim_name = args.sim_name + + if sim_name == "DR1_nomerging": + diffmah_drn = smhm_utils.LCRC_NOMERGING_DIFFMAH_DRN + diffstar_drn = smhm_utils.LCRC_NOMERGING_DIFFSTAR_DRN + binaries_drn = smhm_utils.LCRC_NOMERGING_BINARIES_DRN + diffstar_bnpat = smhm_utils.LCRC_NOMERGING_diffstar_bnpat + elif sim_name == "DR1": + diffmah_drn = smhm_utils.LCRC_DR1_DIFFMAH_DRN + diffstar_drn = smhm_utils.LCRC_DR1_DIFFSTAR_DRN + binaries_drn = smhm_utils.LCRC_DR1_BINARIES_DRN + diffstar_bnpat = smhm_utils.LCRC_DR1_diffstar_bnpat + else: + diffmah_drn = args.diffmah_drn + diffstar_drn = args.diffstar_drn + binaries_drn = smhm_utils.LCRC_NOMERGING_BINARIES_DRN + diffstar_bnpat = smhm_utils.LCRC_NOMERGING_diffstar_bnpat + + # redshift_targets = np.concatenate((np.arange(0,1,0.1), np.arange(1, 2.1, 0.5))) + redshift_targets = args.z_bins + nz, nm = len(redshift_targets), smhm_utils.LOGMH_BINS.size - 1 + nmstar = len(smhm_utils.LOGMSTAR_BINS_PDF) - 1 + nssfr = len(smhm_utils.LOGSSFR_BINS_PDF) - 1 + + # see which subvolumes are available + # Replace the '{}' with regex to match 1 to 3 digits + # Match filenames that have the 'pattern' + regex_str = re.escape(diffstar_bnpat).replace(r"\{\}", r"(\d{1,3})") + pattern = re.compile(f"^{regex_str}$") + matching_files = [f for f in os.listdir(diffstar_drn) if pattern.match(f)] + subvol_avail = len(matching_files) + subvols = [x.split("_")[-1].split(".")[0] for x in matching_files] + subvols = np.sort(np.array(subvols).astype(int)) + subvols_arr = np.array_split(subvols, nranks)[rank] + n_subvol_smdpl = len(subvols) + subvol_used = np.zeros(n_subvol_max).astype(int) + + haloes_data = [] + print("Beginning loop over subvolumes...\n") + + for i in subvols_arr: + gc.collect() + try: + start = time() + _res = smhm_utils.create_target_data( + i, + sim_name, + n_subvol_smdpl, + redshift_targets=redshift_targets, + binaries_drn=binaries_drn, + diffmah_drn=diffmah_drn, + diffstar_drn=diffstar_drn, + diffstar_bnpat=diffstar_bnpat, + ) + ( + wcounts_i, + whist_i, + counts_i, + hist_i, + age_targets, + haloes, + counts_cen_i, + counts_sat_i, + ) = _res + + ( + logmh_id, + logmh_val, + mah_params_samp, + ms_params_samp, + q_params_samp, + upid_samp, + tobs_id, + tobs_val, + redshift_val, + ) = haloes + + _res = smhm_utils.create_pdf_target_data( + i, + sim_name, + redshift_targets, + binaries_drn=binaries_drn, + diffmah_drn=diffmah_drn, + diffstar_drn=diffstar_drn, + diffstar_bnpat=diffstar_bnpat, + ) + + fnout = os.path.join(outdrn, "_tmp_subvol_%d_smdpl_smhm.h5" % i) + with h5py.File(fnout, "w") as hdfout: + hdfout["wcounts_i"] = wcounts_i + hdfout["whist_i"] = whist_i + hdfout["counts_i"] = counts_i + hdfout["hist_i"] = hist_i + hdfout["age_targets"] = age_targets + hdfout["counts_cen_i"] = counts_cen_i + hdfout["counts_sat_i"] = counts_sat_i + hdfout["mstar_wcounts_i"] = _res[0] + hdfout["mstar_counts_i"] = _res[1] + hdfout["mstar_ssfr_wcounts_cent_i"] = _res[2] + hdfout["mstar_ssfr_wcounts_sat_i"] = _res[3] + + hdfout["logmh_id"] = logmh_id + hdfout["logmh_val"] = logmh_val + hdfout["mah_params_samp"] = mah_params_samp + hdfout["ms_params_samp"] = ms_params_samp + hdfout["q_params_samp"] = q_params_samp + hdfout["upid_samp"] = upid_samp + hdfout["tobs_id"] = tobs_id + hdfout["tobs_val"] = tobs_val + hdfout["redshift_val"] = redshift_val + + end = time() + runtime = end - start + print( + f"...computed sumstat counts for subvolume {i}", + "Time: %.2f seconds." % (end - start), + ) + except FileNotFoundError: + print(f"...NO sumstat counts for subvolume {i}") + pass + + comm.Barrier() + + if rank == 0: + print("Collecting all data in rank 0.") + start = time() + + wcounts = np.zeros((nz, nm)) + whist = np.zeros_like(wcounts) + counts = np.zeros_like(wcounts) + hist = np.zeros_like(wcounts) + counts_cen = np.zeros_like(wcounts) + counts_sat = np.zeros_like(wcounts) + + mstar_wcounts = np.zeros((nz, nm, nmstar)) + mstar_counts = np.zeros((nz, nm, nmstar)) + + mstar_ssfr_wcounts_cent = np.zeros((nz, nm, nmstar, nssfr)) + mstar_ssfr_wcounts_sat = np.zeros((nz, nm, nmstar, nssfr)) + + for i in subvols: + fnout = os.path.join(outdrn, "_tmp_subvol_%d_smdpl_smhm.h5" % i) + with h5py.File(fnout, "r") as hdfout: + wcounts = wcounts + hdfout["wcounts_i"][:] + whist = whist + hdfout["whist_i"][:] + counts = counts + hdfout["counts_i"][:] + hist = hist + hdfout["hist_i"][:] + counts_cen = counts_cen + hdfout["counts_cen_i"][:] + counts_sat = counts_sat + hdfout["counts_sat_i"][:] + subvol_used[i] = 1 + mstar_wcounts += hdfout["mstar_wcounts_i"][:] + mstar_counts += hdfout["mstar_counts_i"][:] + mstar_ssfr_wcounts_cent += hdfout["mstar_ssfr_wcounts_cent_i"][:] + mstar_ssfr_wcounts_sat += hdfout["mstar_ssfr_wcounts_sat_i"][:] + + haloes_data.append( + ( + hdfout["logmh_id"][:], + hdfout["logmh_val"][:], + hdfout["mah_params_samp"][:], + hdfout["ms_params_samp"][:], + hdfout["q_params_samp"][:], + hdfout["upid_samp"][:], + hdfout["tobs_id"][:], + hdfout["tobs_val"][:], + hdfout["redshift_val"][:], + ) + ) + + sampled_haloes = smhm_utils.concatenate_samples_haloes(haloes_data) + end = time() + runtime = end - start + print("Ended collecting data. Time: %.2f seconds." % (end - start)) + + print("Saving final target data...") + + fnout = os.path.join(outdrn, "smdpl_smhm.h5") + with h5py.File(fnout, "w") as hdfout: + hdfout["counts_diff"] = wcounts + hdfout["hist_diff"] = whist + hdfout["counts"] = counts + hdfout["hist"] = hist + hdfout["counts_cen"] = counts_cen + hdfout["counts_sat"] = counts_sat + hdfout["smhm_diff"] = whist / wcounts + hdfout["smhm"] = hist / counts + hdfout["logmh_bins"] = smhm_utils.LOGMH_BINS + hdfout["subvol_used"] = subvol_used + hdfout["redshift_targets"] = redshift_targets + hdfout["age_targets"] = age_targets + + ( + logmh_id, + logmh_val, + mah_params_samp, + ms_params_samp, + q_params_samp, + upid_samp, + tobs_id, + tobs_val, + redshift_val, + ) = sampled_haloes + + fnout = os.path.join(outdrn, "smdpl_smhm_samples_haloes.h5") + with h5py.File(fnout, "w") as hdfout: + hdfout["logmh_id"] = logmh_id + hdfout["logmh_val"] = logmh_val + hdfout["mah_params_samp"] = mah_params_samp + hdfout["ms_params_samp"] = ms_params_samp + hdfout["q_params_samp"] = q_params_samp + hdfout["upid_samp"] = upid_samp + hdfout["tobs_id"] = tobs_id + hdfout["tobs_val"] = tobs_val + hdfout["redshift_val"] = redshift_val + + fnout = os.path.join(outdrn, "smdpl_mstar_ssfr.h5") + with h5py.File(fnout, "w") as hdfout: + hdfout["mstar_wcounts"] = mstar_wcounts + hdfout["mstar_counts"] = mstar_counts + hdfout["mstar_ssfr_wcounts_cent"] = mstar_ssfr_wcounts_cent + hdfout["mstar_ssfr_wcounts_sat"] = mstar_ssfr_wcounts_sat + hdfout["logmh_bins"] = smhm_utils.LOGMH_BINS + hdfout["logmstar_bins_pdf"] = smhm_utils.LOGMSTAR_BINS_PDF + hdfout["logssfr_bins_pdf"] = smhm_utils.LOGSSFR_BINS_PDF + hdfout["redshift_targets"] = redshift_targets + hdfout["age_targets"] = age_targets + + n_used = subvol_used.sum() + + # clean up all temporary data + bnpat = os.path.join(outdrn, "_tmp_subvol_*") + command = "rm " + bnpat + subprocess.os.system(command) diff --git a/scripts/diffstarpop_scripts/measure_smhm_tng_script_mpi_mgash.py b/scripts/diffstarpop_scripts/measure_smhm_tng_script_mpi_mgash.py new file mode 100644 index 0000000..97ab9da --- /dev/null +++ b/scripts/diffstarpop_scripts/measure_smhm_tng_script_mpi_mgash.py @@ -0,0 +1,229 @@ +"""This module tabulates for SMDPL""" + +import argparse +import os +from time import time +import subprocess + +import h5py +import numpy as np +import gc + +import smhm_utils_tng_mgash as smhm_utils + +from mpi4py import MPI + +TMP_OUTPATH = "_tmp_subvol_{0}_tng_smhm.h5" + +if __name__ == "__main__": + comm = MPI.COMM_WORLD + rank, nranks = comm.Get_rank(), comm.Get_size() + + parser = argparse.ArgumentParser() + parser.add_argument( + "-diffmah_drn", help="input drn", type=str, default=smhm_utils.BEBOP_TNG_MAH + ) + parser.add_argument( + "-diffstar_drn", + help="input drn", + type=str, + default=smhm_utils.BEBOP_TNG_SFH, + ) + + parser.add_argument("-outdrn", help="output directory", type=str) + parser.add_argument( + "-nchunks", help="Number of chunks", type=int, default=smhm_utils.NCHUNKS + ) + parser.add_argument( + "-z_bins", + nargs="+", + type=float, + help="List of redshift bins", + default=smhm_utils.Z_BINS, + ) + + args = parser.parse_args() + + outdrn = args.outdrn + diffmah_drn = args.diffmah_drn + diffstar_drn = args.diffstar_drn + nchunks = args.nchunks + + redshift_targets = args.z_bins + nz, nm = len(redshift_targets), smhm_utils.LOGMH_BINS.size - 1 + nmstar = len(smhm_utils.LOGMSTAR_BINS_PDF) - 1 + nssfr = len(smhm_utils.LOGSSFR_BINS_PDF) - 1 + + haloes_data = [] + print("Beginning loop over subvolumes...\n") + + start = time() + _res = smhm_utils.create_target_data( + rank, redshift_targets, diffmah_drn=diffmah_drn, diffstar_drn=diffstar_drn + ) + ( + wcounts_i, + whist_i, + counts_i, + hist_i, + age_targets, + haloes, + counts_cen_i, + counts_sat_i, + ) = _res + + ( + logmh_id, + logmh_val, + mah_params_samp, + ms_params_samp, + q_params_samp, + upid_samp, + tobs_id, + tobs_val, + redshift_val, + ) = haloes + + _res = smhm_utils.create_pdf_target_data( + rank, redshift_targets, diffmah_drn=diffmah_drn, diffstar_drn=diffstar_drn + ) + + fnout = os.path.join(outdrn, TMP_OUTPATH.format(rank)) + with h5py.File(fnout, "w") as hdfout: + hdfout["wcounts_i"] = wcounts_i + hdfout["whist_i"] = whist_i + hdfout["counts_i"] = counts_i + hdfout["hist_i"] = hist_i + hdfout["age_targets"] = age_targets + hdfout["counts_cen_i"] = counts_cen_i + hdfout["counts_sat_i"] = counts_sat_i + hdfout["mstar_wcounts_i"] = _res[0] + hdfout["mstar_counts_i"] = _res[1] + hdfout["mstar_ssfr_wcounts_cent_i"] = _res[2] + hdfout["mstar_ssfr_wcounts_sat_i"] = _res[3] + + hdfout["logmh_id"] = logmh_id + hdfout["logmh_val"] = logmh_val + hdfout["mah_params_samp"] = mah_params_samp + hdfout["ms_params_samp"] = ms_params_samp + hdfout["q_params_samp"] = q_params_samp + hdfout["upid_samp"] = upid_samp + hdfout["tobs_id"] = tobs_id + hdfout["tobs_val"] = tobs_val + hdfout["redshift_val"] = redshift_val + + end = time() + runtime = end - start + print( + f"...computed sumstat counts for subvolume {rank}", + "Time: %.2f seconds." % (end - start), + ) + + comm.Barrier() + + if rank == 0: + print("Collecting all data in rank 0.") + start = time() + + wcounts = np.zeros((nz, nm)) + whist = np.zeros_like(wcounts) + counts = np.zeros_like(wcounts) + hist = np.zeros_like(wcounts) + counts_cen = np.zeros_like(wcounts) + counts_sat = np.zeros_like(wcounts) + + mstar_wcounts = np.zeros((nz, nm, nmstar)) + mstar_counts = np.zeros((nz, nm, nmstar)) + + mstar_ssfr_wcounts_cent = np.zeros((nz, nm, nmstar, nssfr)) + mstar_ssfr_wcounts_sat = np.zeros((nz, nm, nmstar, nssfr)) + + for i in range(smhm_utils.NCHUNKS): + fnout = os.path.join(outdrn, TMP_OUTPATH.format(i)) + with h5py.File(fnout, "r") as hdfout: + wcounts = wcounts + hdfout["wcounts_i"][:] + whist = whist + hdfout["whist_i"][:] + counts = counts + hdfout["counts_i"][:] + hist = hist + hdfout["hist_i"][:] + counts_cen = counts_cen + hdfout["counts_cen_i"][:] + counts_sat = counts_sat + hdfout["counts_sat_i"][:] + mstar_wcounts += hdfout["mstar_wcounts_i"][:] + mstar_counts += hdfout["mstar_counts_i"][:] + mstar_ssfr_wcounts_cent += hdfout["mstar_ssfr_wcounts_cent_i"][:] + mstar_ssfr_wcounts_sat += hdfout["mstar_ssfr_wcounts_sat_i"][:] + + haloes_data.append( + ( + hdfout["logmh_id"][:], + hdfout["logmh_val"][:], + hdfout["mah_params_samp"][:], + hdfout["ms_params_samp"][:], + hdfout["q_params_samp"][:], + hdfout["upid_samp"][:], + hdfout["tobs_id"][:], + hdfout["tobs_val"][:], + hdfout["redshift_val"][:], + ) + ) + + sampled_haloes = smhm_utils.concatenate_samples_haloes(haloes_data) + end = time() + runtime = end - start + print("Ended collecting data. Time: %.2f seconds." % (end - start)) + + print("Saving final target data...") + + fnout = os.path.join(outdrn, "smdpl_smhm.h5") + with h5py.File(fnout, "w") as hdfout: + hdfout["counts_diff"] = wcounts + hdfout["hist_diff"] = whist + hdfout["counts"] = counts + hdfout["hist"] = hist + hdfout["counts_cen"] = counts_cen + hdfout["counts_sat"] = counts_sat + hdfout["smhm_diff"] = whist / wcounts + hdfout["smhm"] = hist / counts + hdfout["logmh_bins"] = smhm_utils.LOGMH_BINS + hdfout["redshift_targets"] = redshift_targets + hdfout["age_targets"] = age_targets + + ( + logmh_id, + logmh_val, + mah_params_samp, + ms_params_samp, + q_params_samp, + upid_samp, + tobs_id, + tobs_val, + redshift_val, + ) = sampled_haloes + + fnout = os.path.join(outdrn, "smdpl_smhm_samples_haloes.h5") + with h5py.File(fnout, "w") as hdfout: + hdfout["logmh_id"] = logmh_id + hdfout["logmh_val"] = logmh_val + hdfout["mah_params_samp"] = mah_params_samp + hdfout["ms_params_samp"] = ms_params_samp + hdfout["q_params_samp"] = q_params_samp + hdfout["upid_samp"] = upid_samp + hdfout["tobs_id"] = tobs_id + hdfout["tobs_val"] = tobs_val + hdfout["redshift_val"] = redshift_val + + fnout = os.path.join(outdrn, "smdpl_mstar_ssfr.h5") + with h5py.File(fnout, "w") as hdfout: + hdfout["mstar_wcounts"] = mstar_wcounts + hdfout["mstar_counts"] = mstar_counts + hdfout["mstar_ssfr_wcounts_cent"] = mstar_ssfr_wcounts_cent + hdfout["mstar_ssfr_wcounts_sat"] = mstar_ssfr_wcounts_sat + hdfout["logmh_bins"] = smhm_utils.LOGMH_BINS + hdfout["logmstar_bins_pdf"] = smhm_utils.LOGMSTAR_BINS_PDF + hdfout["logssfr_bins_pdf"] = smhm_utils.LOGSSFR_BINS_PDF + hdfout["redshift_targets"] = redshift_targets + hdfout["age_targets"] = age_targets + + # clean up all temporary data + bnpat = os.path.join(outdrn, "_tmp_subvol_*") + command = "rm " + bnpat + subprocess.os.system(command) diff --git a/scripts/diffstarpop_scripts/smhm_utils_galacticus_mgash.py b/scripts/diffstarpop_scripts/smhm_utils_galacticus_mgash.py new file mode 100644 index 0000000..395602f --- /dev/null +++ b/scripts/diffstarpop_scripts/smhm_utils_galacticus_mgash.py @@ -0,0 +1,516 @@ +""" """ + +import os + +import h5py +import numpy as np +from diffmah.diffmah_kernels import DEFAULT_MAH_PARAMS, mah_halopop +from diffsky.diffndhist import tw_ndhist_weighted +from diffstar.defaults_mgash_model import DEFAULT_DIFFSTAR_PARAMS, LGT0, T_TABLE_MIN +from diffstar.sfh_model_mgash import calc_sfh_galpop +from diffstar.data_loaders.load_galacticus_sfh import load_galacticus_diffstar_data +from scipy.stats import binned_statistic +from astropy.cosmology import Planck13 +from umachine_pyio.load_mock import load_mock_from_binaries + +BEBOP_GALAC = "/lcrc/project/halotools/Galacticus/diffstarpop_data/" +BEBOP_GALAC_SFH = "/lcrc/project/halotools/alarcon/results/mgash/Galacticus/" + +LGMH_MIN, LGMH_MAX = 11, 14.50 +N_LGM_BINS = 12 +LOGMH_BINS = np.linspace(LGMH_MIN, LGMH_MAX, N_LGM_BINS) +LOGMSTAR_BINS_PDF = np.linspace(7.0, 13.0, 26) +LOGSSFR_BINS_PDF = np.linspace(-13, -8, 30) + +Z_BINS = [0.0, 0.5, 1.0, 1.5, 2.0] + +T0_SMDPL = 13.7976158 + +N_HALOS_MAX = 20_000 + + +def load_diffstar_sfh_tables( + sfh_type, + diffmah_drn=BEBOP_GALAC, + diffstar_drn=BEBOP_GALAC, + lgt0=LGT0, + n_times=200, +): + + data = load_galacticus_diffstar_data( + BEBOP_GALAC, diffstar_drn=diffstar_drn, diffmah_drn=diffmah_drn + ) + + diffmah_data = data.diffmah_fit_data + if sfh_type == "in_situ": + diffstar_data = data.diffstar_in_situ_fit_data + elif sfh_type == "in_plus_ex_situ": + diffstar_data = data.diffstar_in_plus_ex_situ_fit_data + else: + raise NotImplementedError + + has_fit = ( + (diffmah_data["loss"] > 0.0) + & (diffstar_data["loss"] > 0.0) + & (diffstar_data["success"] == 1) + ) + + is_cen = data.galcus_sfh_data["is_cen"][has_fit] + + mah_params = DEFAULT_MAH_PARAMS._make( + [diffmah_data[key][has_fit] for key in DEFAULT_MAH_PARAMS._fields] + ) + + ms_params = DEFAULT_DIFFSTAR_PARAMS.ms_params._make( + [ + diffstar_data[key][has_fit] + for key in DEFAULT_DIFFSTAR_PARAMS.ms_params._fields + ] + ) + q_params = DEFAULT_DIFFSTAR_PARAMS.q_params._make( + [ + diffstar_data[key][has_fit] + for key in DEFAULT_DIFFSTAR_PARAMS.q_params._fields + ] + ) + + sfh_params = DEFAULT_DIFFSTAR_PARAMS._make((ms_params, q_params)) + + t_0 = 10**lgt0 + t_table = np.linspace(T_TABLE_MIN, t_0, n_times) + + __, log_mah_table = mah_halopop(mah_params, t_table, LGT0) + + sfh_table, smh_table = calc_sfh_galpop( + sfh_params, mah_params, t_table, lgt0=LGT0, return_smh=True + ) + log_sfh_table = np.log10(sfh_table) + log_smh_table = np.log10(smh_table) + log_ssfrh_table = log_sfh_table - log_smh_table + + out = ( + t_table, + log_mah_table, + log_smh_table, + log_ssfrh_table, + mah_params, + ms_params, + q_params, + is_cen, + has_fit, + ) + + return out + + +def compute_diff_histograms_atz(logmh_bins, log_mah_table, log_smh_table): + + n_halos = log_smh_table.shape[0] + + nddata = log_mah_table.reshape((-1, 1)) + + sigma = np.mean(np.diff(logmh_bins)) + np.zeros(n_halos) + ndsig = sigma.reshape((-1, 1)) + + ydata = log_smh_table.reshape((-1, 1)) + _ones = np.ones_like(ydata) + + ndbins_lo = logmh_bins[:-1].reshape((-1, 1)) + ndbins_hi = logmh_bins[1:].reshape((-1, 1)) + + whist = tw_ndhist_weighted(nddata, ndsig, ydata, ndbins_lo, ndbins_hi) + wcounts = tw_ndhist_weighted(nddata, ndsig, _ones, ndbins_lo, ndbins_hi) + + return wcounts, whist + + +def compute_histograms_atz(logmh_bins, log_mah_table, log_smh_table): + count = binned_statistic( + log_mah_table, values=log_smh_table, bins=logmh_bins, statistic="count" + )[0] + whist = binned_statistic( + log_mah_table, values=log_smh_table, bins=logmh_bins, statistic="sum" + )[0] + return count, whist + + +def get_redshift_from_age(age): + z_table = np.linspace(0, 10, 2000)[::-1] + t_table = Planck13.age(z_table).value + redshift_from_age = np.interp(age, t_table, z_table) + return redshift_from_age + + +def return_target_redshfit_index(t_table, redshift_targets): + z_table = get_redshift_from_age(t_table) + return np.digitize(redshift_targets, z_table) + + +def sample_halos( + logmh_bins, + log_mah, + log_smh, + mah_params, + ms_params, + q_params, + upid, +): + ndbins_lo = logmh_bins[:-1] + ndbins_hi = logmh_bins[1:] + arange_arr = np.arange(len(log_mah)) + logmh_id = [] + logmh_val = [] + mah_params_samp = [] + ms_params_samp = [] + q_params_samp = [] + upid_samp = [] + + mah_params = np.array(mah_params).T + ms_params = np.array(ms_params).T + q_params = np.array(q_params).T + + for i in range(len(ndbins_lo)): + sel = (log_mah >= ndbins_lo[i]) & (log_mah < ndbins_hi[i]) + if sel.sum() == 0: + continue + sel_num = int(min(N_HALOS_MAX, sel.sum())) + sel = np.random.choice(arange_arr[sel], sel_num, replace=False) + logmh_id.append(np.ones_like(sel) * i) + logmh_val.append(np.ones_like(sel) * ((ndbins_lo[i] + ndbins_hi[i]) / 2.0)) + mah_params_samp.append(mah_params[sel]) + ms_params_samp.append(ms_params[sel]) + q_params_samp.append(q_params[sel]) + upid_samp.append(upid[sel]) + + logmh_id = np.concatenate(logmh_id) + logmh_val = np.concatenate(logmh_val) + mah_params_samp = np.concatenate(mah_params_samp) + ms_params_samp = np.concatenate(ms_params_samp) + q_params_samp = np.concatenate(q_params_samp) + upid_samp = np.concatenate(upid_samp) + + mah_params_samp = DEFAULT_MAH_PARAMS._make(mah_params_samp.T) + ms_params_samp = DEFAULT_DIFFSTAR_PARAMS.ms_params._make(ms_params_samp.T) + q_params_samp = DEFAULT_DIFFSTAR_PARAMS.q_params._make(q_params_samp.T) + out = ( + logmh_id, + logmh_val, + mah_params_samp, + ms_params_samp, + q_params_samp, + upid_samp, + ) + return out + + +def create_target_data( + sfh_type, + redshift_targets=Z_BINS, + diffmah_drn=BEBOP_GALAC, + diffstar_drn=BEBOP_GALAC, + lgt0=LGT0, + logmh_bins=LOGMH_BINS, +): + _res = load_diffstar_sfh_tables( + sfh_type, + diffmah_drn=diffmah_drn, + diffstar_drn=diffstar_drn, + lgt0=lgt0, + ) + ( + t_table, + log_mah_table, + log_smh_table, + log_ssfrh_table, + mah_params, + ms_params, + q_params, + is_cen, + has_fit, + ) = _res + + _path = os.path.join(diffmah_drn, "tarr_disk.npy") + tarr = np.load(_path) + + tids = return_target_redshfit_index(t_table, redshift_targets) + tids_galac = return_target_redshfit_index(tarr, redshift_targets) + + nz, nm = len(redshift_targets), len(logmh_bins) - 1 + + wcounts_zid = np.zeros((nz, nm)) + whist_zid = np.zeros((nz, nm)) + counts_zid = np.zeros((nz, nm)) + hist_zid = np.zeros((nz, nm)) + + counts_zid_cen = np.zeros((nz, nm)) + counts_zid_sat = np.zeros((nz, nm)) + + is_central = is_cen == 1 + is_satell = is_cen == 0 + + for i, tid in enumerate(tids): + _res = compute_diff_histograms_atz( + logmh_bins, log_mah_table[:, tid], log_smh_table[:, tid] + ) + wcounts_zid[i] = _res[0] + whist_zid[i] = _res[1] + + _res = compute_histograms_atz( + logmh_bins, log_mah_table[:, tid], log_smh_table[:, tid] + ) + counts_zid[i] = _res[0] + hist_zid[i] = _res[1] + + counts_zid_cen[i] = compute_histograms_atz( + logmh_bins, + log_mah_table[:, tid][is_central], + log_smh_table[:, tid][is_central], + )[0] + counts_zid_sat[i] = compute_histograms_atz( + logmh_bins, + log_mah_table[:, tid][is_satell], + log_smh_table[:, tid][is_satell], + )[0] + + data = [] + + final_upid = is_cen.copy() + final_upid[is_central] = -1 + + for i, tid in enumerate(tids): + _res = sample_halos( + logmh_bins, + log_mah_table[:, tid], + log_smh_table[:, tid], + mah_params, + ms_params, + q_params, + final_upid, + ) + data.append( + ( + *_res, + np.ones_like(_res[0]) * i, + np.ones_like(_res[0]) * t_table[tid], + np.ones_like(_res[0]) * redshift_targets[i], + ) + ) + + haloes = concatenate_samples_haloes(data) + + out = ( + wcounts_zid, + whist_zid, + counts_zid, + hist_zid, + t_table[tids], + haloes, + counts_zid_cen, + counts_zid_sat, + ) + + return out + + +def concatenate_samples_haloes(data): + logmh_id = [] + logmh_val = [] + redshift_val = [] + tobs_id = [] + tobs_val = [] + mah_params_samp = [] + ms_params_samp = [] + q_params_samp = [] + upid_samp = [] + + for subdata in data: + + logmh_id.append(subdata[0]) + logmh_val.append(subdata[1]) + mah_params_samp.append(np.array(subdata[2]).T) + ms_params_samp.append(np.array(subdata[3]).T) + q_params_samp.append(np.array(subdata[4]).T) + upid_samp.append(subdata[5]) + tobs_id.append(subdata[6]) + tobs_val.append(subdata[7]) + redshift_val.append(subdata[8]) + + logmh_id = np.concatenate(logmh_id) + logmh_val = np.concatenate(logmh_val) + mah_params_samp = np.concatenate(mah_params_samp) + ms_params_samp = np.concatenate(ms_params_samp) + q_params_samp = np.concatenate(q_params_samp) + upid_samp = np.concatenate(upid_samp) + tobs_id = np.concatenate(tobs_id) + tobs_val = np.concatenate(tobs_val) + redshift_val = np.concatenate(redshift_val) + + mah_params_samp = DEFAULT_MAH_PARAMS._make(mah_params_samp.T) + ms_params_samp = DEFAULT_DIFFSTAR_PARAMS.ms_params._make(ms_params_samp.T) + q_params_samp = DEFAULT_DIFFSTAR_PARAMS.q_params._make(q_params_samp.T) + + haloes = ( + logmh_id, + logmh_val, + mah_params_samp, + ms_params_samp, + q_params_samp, + upid_samp, + tobs_id, + tobs_val, + redshift_val, + ) + return haloes + + +def compute_diff_histograms_mstar_atmobs_z( + logmstar_bins, + log_smh_table, +): + + n_halos = log_smh_table.shape[0] + + nddata = log_smh_table.reshape((-1, 1)) + + sigma = np.mean(np.diff(logmstar_bins)) + np.zeros(n_halos) + ndsig = sigma.reshape((-1, 1)) + + ydata = log_smh_table.reshape((-1, 1)) + _ones = np.ones_like(ydata) + + ndbins_lo = logmstar_bins[:-1].reshape((-1, 1)) + ndbins_hi = logmstar_bins[1:].reshape((-1, 1)) + + wcounts = tw_ndhist_weighted(nddata, ndsig, _ones, ndbins_lo, ndbins_hi) + + counts = np.histogram(ydata, logmstar_bins)[0] + + return wcounts, counts + + +def compute_diff_histograms_mstar_ssfr_atz( + log_smh_table, + log_ssfr_table, + ndbins_lo, + ndbins_hi, + logmstar_bins, + logssfr_bins, +): + n_halos = log_smh_table.shape[0] + + sigma_mstar = np.mean(np.diff(logmstar_bins)) + np.zeros(n_halos) + sigma_ssfr = np.mean(np.diff(logssfr_bins)) + np.zeros(n_halos) + + ndsig = np.ones((n_halos, 2)) + ndsig[:, 0] = sigma_mstar + ndsig[:, 1] = sigma_ssfr + + nddata = np.array([log_smh_table, log_ssfr_table]).T + + _ones = np.ones(n_halos) + + wcounts = tw_ndhist_weighted(nddata, ndsig, _ones, ndbins_lo, ndbins_hi) + + return wcounts + + +def create_pdf_target_data( + sfh_type, + redshift_targets=Z_BINS, + diffmah_drn=BEBOP_GALAC, + diffstar_drn=BEBOP_GALAC, + lgt0=LGT0, + logmh_bins=LOGMH_BINS, + logmstar_bins_pdf=LOGMSTAR_BINS_PDF, + logssfr_bins_pdf=LOGSSFR_BINS_PDF, +): + _res = load_diffstar_sfh_tables( + sfh_type, + diffmah_drn=diffmah_drn, + diffstar_drn=diffstar_drn, + lgt0=lgt0, + ) + ( + t_table, + log_mah_table, + log_smh_table, + log_ssfrh_table, + mah_params, + ms_params, + q_params, + is_cen, + has_fit, + ) = _res + + log_ssfrh_table = np.clip(log_ssfrh_table, -12.0, None) + + _path = os.path.join(diffmah_drn, "tarr_disk.npy") + tarr = np.load(_path) + + tids = return_target_redshfit_index(t_table, redshift_targets) + tids_galac = return_target_redshfit_index(tarr, redshift_targets) + + nz, nm = len(redshift_targets), len(logmh_bins) - 1 + nmstar = len(logmstar_bins_pdf) - 1 + nssfr = len(logssfr_bins_pdf) - 1 + + ndbins_lo = [] + ndbins_hi = [] + for i in range(len(logmstar_bins_pdf) - 1): + for j in range(len(logssfr_bins_pdf) - 1): + ndbins_lo.append([logmstar_bins_pdf[i], logssfr_bins_pdf[j]]) + ndbins_hi.append([logmstar_bins_pdf[i + 1], logssfr_bins_pdf[j + 1]]) + ndbins_lo = np.array(ndbins_lo) + ndbins_hi = np.array(ndbins_hi) + + mstar_wcounts = np.zeros((nz, nm, nmstar)) + mstar_counts = np.zeros((nz, nm, nmstar)) + + mstar_ssfr_wcounts_cent = np.zeros((nz, nm, nmstar, nssfr)) + mstar_ssfr_wcounts_sat = np.zeros((nz, nm, nmstar, nssfr)) + + is_central = is_cen == 1 + is_satell = is_cen == 0 + for i, tid in enumerate(tids): + + for j in range(nm): + mobs_sel = (log_mah_table[:, tid] > logmh_bins[j]) & ( + log_mah_table[:, tid] < logmh_bins[j + 1] + ) + _res = compute_diff_histograms_mstar_atmobs_z( + logmstar_bins_pdf, + log_smh_table[mobs_sel][:, tid], + ) + mstar_wcounts[i, j] = _res[0] + mstar_counts[i, j] = _res[1] + + mobs_sel_cent = mobs_sel & is_central + _res = compute_diff_histograms_mstar_ssfr_atz( + log_smh_table[mobs_sel_cent][:, tid], + log_ssfrh_table[mobs_sel_cent][:, tid], + ndbins_lo, + ndbins_hi, + logmstar_bins_pdf, + logssfr_bins_pdf, + ) + mstar_ssfr_wcounts_cent[i, j] = _res.reshape((nmstar, nssfr)) + + mobs_sel_sat = mobs_sel & is_satell + _res = compute_diff_histograms_mstar_ssfr_atz( + log_smh_table[mobs_sel_sat][:, tid], + log_ssfrh_table[mobs_sel_sat][:, tid], + ndbins_lo, + ndbins_hi, + logmstar_bins_pdf, + logssfr_bins_pdf, + ) + mstar_ssfr_wcounts_sat[i, j] = _res.reshape((nmstar, nssfr)) + + out = ( + mstar_wcounts, + mstar_counts, + mstar_ssfr_wcounts_cent, + mstar_ssfr_wcounts_sat, + ) + + return out diff --git a/scripts/diffstarpop_scripts/smhm_utils_smdpl_mgash.py b/scripts/diffstarpop_scripts/smhm_utils_smdpl_mgash.py new file mode 100644 index 0000000..aa4c36f --- /dev/null +++ b/scripts/diffstarpop_scripts/smhm_utils_smdpl_mgash.py @@ -0,0 +1,650 @@ +""" """ + +import os +import re +import h5py +import numpy as np +from diffmah.diffmah_kernels import DEFAULT_MAH_PARAMS, mah_halopop +from diffsky.diffndhist import tw_ndhist_weighted +from diffstar.defaults_mgash_model import DEFAULT_DIFFSTAR_PARAMS, LGT0, T_TABLE_MIN +from diffstar.sfh_model_mgash import calc_sfh_galpop +from scipy.stats import binned_statistic +from astropy.cosmology import Planck13 +from umachine_pyio.load_mock import load_mock_from_binaries + +LCRC_NOMERGING_DIFFSTAR_DRN = ( + "/lcrc/project/halotools/alarcon/results/mgash/UniverseMachine/DR1_nomerging/" +) +LCRC_NOMERGING_DIFFMAH_DRN = ( + "/lcrc/project/halotools/SMDPL/dr1_no_merging_upidh/diffmah_tpeak_fits/" +) +LCRC_NOMERGING_BINARIES_DRN = ( + "/lcrc/project/halotools/SMDPL/dr1_no_merging_upidh/sfh_binary_catalogs/a_1.000000/" +) + +LCRC_DR1_DIFFSTAR_DRN = ( + "/lcrc/project/halotools/alarcon/results/mgash/UniverseMachine/DR1/" +) +LCRC_DR1_DIFFMAH_DRN = "/lcrc/project/halotools/UniverseMachine/SMDPL/sfh_binaries_dr1_bestfit/diffmah_tpeak_fits/" +LCRC_DR1_BINARIES_DRN = ( + "/lcrc/project/halotools/UniverseMachine/SMDPL/sfh_binaries_dr1_bestfit/a_1.000000/" +) +LCRC_NOMERGING_diffstar_bnpat = "diffstar_fits_subvol_{}.hdf5" +LCRC_DR1_diffstar_bnpat = "diffstar_fits_subvol_{}.hdf5" +LCRC_NOMERGING_diffmah_bnpat = "subvol_{}_diffmah_fits.h5" + +TASSO_DIFFSTAR_DRN = "/Users/aphearin/work/DATA/diffstar_data/SMDPL/" +N_SUBVOL_SMDPL = 576 + +LGMH_MIN, LGMH_MAX = 11, 14.75 +N_LGM_BINS = 12 +LOGMH_BINS = np.linspace(LGMH_MIN, LGMH_MAX, N_LGM_BINS) + +LOGMSTAR_BINS_PDF = np.linspace(7.0, 13.0, 26) +LOGSSFR_BINS_PDF = np.linspace(-13.0, -8.0, 30) + +Z_BINS = [0.0, 0.5, 1.0, 1.5, 2.0] + +T0_SMDPL = 13.7976158 + +N_HALOS_MAX = 20_000 +N_HALOS_PER_SUBVOL = N_HALOS_MAX // N_SUBVOL_SMDPL + + +def _load_flat_hdf5(fn): + data = dict() + with h5py.File(fn, "r") as hdf: + for key in hdf.keys(): + data[key] = hdf[key][...] + return data + + +def return_subvol_str_diffmah(subvol, sim_name, diffstar_drn, diffstar_bnpat): + regex_str = re.escape(diffstar_bnpat).replace(r"\{\}", r"(\d{1,3})") + pattern = re.compile(f"^{regex_str}$") + matching_files = [f for f in os.listdir(diffstar_drn) if pattern.match(f)] + if sim_name == "DR1_nomerging": + subvols = [x.split("_")[1] for x in matching_files] + elif sim_name == "DR1": + subvols = [x.split("_")[-1].split(".")[0] for x in matching_files] + subvols_len = np.array([len(x) for x in subvols]) + + if np.any(subvols_len == 1): + subvol_str = f"{subvol:d}" + elif np.all(subvols_len == subvols_len.max()): + nchar_subvol = subvols_len.max() + subvol_str = f"{subvol:0{nchar_subvol}d}" + return subvol_str + + +def return_subvol_str(subvol, sim_name, diffstar_drn, diffstar_bnpat): + regex_str = re.escape(diffstar_bnpat).replace(r"\{\}", r"(\d{1,3})") + pattern = re.compile(f"^{regex_str}$") + matching_files = [f for f in os.listdir(diffstar_drn) if pattern.match(f)] + subvols = [x.split("_")[-1].split(".")[0] for x in matching_files] + subvols_len = np.array([len(x) for x in subvols]) + + if np.any(subvols_len == 1): + subvol_str = f"{subvol:d}" + elif np.all(subvols_len == subvols_len.max()): + nchar_subvol = subvols_len.max() + subvol_str = f"{subvol:0{nchar_subvol}d}" + return subvol_str + + +def load_diffstar_subvolume( + subvol, + sim_name, + n_subvol_tot=N_SUBVOL_SMDPL, + diffmah_drn=TASSO_DIFFSTAR_DRN, + diffstar_drn=TASSO_DIFFSTAR_DRN, + diffstar_bnpat=LCRC_NOMERGING_diffstar_bnpat, +): + # nchar_subvol = len(str(n_subvol_tot)) + subvol_str = return_subvol_str(subvol, sim_name, diffstar_drn, diffstar_bnpat) + + diffstar_bn = diffstar_bnpat.format(subvol_str) + diffstar_fn = os.path.join(diffstar_drn, diffstar_bn) + diffstar_data = _load_flat_hdf5(diffstar_fn) + + if sim_name == "DR1_nomerging": + subvol_str = return_subvol_str_diffmah( + subvol, sim_name, diffmah_drn, LCRC_NOMERGING_diffmah_bnpat + ) + diffmah_bn = LCRC_NOMERGING_diffmah_bnpat.format(subvol_str).replace( + "diffstar", "diffmah" + ) + elif sim_name == "DR1": + diffmah_bn = diffstar_bn.replace("diffstar", "diffmah") + diffmah_fn = os.path.join(diffmah_drn, diffmah_bn) + diffmah_data = _load_flat_hdf5(diffmah_fn) + + return diffmah_data, diffstar_data + + +def load_diffstar_sfh_tables( + subvol, + sim_name, + n_subvol_tot=N_SUBVOL_SMDPL, + diffmah_drn=TASSO_DIFFSTAR_DRN, + diffstar_drn=TASSO_DIFFSTAR_DRN, + diffstar_bnpat=LCRC_NOMERGING_diffstar_bnpat, + lgt0=LGT0, + n_times=200, +): + diffmah_data, diffstar_data = load_diffstar_subvolume( + subvol, + sim_name, + n_subvol_tot=n_subvol_tot, + diffmah_drn=diffmah_drn, + diffstar_drn=diffstar_drn, + diffstar_bnpat=diffstar_bnpat, + ) + has_fit = (diffmah_data["loss"] > 0.0) & (diffstar_data["success"] == 1) + mah_params = DEFAULT_MAH_PARAMS._make( + [diffmah_data[key][has_fit] for key in DEFAULT_MAH_PARAMS._fields] + ) + + ms_params = DEFAULT_DIFFSTAR_PARAMS.ms_params._make( + [ + diffstar_data[key][has_fit] + for key in DEFAULT_DIFFSTAR_PARAMS.ms_params._fields + ] + ) + q_params = DEFAULT_DIFFSTAR_PARAMS.q_params._make( + [ + diffstar_data[key][has_fit] + for key in DEFAULT_DIFFSTAR_PARAMS.q_params._fields + ] + ) + sfh_params = DEFAULT_DIFFSTAR_PARAMS._make((ms_params, q_params)) + + t_0 = 10**lgt0 + t_table = np.linspace(T_TABLE_MIN, t_0, n_times) + + __, log_mah_table = mah_halopop(mah_params, t_table, LGT0) + + sfh_table, smh_table = calc_sfh_galpop( + sfh_params, mah_params, t_table, lgt0=LGT0, return_smh=True + ) + log_sfh_table = np.log10(sfh_table) + log_smh_table = np.log10(smh_table) + log_ssfrh_table = log_sfh_table - log_smh_table + + out = ( + t_table, + log_mah_table, + log_smh_table, + log_ssfrh_table, + mah_params, + ms_params, + q_params, + has_fit, + ) + + return out + + +def compute_weighted_histograms_z0( + subvol, + sim_name, + n_subvol_tot=N_SUBVOL_SMDPL, + diffmah_drn=TASSO_DIFFSTAR_DRN, + diffstar_drn=TASSO_DIFFSTAR_DRN, + diffstar_bnpat=LCRC_NOMERGING_diffstar_bnpat, + lgt0=LGT0, + logmh_bins=LOGMH_BINS, +): + _res = load_diffstar_sfh_tables( + subvol, + sim_name, + n_subvol_tot=n_subvol_tot, + diffmah_drn=diffmah_drn, + diffstar_drn=diffstar_drn, + diffstar_bnpat=diffstar_bnpat, + lgt0=lgt0, + ) + t_table, log_mah_table, log_smh_table, log_ssfrh_table = _res[:4] + + n_halos = log_smh_table.shape[0] + + nddata = log_mah_table[:, -1].reshape((-1, 1)) + + sigma = np.mean(np.diff(logmh_bins)) + np.zeros(n_halos) + ndsig = sigma.reshape((-1, 1)) + + ydata = log_smh_table[:, -1].reshape((-1, 1)) + _ones = np.ones_like(ydata) + + ndbins_lo = logmh_bins[:-1].reshape((-1, 1)) + ndbins_hi = logmh_bins[1:].reshape((-1, 1)) + + whist = tw_ndhist_weighted(nddata, ndsig, ydata, ndbins_lo, ndbins_hi) + wcounts = tw_ndhist_weighted(nddata, ndsig, _ones, ndbins_lo, ndbins_hi) + + return wcounts, whist + + +def compute_diff_histograms_atz(logmh_bins, log_mah_table, log_smh_table): + + n_halos = log_smh_table.shape[0] + + nddata = log_mah_table.reshape((-1, 1)) + + sigma = np.mean(np.diff(logmh_bins)) + np.zeros(n_halos) + ndsig = sigma.reshape((-1, 1)) + + ydata = log_smh_table.reshape((-1, 1)) + _ones = np.ones_like(ydata) + + ndbins_lo = logmh_bins[:-1].reshape((-1, 1)) + ndbins_hi = logmh_bins[1:].reshape((-1, 1)) + + whist = tw_ndhist_weighted(nddata, ndsig, ydata, ndbins_lo, ndbins_hi) + wcounts = tw_ndhist_weighted(nddata, ndsig, _ones, ndbins_lo, ndbins_hi) + + return wcounts, whist + + +def compute_histograms_atz(logmh_bins, log_mah_table, log_smh_table): + count = binned_statistic( + log_mah_table, values=log_smh_table, bins=logmh_bins, statistic="count" + )[0] + whist = binned_statistic( + log_mah_table, values=log_smh_table, bins=logmh_bins, statistic="sum" + )[0] + return count, whist + + +def get_redshift_from_age(age): + z_table = np.linspace(0, 10, 2000)[::-1] + t_table = Planck13.age(z_table).value + redshift_from_age = np.interp(age, t_table, z_table) + return redshift_from_age + + +def return_target_redshfit_index(t_table, redshift_targets): + z_table = get_redshift_from_age(t_table) + return np.digitize(redshift_targets, z_table) + + +def sample_halos( + n_subvol_smdpl, + logmh_bins, + log_mah, + log_smh, + mah_params, + ms_params, + q_params, + upid, +): + ndbins_lo = logmh_bins[:-1] + ndbins_hi = logmh_bins[1:] + arange_arr = np.arange(len(log_mah)) + logmh_id = [] + logmh_val = [] + mah_params_samp = [] + ms_params_samp = [] + q_params_samp = [] + upid_samp = [] + + mah_params = np.array(mah_params).T + ms_params = np.array(ms_params).T + q_params = np.array(q_params).T + + n_halos_per_subvol = N_HALOS_MAX // n_subvol_smdpl + + for i in range(len(ndbins_lo)): + sel = (log_mah >= ndbins_lo[i]) & (log_mah < ndbins_hi[i]) + sel_num = int(min(n_halos_per_subvol, sel.sum())) + sel = np.random.choice(arange_arr[sel], sel_num, replace=False) + logmh_id.append(np.ones_like(sel) * i) + logmh_val.append(np.ones_like(sel) * ((ndbins_lo[i] + ndbins_hi[i]) / 2.0)) + mah_params_samp.append(mah_params[sel]) + ms_params_samp.append(ms_params[sel]) + q_params_samp.append(q_params[sel]) + upid_samp.append(upid[sel]) + + logmh_id = np.concatenate(logmh_id) + logmh_val = np.concatenate(logmh_val) + mah_params_samp = np.concatenate(mah_params_samp) + ms_params_samp = np.concatenate(ms_params_samp) + q_params_samp = np.concatenate(q_params_samp) + upid_samp = np.concatenate(upid_samp) + + mah_params_samp = DEFAULT_MAH_PARAMS._make(mah_params_samp.T) + ms_params_samp = DEFAULT_DIFFSTAR_PARAMS.ms_params._make(ms_params_samp.T) + q_params_samp = DEFAULT_DIFFSTAR_PARAMS.q_params._make(q_params_samp.T) + out = ( + logmh_id, + logmh_val, + mah_params_samp, + ms_params_samp, + q_params_samp, + upid_samp, + ) + return out + + +def create_target_data( + subvol, + sim_name, + n_subvol_smdpl, + redshift_targets=Z_BINS, + n_subvol_tot=N_SUBVOL_SMDPL, + binaries_drn=LCRC_NOMERGING_BINARIES_DRN, + diffmah_drn=LCRC_NOMERGING_DIFFMAH_DRN, + diffstar_drn=LCRC_NOMERGING_DIFFSTAR_DRN, + diffstar_bnpat=LCRC_NOMERGING_diffstar_bnpat, + lgt0=LGT0, + logmh_bins=LOGMH_BINS, +): + _res = load_diffstar_sfh_tables( + subvol, + sim_name, + n_subvol_tot=n_subvol_tot, + diffmah_drn=diffmah_drn, + diffstar_drn=diffstar_drn, + diffstar_bnpat=diffstar_bnpat, + lgt0=lgt0, + ) + ( + t_table, + log_mah_table, + log_smh_table, + log_ssfrh_table, + mah_params, + ms_params, + q_params, + has_fit, + ) = _res + + galprops = ["halo_id", "upid"] + halos = load_mock_from_binaries( + np.atleast_1d(subvol), root_dirname=binaries_drn, galprops=galprops + ) + upid = np.array(halos["upid"])[has_fit] + + tids = return_target_redshfit_index(t_table, redshift_targets) + + nz, nm = len(redshift_targets), len(logmh_bins) - 1 + + wcounts_zid = np.zeros((nz, nm)) + whist_zid = np.zeros((nz, nm)) + counts_zid = np.zeros((nz, nm)) + hist_zid = np.zeros((nz, nm)) + + counts_zid_cen = np.zeros((nz, nm)) + counts_zid_sat = np.zeros((nz, nm)) + + for i, tid in enumerate(tids): + _res = compute_diff_histograms_atz( + logmh_bins, log_mah_table[:, tid], log_smh_table[:, tid] + ) + wcounts_zid[i] = _res[0] + whist_zid[i] = _res[1] + + _res = compute_histograms_atz( + logmh_bins, log_mah_table[:, tid], log_smh_table[:, tid] + ) + counts_zid[i] = _res[0] + hist_zid[i] = _res[1] + + is_central = upid == -1 + counts_zid_cen[i] = compute_histograms_atz( + logmh_bins, + log_mah_table[:, tid][is_central], + log_smh_table[:, tid][is_central], + )[0] + counts_zid_sat[i] = compute_histograms_atz( + logmh_bins, + log_mah_table[:, tid][~is_central], + log_smh_table[:, tid][~is_central], + )[0] + + data = [] + + for i, tid in enumerate(tids): + _res = sample_halos( + n_subvol_smdpl, + logmh_bins, + log_mah_table[:, tid], + log_smh_table[:, tid], + mah_params, + ms_params, + q_params, + upid, + ) + data.append( + ( + *_res, + np.ones_like(_res[0]) * i, + np.ones_like(_res[0]) * t_table[tid], + np.ones_like(_res[0]) * redshift_targets[i], + ) + ) + + haloes = concatenate_samples_haloes(data) + + out = ( + wcounts_zid, + whist_zid, + counts_zid, + hist_zid, + t_table[tids], + haloes, + counts_zid_cen, + counts_zid_sat, + ) + + return out + + +def concatenate_samples_haloes(data): + logmh_id = [] + logmh_val = [] + redshift_val = [] + tobs_id = [] + tobs_val = [] + mah_params_samp = [] + ms_params_samp = [] + q_params_samp = [] + upid_samp = [] + + for subdata in data: + + logmh_id.append(subdata[0]) + logmh_val.append(subdata[1]) + mah_params_samp.append(np.array(subdata[2]).T) + ms_params_samp.append(np.array(subdata[3]).T) + q_params_samp.append(np.array(subdata[4]).T) + upid_samp.append(subdata[5]) + tobs_id.append(subdata[6]) + tobs_val.append(subdata[7]) + redshift_val.append(subdata[8]) + + logmh_id = np.concatenate(logmh_id) + logmh_val = np.concatenate(logmh_val) + mah_params_samp = np.concatenate(mah_params_samp) + ms_params_samp = np.concatenate(ms_params_samp) + q_params_samp = np.concatenate(q_params_samp) + upid_samp = np.concatenate(upid_samp) + tobs_id = np.concatenate(tobs_id) + tobs_val = np.concatenate(tobs_val) + redshift_val = np.concatenate(redshift_val) + + mah_params_samp = DEFAULT_MAH_PARAMS._make(mah_params_samp.T) + ms_params_samp = DEFAULT_DIFFSTAR_PARAMS.ms_params._make(ms_params_samp.T) + q_params_samp = DEFAULT_DIFFSTAR_PARAMS.q_params._make(q_params_samp.T) + + haloes = ( + logmh_id, + logmh_val, + mah_params_samp, + ms_params_samp, + q_params_samp, + upid_samp, + tobs_id, + tobs_val, + redshift_val, + ) + return haloes + + +def compute_diff_histograms_mstar_atmobs_z( + logmstar_bins, + log_smh_table, +): + + n_halos = log_smh_table.shape[0] + + nddata = log_smh_table.reshape((-1, 1)) + + sigma = np.mean(np.diff(logmstar_bins)) + np.zeros(n_halos) + ndsig = sigma.reshape((-1, 1)) + + ydata = log_smh_table.reshape((-1, 1)) + _ones = np.ones_like(ydata) + + ndbins_lo = logmstar_bins[:-1].reshape((-1, 1)) + ndbins_hi = logmstar_bins[1:].reshape((-1, 1)) + + wcounts = tw_ndhist_weighted(nddata, ndsig, _ones, ndbins_lo, ndbins_hi) + + counts = np.histogram(ydata, logmstar_bins)[0] + + return wcounts, counts + + +def compute_diff_histograms_mstar_ssfr_atz( + log_smh_table, + log_ssfr_table, + ndbins_lo, + ndbins_hi, + logmstar_bins, + logssfr_bins, +): + n_halos = log_smh_table.shape[0] + + sigma_mstar = np.mean(np.diff(logmstar_bins)) + np.zeros(n_halos) + sigma_ssfr = np.mean(np.diff(logssfr_bins)) + np.zeros(n_halos) + + ndsig = np.ones((n_halos, 2)) + ndsig[:, 0] = sigma_mstar + ndsig[:, 1] = sigma_ssfr + + nddata = np.array([log_smh_table, log_ssfr_table]).T + + _ones = np.ones(n_halos) + + wcounts = tw_ndhist_weighted(nddata, ndsig, _ones, ndbins_lo, ndbins_hi) + + return wcounts + + +def create_pdf_target_data( + subvol, + sim_name, + redshift_targets=Z_BINS, + n_subvol_tot=N_SUBVOL_SMDPL, + binaries_drn=LCRC_NOMERGING_BINARIES_DRN, + diffmah_drn=LCRC_NOMERGING_DIFFMAH_DRN, + diffstar_drn=LCRC_NOMERGING_DIFFSTAR_DRN, + diffstar_bnpat=LCRC_NOMERGING_diffstar_bnpat, + lgt0=LGT0, + logmh_bins=LOGMH_BINS, + logmstar_bins_pdf=LOGMSTAR_BINS_PDF, + logssfr_bins_pdf=LOGSSFR_BINS_PDF, +): + _res = load_diffstar_sfh_tables( + subvol, + sim_name, + n_subvol_tot=n_subvol_tot, + diffmah_drn=diffmah_drn, + diffstar_drn=diffstar_drn, + diffstar_bnpat=diffstar_bnpat, + lgt0=lgt0, + ) + ( + t_table, + log_mah_table, + log_smh_table, + log_ssfrh_table, + mah_params, + ms_params, + q_params, + has_fit, + ) = _res + + log_ssfrh_table = np.clip(log_ssfrh_table, -12.0, None) + + galprops = ["halo_id", "upid"] + halos = load_mock_from_binaries( + np.atleast_1d(subvol), root_dirname=binaries_drn, galprops=galprops + ) + upid = np.array(halos["upid"])[has_fit] + is_central = upid == -1 + + tids = return_target_redshfit_index(t_table, redshift_targets) + + nz, nm = len(redshift_targets), len(logmh_bins) - 1 + nmstar = len(logmstar_bins_pdf) - 1 + nssfr = len(logssfr_bins_pdf) - 1 + + ndbins_lo = [] + ndbins_hi = [] + for i in range(len(logmstar_bins_pdf) - 1): + for j in range(len(logssfr_bins_pdf) - 1): + ndbins_lo.append([logmstar_bins_pdf[i], logssfr_bins_pdf[j]]) + ndbins_hi.append([logmstar_bins_pdf[i + 1], logssfr_bins_pdf[j + 1]]) + ndbins_lo = np.array(ndbins_lo) + ndbins_hi = np.array(ndbins_hi) + + mstar_wcounts = np.zeros((nz, nm, nmstar)) + mstar_counts = np.zeros((nz, nm, nmstar)) + + mstar_ssfr_wcounts_cent = np.zeros((nz, nm, nmstar, nssfr)) + mstar_ssfr_wcounts_sat = np.zeros((nz, nm, nmstar, nssfr)) + + for i, tid in enumerate(tids): + for j in range(nm): + mobs_sel = (log_mah_table[:, tid] > logmh_bins[j]) & ( + log_mah_table[:, tid] < logmh_bins[j + 1] + ) + _res = compute_diff_histograms_mstar_atmobs_z( + logmstar_bins_pdf, + log_smh_table[mobs_sel][:, tid], + ) + mstar_wcounts[i, j] = _res[0] + mstar_counts[i, j] = _res[1] + + mobs_sel_cent = mobs_sel & is_central + _res = compute_diff_histograms_mstar_ssfr_atz( + log_smh_table[mobs_sel_cent][:, tid], + log_ssfrh_table[mobs_sel_cent][:, tid], + ndbins_lo, + ndbins_hi, + logmstar_bins_pdf, + logssfr_bins_pdf, + ) + mstar_ssfr_wcounts_cent[i, j] = _res.reshape((nmstar, nssfr)) + + mobs_sel_sat = mobs_sel & (~is_central) + _res = compute_diff_histograms_mstar_ssfr_atz( + log_smh_table[mobs_sel_sat][:, tid], + log_ssfrh_table[mobs_sel_sat][:, tid], + ndbins_lo, + ndbins_hi, + logmstar_bins_pdf, + logssfr_bins_pdf, + ) + mstar_ssfr_wcounts_sat[i, j] = _res.reshape((nmstar, nssfr)) + + out = ( + mstar_wcounts, + mstar_counts, + mstar_ssfr_wcounts_cent, + mstar_ssfr_wcounts_sat, + ) + + return out diff --git a/scripts/diffstarpop_scripts/smhm_utils_tng_mgash.py b/scripts/diffstarpop_scripts/smhm_utils_tng_mgash.py new file mode 100644 index 0000000..eba419f --- /dev/null +++ b/scripts/diffstarpop_scripts/smhm_utils_tng_mgash.py @@ -0,0 +1,585 @@ +""" """ + +import os + +import h5py +import numpy as np +from diffmah.diffmah_kernels import DEFAULT_MAH_PARAMS, mah_halopop +from diffsky.diffndhist import tw_ndhist_weighted +from diffstar.defaults_mgash_model import DEFAULT_DIFFSTAR_PARAMS, LGT0, T_TABLE_MIN +from diffstar.sfh_model_mgash import calc_sfh_galpop +from scipy.stats import binned_statistic +from astropy.cosmology import Planck13 +from umachine_pyio.load_mock import load_mock_from_binaries + +BEBOP_TNG = "/lcrc/project/halotools/alarcon/data/" +BEBOP_TNG_MAH = "/lcrc/project/halotools/alarcon/results/tng_diffmah_tpeak/" +BEBOP_TNG_SFH = "/lcrc/project/halotools/alarcon/results/mgash/TNG/" + + +LGMH_MIN, LGMH_MAX = 11, 14.75 +N_LGM_BINS = 12 +LOGMH_BINS = np.linspace(LGMH_MIN, LGMH_MAX, N_LGM_BINS) + +LOGMSTAR_BINS_PDF = np.linspace(7.0, 13.0, 26) +LOGSSFR_BINS_PDF = np.linspace(-13.0, -8.0, 30) + + +Z_BINS = [0.0, 0.5, 1.0, 1.5, 2.0] + +T0_SMDPL = 13.7976158 + +NCHUNKS = 20 + +N_HALOS_MAX = 20_000 +N_HALOS_PER_SUBVOL = N_HALOS_MAX // NCHUNKS + + +def load_tng_chunk_data(subvol, indir=BEBOP_TNG): + + fn = os.path.join(indir, "tng_diffmah.npy") + halos = np.load(fn) + upid = halos["cen1_sat0"] + + nhalos_tot = len(upid) + + _a = np.arange(0, nhalos_tot).astype("i8") + indx = np.array_split(_a, NCHUNKS)[subvol] + + return upid[indx] + + +def _load_flat_hdf5(fn): + data = dict() + with h5py.File(fn, "r") as hdf: + for key in hdf.keys(): + data[key] = hdf[key][...] + return data + + +def load_diffdata( + diffmah_drn=BEBOP_TNG_MAH, + diffstar_drn=BEBOP_TNG_SFH, +): + diffstar_bn = "diffstar_tng_fits.hdf5" + diffstar_fn = os.path.join(diffstar_drn, diffstar_bn) + diffstar_data = _load_flat_hdf5(diffstar_fn) + + diffmah_bn = diffstar_bn.replace("diffstar", "diffmah") + diffmah_fn = os.path.join(diffmah_drn, diffmah_bn) + diffmah_data = _load_flat_hdf5(diffmah_fn) + + return diffmah_data, diffstar_data + + +def load_diffstar_sfh_tables( + subvol, + diffmah_drn=BEBOP_TNG_MAH, + diffstar_drn=BEBOP_TNG_SFH, + lgt0=LGT0, + n_times=200, +): + diffmah_data, diffstar_data = load_diffdata( + diffmah_drn=diffmah_drn, + diffstar_drn=diffstar_drn, + ) + nhalos_tot = len(diffmah_data["logm0"][...]) + _a = np.arange(0, nhalos_tot).astype("i8") + indx = np.array_split(_a, NCHUNKS)[subvol] + + has_fit = ( + (diffmah_data["loss"][indx] > 0.0) + & (diffstar_data["loss"][indx] > 0.0) + & (diffstar_data["success"][indx] == 1) + ) + + mah_params = DEFAULT_MAH_PARAMS._make( + [diffmah_data[key][indx][has_fit] for key in DEFAULT_MAH_PARAMS._fields] + ) + + ms_params = DEFAULT_DIFFSTAR_PARAMS.ms_params._make( + [ + diffstar_data[key][indx][has_fit] + for key in DEFAULT_DIFFSTAR_PARAMS.ms_params._fields + ] + ) + q_params = DEFAULT_DIFFSTAR_PARAMS.q_params._make( + [ + diffstar_data[key][indx][has_fit] + for key in DEFAULT_DIFFSTAR_PARAMS.q_params._fields + ] + ) + sfh_params = DEFAULT_DIFFSTAR_PARAMS._make((ms_params, q_params)) + + t_0 = 10**lgt0 + t_table = np.linspace(T_TABLE_MIN, t_0, n_times) + + __, log_mah_table = mah_halopop(mah_params, t_table, LGT0) + + sfh_table, smh_table = calc_sfh_galpop( + sfh_params, mah_params, t_table, lgt0=LGT0, return_smh=True + ) + log_sfh_table = np.log10(sfh_table) + log_smh_table = np.log10(smh_table) + log_ssfrh_table = log_sfh_table - log_smh_table + + out = ( + t_table, + log_mah_table, + log_smh_table, + log_ssfrh_table, + mah_params, + ms_params, + q_params, + has_fit, + ) + + return out + + +def compute_weighted_histograms_z0( + subvol, + diffmah_drn=BEBOP_TNG_MAH, + diffstar_drn=BEBOP_TNG_SFH, + lgt0=LGT0, + logmh_bins=LOGMH_BINS, +): + _res = load_diffstar_sfh_tables( + subvol, + diffmah_drn=diffmah_drn, + diffstar_drn=diffstar_drn, + lgt0=lgt0, + ) + t_table, log_mah_table, log_smh_table, log_ssfrh_table = _res[:4] + + n_halos = log_smh_table.shape[0] + + nddata = log_mah_table[:, -1].reshape((-1, 1)) + + sigma = np.mean(np.diff(logmh_bins)) + np.zeros(n_halos) + ndsig = sigma.reshape((-1, 1)) + + ydata = log_smh_table[:, -1].reshape((-1, 1)) + _ones = np.ones_like(ydata) + + ndbins_lo = logmh_bins[:-1].reshape((-1, 1)) + ndbins_hi = logmh_bins[1:].reshape((-1, 1)) + + whist = tw_ndhist_weighted(nddata, ndsig, ydata, ndbins_lo, ndbins_hi) + wcounts = tw_ndhist_weighted(nddata, ndsig, _ones, ndbins_lo, ndbins_hi) + + return wcounts, whist + + +def compute_diff_histograms_atz(logmh_bins, log_mah_table, log_smh_table): + + n_halos = log_smh_table.shape[0] + + nddata = log_mah_table.reshape((-1, 1)) + + sigma = np.mean(np.diff(logmh_bins)) + np.zeros(n_halos) + ndsig = sigma.reshape((-1, 1)) + + ydata = log_smh_table.reshape((-1, 1)) + _ones = np.ones_like(ydata) + + ndbins_lo = logmh_bins[:-1].reshape((-1, 1)) + ndbins_hi = logmh_bins[1:].reshape((-1, 1)) + + whist = tw_ndhist_weighted(nddata, ndsig, ydata, ndbins_lo, ndbins_hi) + wcounts = tw_ndhist_weighted(nddata, ndsig, _ones, ndbins_lo, ndbins_hi) + + return wcounts, whist + + +def compute_histograms_atz(logmh_bins, log_mah_table, log_smh_table): + count = binned_statistic( + log_mah_table, values=log_smh_table, bins=logmh_bins, statistic="count" + )[0] + whist = binned_statistic( + log_mah_table, values=log_smh_table, bins=logmh_bins, statistic="sum" + )[0] + return count, whist + + +def get_redshift_from_age(age): + z_table = np.linspace(0, 10, 2000)[::-1] + t_table = Planck13.age(z_table).value + redshift_from_age = np.interp(age, t_table, z_table) + return redshift_from_age + + +def return_target_redshfit_index(t_table, redshift_targets): + z_table = get_redshift_from_age(t_table) + return np.digitize(redshift_targets, z_table) + + +def sample_halos( + logmh_bins, + log_mah, + log_smh, + mah_params, + ms_params, + q_params, + upid, +): + ndbins_lo = logmh_bins[:-1] + ndbins_hi = logmh_bins[1:] + arange_arr = np.arange(len(log_mah)) + logmh_id = [] + logmh_val = [] + mah_params_samp = [] + ms_params_samp = [] + q_params_samp = [] + upid_samp = [] + + mah_params = np.array(mah_params).T + ms_params = np.array(ms_params).T + q_params = np.array(q_params).T + + for i in range(len(ndbins_lo)): + sel = (log_mah >= ndbins_lo[i]) & (log_mah < ndbins_hi[i]) + sel_num = int(min(N_HALOS_PER_SUBVOL, sel.sum())) + sel = np.random.choice(arange_arr[sel], sel_num, replace=False) + logmh_id.append(np.ones_like(sel) * i) + logmh_val.append(np.ones_like(sel) * ((ndbins_lo[i] + ndbins_hi[i]) / 2.0)) + mah_params_samp.append(mah_params[sel]) + ms_params_samp.append(ms_params[sel]) + q_params_samp.append(q_params[sel]) + upid_samp.append(upid[sel]) + + logmh_id = np.concatenate(logmh_id) + logmh_val = np.concatenate(logmh_val) + mah_params_samp = np.concatenate(mah_params_samp) + ms_params_samp = np.concatenate(ms_params_samp) + q_params_samp = np.concatenate(q_params_samp) + upid_samp = np.concatenate(upid_samp) + + mah_params_samp = DEFAULT_MAH_PARAMS._make(mah_params_samp.T) + ms_params_samp = DEFAULT_DIFFSTAR_PARAMS.ms_params._make(ms_params_samp.T) + q_params_samp = DEFAULT_DIFFSTAR_PARAMS.q_params._make(q_params_samp.T) + out = ( + logmh_id, + logmh_val, + mah_params_samp, + ms_params_samp, + q_params_samp, + upid_samp, + ) + return out + + +def create_target_data( + subvol, + redshift_targets=Z_BINS, + diffmah_drn=BEBOP_TNG_MAH, + diffstar_drn=BEBOP_TNG_SFH, + lgt0=LGT0, + logmh_bins=LOGMH_BINS, +): + _res = load_diffstar_sfh_tables( + subvol, + diffmah_drn=diffmah_drn, + diffstar_drn=diffstar_drn, + lgt0=lgt0, + ) + ( + t_table, + log_mah_table, + log_smh_table, + log_ssfrh_table, + mah_params, + ms_params, + q_params, + has_fit, + ) = _res + + upid = load_tng_chunk_data(subvol)[has_fit] + tng_t = np.load(os.path.join(BEBOP_TNG, "tng_cosmic_time.npy")) + + tids = return_target_redshfit_index(t_table, redshift_targets) + tids_tng = return_target_redshfit_index(tng_t, redshift_targets) + + nz, nm = len(redshift_targets), len(logmh_bins) - 1 + + wcounts_zid = np.zeros((nz, nm)) + whist_zid = np.zeros((nz, nm)) + counts_zid = np.zeros((nz, nm)) + hist_zid = np.zeros((nz, nm)) + + counts_zid_cen = np.zeros((nz, nm)) + counts_zid_sat = np.zeros((nz, nm)) + + is_central = upid[:, -1] == 1 + is_satell = upid[:, -1] == 0 + + for i, tid in enumerate(tids): + _res = compute_diff_histograms_atz( + logmh_bins, log_mah_table[:, tid], log_smh_table[:, tid] + ) + wcounts_zid[i] = _res[0] + whist_zid[i] = _res[1] + + _res = compute_histograms_atz( + logmh_bins, log_mah_table[:, tid], log_smh_table[:, tid] + ) + counts_zid[i] = _res[0] + hist_zid[i] = _res[1] + + # is_central = upid[:, tids_tng[i]] == 1 + counts_zid_cen[i] = compute_histograms_atz( + logmh_bins, + log_mah_table[:, tid][is_central], + log_smh_table[:, tid][is_central], + )[0] + # is_satell = upid[:, tids_tng[i]] == 0 + counts_zid_sat[i] = compute_histograms_atz( + logmh_bins, + log_mah_table[:, tid][is_satell], + log_smh_table[:, tid][is_satell], + )[0] + + data = [] + + final_upid = upid[:, -1].copy() + final_upid[is_central] = -1 + + for i, tid in enumerate(tids): + _res = sample_halos( + logmh_bins, + log_mah_table[:, tid], + log_smh_table[:, tid], + mah_params, + ms_params, + q_params, + final_upid, + ) + data.append( + ( + *_res, + np.ones_like(_res[0]) * i, + np.ones_like(_res[0]) * t_table[tid], + np.ones_like(_res[0]) * redshift_targets[i], + ) + ) + + haloes = concatenate_samples_haloes(data) + + out = ( + wcounts_zid, + whist_zid, + counts_zid, + hist_zid, + t_table[tids], + haloes, + counts_zid_cen, + counts_zid_sat, + ) + + return out + + +def concatenate_samples_haloes(data): + logmh_id = [] + logmh_val = [] + redshift_val = [] + tobs_id = [] + tobs_val = [] + mah_params_samp = [] + ms_params_samp = [] + q_params_samp = [] + upid_samp = [] + + for subdata in data: + + logmh_id.append(subdata[0]) + logmh_val.append(subdata[1]) + mah_params_samp.append(np.array(subdata[2]).T) + ms_params_samp.append(np.array(subdata[3]).T) + q_params_samp.append(np.array(subdata[4]).T) + upid_samp.append(subdata[5]) + tobs_id.append(subdata[6]) + tobs_val.append(subdata[7]) + redshift_val.append(subdata[8]) + + logmh_id = np.concatenate(logmh_id) + logmh_val = np.concatenate(logmh_val) + mah_params_samp = np.concatenate(mah_params_samp) + ms_params_samp = np.concatenate(ms_params_samp) + q_params_samp = np.concatenate(q_params_samp) + upid_samp = np.concatenate(upid_samp) + tobs_id = np.concatenate(tobs_id) + tobs_val = np.concatenate(tobs_val) + redshift_val = np.concatenate(redshift_val) + + mah_params_samp = DEFAULT_MAH_PARAMS._make(mah_params_samp.T) + ms_params_samp = DEFAULT_DIFFSTAR_PARAMS.ms_params._make(ms_params_samp.T) + q_params_samp = DEFAULT_DIFFSTAR_PARAMS.q_params._make(q_params_samp.T) + + haloes = ( + logmh_id, + logmh_val, + mah_params_samp, + ms_params_samp, + q_params_samp, + upid_samp, + tobs_id, + tobs_val, + redshift_val, + ) + return haloes + + +def compute_diff_histograms_mstar_atmobs_z( + logmstar_bins, + log_smh_table, +): + + n_halos = log_smh_table.shape[0] + + nddata = log_smh_table.reshape((-1, 1)) + + sigma = np.mean(np.diff(logmstar_bins)) + np.zeros(n_halos) + ndsig = sigma.reshape((-1, 1)) + + ydata = log_smh_table.reshape((-1, 1)) + _ones = np.ones_like(ydata) + + ndbins_lo = logmstar_bins[:-1].reshape((-1, 1)) + ndbins_hi = logmstar_bins[1:].reshape((-1, 1)) + + wcounts = tw_ndhist_weighted(nddata, ndsig, _ones, ndbins_lo, ndbins_hi) + + counts = np.histogram(ydata, logmstar_bins)[0] + + return wcounts, counts + + +def compute_diff_histograms_mstar_ssfr_atz( + log_smh_table, + log_ssfr_table, + ndbins_lo, + ndbins_hi, + logmstar_bins, + logssfr_bins, +): + n_halos = log_smh_table.shape[0] + + sigma_mstar = np.mean(np.diff(logmstar_bins)) + np.zeros(n_halos) + sigma_ssfr = np.mean(np.diff(logssfr_bins)) + np.zeros(n_halos) + + ndsig = np.ones((n_halos, 2)) + ndsig[:, 0] = sigma_mstar + ndsig[:, 1] = sigma_ssfr + + nddata = np.array([log_smh_table, log_ssfr_table]).T + + _ones = np.ones(n_halos) + + wcounts = tw_ndhist_weighted(nddata, ndsig, _ones, ndbins_lo, ndbins_hi) + + return wcounts + + +def create_pdf_target_data( + subvol, + redshift_targets=Z_BINS, + diffmah_drn=BEBOP_TNG_MAH, + diffstar_drn=BEBOP_TNG_SFH, + lgt0=LGT0, + logmh_bins=LOGMH_BINS, + logmstar_bins_pdf=LOGMSTAR_BINS_PDF, + logssfr_bins_pdf=LOGSSFR_BINS_PDF, +): + _res = load_diffstar_sfh_tables( + subvol, + diffmah_drn=diffmah_drn, + diffstar_drn=diffstar_drn, + lgt0=lgt0, + ) + ( + t_table, + log_mah_table, + log_smh_table, + log_ssfrh_table, + mah_params, + ms_params, + q_params, + has_fit, + ) = _res + + log_ssfrh_table = np.clip(log_ssfrh_table, -12.0, None) + + upid = load_tng_chunk_data(subvol)[has_fit] + + tng_t = np.load(os.path.join(BEBOP_TNG, "tng_cosmic_time.npy")) + + tids = return_target_redshfit_index(t_table, redshift_targets) + tids_tng = return_target_redshfit_index(tng_t, redshift_targets) + + nz, nm = len(redshift_targets), len(logmh_bins) - 1 + nmstar = len(logmstar_bins_pdf) - 1 + nssfr = len(logssfr_bins_pdf) - 1 + + ndbins_lo = [] + ndbins_hi = [] + for i in range(len(logmstar_bins_pdf) - 1): + for j in range(len(logssfr_bins_pdf) - 1): + ndbins_lo.append([logmstar_bins_pdf[i], logssfr_bins_pdf[j]]) + ndbins_hi.append([logmstar_bins_pdf[i + 1], logssfr_bins_pdf[j + 1]]) + ndbins_lo = np.array(ndbins_lo) + ndbins_hi = np.array(ndbins_hi) + + mstar_wcounts = np.zeros((nz, nm, nmstar)) + mstar_counts = np.zeros((nz, nm, nmstar)) + + mstar_ssfr_wcounts_cent = np.zeros((nz, nm, nmstar, nssfr)) + mstar_ssfr_wcounts_sat = np.zeros((nz, nm, nmstar, nssfr)) + + is_central = upid[:, -1] == 1 + is_satell = upid[:, -1] == 0 + for i, tid in enumerate(tids): + # is_central = upid[:, tids_tng[i]] == 1 + # is_satell = upid[:, tids_tng[i]] == 0 + + for j in range(nm): + mobs_sel = (log_mah_table[:, tid] > logmh_bins[j]) & ( + log_mah_table[:, tid] < logmh_bins[j + 1] + ) + _res = compute_diff_histograms_mstar_atmobs_z( + logmstar_bins_pdf, + log_smh_table[mobs_sel][:, tid], + ) + mstar_wcounts[i, j] = _res[0] + mstar_counts[i, j] = _res[1] + + mobs_sel_cent = mobs_sel & is_central + _res = compute_diff_histograms_mstar_ssfr_atz( + log_smh_table[mobs_sel_cent][:, tid], + log_ssfrh_table[mobs_sel_cent][:, tid], + ndbins_lo, + ndbins_hi, + logmstar_bins_pdf, + logssfr_bins_pdf, + ) + mstar_ssfr_wcounts_cent[i, j] = _res.reshape((nmstar, nssfr)) + + mobs_sel_sat = mobs_sel & is_satell + _res = compute_diff_histograms_mstar_ssfr_atz( + log_smh_table[mobs_sel_sat][:, tid], + log_ssfrh_table[mobs_sel_sat][:, tid], + ndbins_lo, + ndbins_hi, + logmstar_bins_pdf, + logssfr_bins_pdf, + ) + mstar_ssfr_wcounts_sat[i, j] = _res.reshape((nmstar, nssfr)) + + out = ( + mstar_wcounts, + mstar_counts, + mstar_ssfr_wcounts_cent, + mstar_ssfr_wcounts_sat, + ) + + return out