diff --git a/.github/workflows/linting.yaml b/.github/workflows/linting.yaml index 16ec46d..6fe30b6 100644 --- a/.github/workflows/linting.yaml +++ b/.github/workflows/linting.yaml @@ -1,4 +1,4 @@ -name: linting +name: Flake8 test of source code on: push: @@ -8,7 +8,7 @@ on: jobs: tests: - name: lintest + name: flake8 diffstar runs-on: "ubuntu-latest" steps: diff --git a/.github/workflows/monthly-warning-test.yaml b/.github/workflows/monthly-warning-test.yaml index a371090..32dc135 100644 --- a/.github/workflows/monthly-warning-test.yaml +++ b/.github/workflows/monthly-warning-test.yaml @@ -1,4 +1,4 @@ -name: Test for Warnings +name: Monthly test for warnings on: workflow_dispatch: null @@ -8,7 +8,7 @@ on: jobs: tests: - name: tests + name: pytest with diffmah/dsps/diffsky@main runs-on: "ubuntu-latest" steps: diff --git a/.github/workflows/test_releases.yaml b/.github/workflows/test_releases.yaml index 05b2de9..366d870 100644 --- a/.github/workflows/test_releases.yaml +++ b/.github/workflows/test_releases.yaml @@ -1,4 +1,4 @@ -name: tests +name: Test against latest diffstuff releases on: workflow_dispatch: null @@ -9,7 +9,7 @@ on: jobs: tests: - name: tests + name: pytest with latest releases on conda-forge runs-on: "ubuntu-latest" steps: diff --git a/.github/workflows/tests_cron.yaml b/.github/workflows/tests_cron.yaml index 5bed46d..9e6f451 100644 --- a/.github/workflows/tests_cron.yaml +++ b/.github/workflows/tests_cron.yaml @@ -1,4 +1,4 @@ -name: test_main_branch_dependencies +name: Weekly cron testing on: workflow_dispatch: null @@ -12,7 +12,7 @@ on: jobs: tests: - name: tests + name: pytest with diffmah/dsps/diffsky@main runs-on: "ubuntu-latest" steps: diff --git a/CHANGES.rst b/CHANGES.rst index 5bb2eaf..72b96c3 100644 --- a/CHANGES.rst +++ b/CHANGES.rst @@ -1,3 +1,8 @@ +1.0.2 (unreleased) +------------------ +- Monte Carlo SFH generators now have required kwargs lgt0 and fb (https://github.com/ArgonneCPAC/diffstar/pull/108) + + 1.0.1 (2025-11-02) ------------------ - Update scaling relations and recalibrate default parameters (https://github.com/ArgonneCPAC/diffstar/pull/106) diff --git a/diffstar/diffstarpop/loss_kernels/mstar_ssfr_loss_mgash_anyz.py b/diffstar/diffstarpop/loss_kernels/mstar_ssfr_loss_mgash_anyz.py index a576376..971d5bf 100644 --- a/diffstar/diffstarpop/loss_kernels/mstar_ssfr_loss_mgash_anyz.py +++ b/diffstar/diffstarpop/loss_kernels/mstar_ssfr_loss_mgash_anyz.py @@ -1,18 +1,18 @@ """ """ from diffsky.diffndhist import tw_ndhist_weighted -from diffstar.utils import cumulative_mstar_formed from jax import jit as jjit from jax import numpy as jnp from jax import value_and_grad, vmap +from diffstar.utils import cumulative_mstar_formed + from ..kernels.defaults_mgash import ( DEFAULT_DIFFSTARPOP_U_PARAMS, get_bounded_diffstarpop_params, ) from ..mc_diffstarpop_mgash import mc_diffstar_sfh_galpop - N_TIMES = 20 _A = (None, 0) @@ -51,6 +51,8 @@ def _mc_diffstar_sfh_galpop_vmap_kern( gyr_since_infall, ran_key, tobs_target, + lgt0, + fb, ): tarr = jnp.logspace(-1, jnp.log10(tobs_target), N_TIMES) res = mc_diffstar_sfh_galpop( @@ -63,11 +65,13 @@ def _mc_diffstar_sfh_galpop_vmap_kern( gyr_since_infall, ran_key, tarr, + lgt0=lgt0, + fb=fb, ) return res -_U = (None, *[0] * 8) +_U = (None, *[0] * 8, None, None) mc_diffstar_sfh_galpop_vmap = jjit(vmap(_mc_diffstar_sfh_galpop_vmap_kern, in_axes=_U)) @@ -117,6 +121,8 @@ def mstar_kern_tobs(u_params, loss_data): gyr_since_infall, ran_key, tobs_target, + lgt0, + fb, logmstar_bins, target_mstar_pdf, ) = loss_data @@ -133,6 +139,8 @@ def mstar_kern_tobs(u_params, loss_data): gyr_since_infall, ran_key, tobs_target, + lgt0, + fb, ) diffstar_params_ms, diffstar_params_q, sfh_ms, sfh_q, frac_q, mc_is_q = _res @@ -233,6 +241,8 @@ def mstar_ssfr_kern_tobs(u_params, loss_data): gyr_since_infall, ran_key, tobs_target, + lgt0, + fb, ndbins_lo, ndbins_hi, logmstar_bins, @@ -255,6 +265,8 @@ def mstar_ssfr_kern_tobs(u_params, loss_data): gyr_since_infall, ran_key, tobs_target, + lgt0, + fb, ) diffstar_params_ms, diffstar_params_q, sfh_ms, sfh_q, frac_q, mc_is_q = _res @@ -347,6 +359,8 @@ def mstar_ssfr_sat_kern_tobs(u_params, loss_data): gyr_since_infall, ran_key, tobs_target, + lgt0, + fb, ndbins_lo, ndbins_hi, logmstar_bins, @@ -369,6 +383,8 @@ def mstar_ssfr_sat_kern_tobs(u_params, loss_data): gyr_since_infall, ran_key, tobs_target, + lgt0, + fb, ) diffstar_params_ms, diffstar_params_q, sfh_ms, sfh_q, frac_q, mc_is_q = _res diff --git a/diffstar/diffstarpop/loss_kernels/tests/test_mstar_ssfr_loss_mgash_anyz.py b/diffstar/diffstarpop/loss_kernels/tests/test_mstar_ssfr_loss_mgash_anyz.py index 3aa8923..8fc0fb7 100644 --- a/diffstar/diffstarpop/loss_kernels/tests/test_mstar_ssfr_loss_mgash_anyz.py +++ b/diffstar/diffstarpop/loss_kernels/tests/test_mstar_ssfr_loss_mgash_anyz.py @@ -91,6 +91,8 @@ def test_h5_data_shapes_and_sanity(loss_data_mstar, loss_data_ssfr, loss_data_ss gyr_since_infall_data, ran_key_data, t_obs_targets, + lgt0, + fb, logmstar_bins_pdf, mstar_counts_target, ) = loss_data_mstar @@ -122,6 +124,8 @@ def test_h5_data_shapes_and_sanity(loss_data_mstar, loss_data_ssfr, loss_data_ss _, _, _, + _, + _, ndbins_lo, ndbins_hi, logmstar_bins_pdf2, @@ -148,6 +152,8 @@ def test_h5_data_shapes_and_sanity(loss_data_mstar, loss_data_ssfr, loss_data_ss _, _, _, + _, + _, ndbins_lo_s, ndbins_hi_s, logmstar_bins_pdf_s, diff --git a/diffstar/diffstarpop/loss_kernels/tests/testing_data/loss_kernels_testing_data_10halos.h5 b/diffstar/diffstarpop/loss_kernels/tests/testing_data/loss_kernels_testing_data_10halos.h5 index b375b97..47b080e 100644 Binary files a/diffstar/diffstarpop/loss_kernels/tests/testing_data/loss_kernels_testing_data_10halos.h5 and b/diffstar/diffstarpop/loss_kernels/tests/testing_data/loss_kernels_testing_data_10halos.h5 differ diff --git a/diffstar/diffstarpop/mc_diffstarpop_mgash.py b/diffstar/diffstarpop/mc_diffstarpop_mgash.py index 23e63d7..db7fd6d 100644 --- a/diffstar/diffstarpop/mc_diffstarpop_mgash.py +++ b/diffstar/diffstarpop/mc_diffstarpop_mgash.py @@ -6,7 +6,7 @@ from jax import random as jran from jax import vmap -from ..defaults import FB, LGT0, get_bounded_diffstar_params +from ..defaults import get_bounded_diffstar_params from ..sfh_model import calc_sfh_galpop, calc_sfh_singlegal from .kernels.diffstarpop_mgash import mc_diffstar_u_params_singlegal_kernel @@ -35,8 +35,9 @@ def mc_diffstar_sfh_singlegal( gyr_since_infall, ran_key, tarr, - lgt0=LGT0, - fb=FB, + *, + lgt0, + fb, ): """Monte Carlo realization of a single point in Diffstar parameter space, along with the computation of SFH for this point. @@ -376,8 +377,9 @@ def mc_diffstar_sfh_galpop( gyr_since_infall, ran_key, tarr, - lgt0=LGT0, - fb=FB, + *, + lgt0, + fb, ): """Monte Carlo realization of a single point in Diffstar parameter space, along with the computation of SFH for this point. diff --git a/diffstar/diffstarpop/tests/test_gradients.py b/diffstar/diffstarpop/tests/test_gradients.py index b09e728..c62b841 100644 --- a/diffstar/diffstarpop/tests/test_gradients.py +++ b/diffstar/diffstarpop/tests/test_gradients.py @@ -105,7 +105,7 @@ def test_all_diffstarpop_u_param_gradients_are_nonzero(): default_sfh_q, frac_q, mc_is_q, - ) = mc_diffstar_sfh_galpop(*args) + ) = mc_diffstar_sfh_galpop(*args, lgt0=1.14, fb=0.156) assert default_sfh_q.shape == (n_halos, ntimes) assert np.all(np.isfinite(default_sfh_q)) @@ -140,7 +140,7 @@ def test_all_diffstarpop_u_param_gradients_are_nonzero(): alt_sfh_q, alt_frac_q, mc_is_q, - ) = mc_diffstar_sfh_galpop(*args) + ) = mc_diffstar_sfh_galpop(*args, lgt0=1.14, fb=0.156) assert alt_sfh_q.shape == (n_halos, ntimes) assert np.all(np.isfinite(alt_sfh_q)) @@ -167,7 +167,7 @@ def _loss(u_params): pred_sfh_q, pred_frac_q, mc_is_q, - ) = mc_diffstar_sfh_galpop(*args) + ) = mc_diffstar_sfh_galpop(*args, lgt0=1.14, fb=0.156) pred_mean_sfh_total = jnp.mean( pred_frac_q[:, None] * pred_sfh_q + (1.0 - pred_frac_q[:, None]) * pred_sfh_ms, diff --git a/diffstar/diffstarpop/tests/test_mc_diffstarpop_mgash.py b/diffstar/diffstarpop/tests/test_mc_diffstarpop_mgash.py index eccbc68..4b46c52 100644 --- a/diffstar/diffstarpop/tests/test_mc_diffstarpop_mgash.py +++ b/diffstar/diffstarpop/tests/test_mc_diffstarpop_mgash.py @@ -55,7 +55,7 @@ def test_mc_diffstar_sfh_singlegal_evaluates(): ran_key, tarr, ) - _res = mcdsp.mc_diffstar_sfh_singlegal(*args) + _res = mcdsp.mc_diffstar_sfh_singlegal(*args, lgt0=1.14, fb=0.156) params_ms, params_q, sfh_ms, sfh_q, frac_q, mc_is_q = _res assert np.all(frac_q >= 0) assert np.all(frac_q <= 1) @@ -147,6 +147,8 @@ def test_mc_diffstar_sfh_galpop(): gyr_since_infall, ran_key, t_table, + lgt0=1.14, + fb=0.156, ) sfh_q, sfh_ms, frac_q = _res[2:5] diff --git a/diffstar/sfh_model.py b/diffstar/sfh_model.py index 0cfdf71..d35ee12 100644 --- a/diffstar/sfh_model.py +++ b/diffstar/sfh_model.py @@ -7,7 +7,6 @@ from jax import numpy as jnp from jax import vmap -from .defaults import FB, LGT0 from .kernels.history_kernel_builders import _sfh_galpop_kern, _sfh_singlegal_kern from .utils import cumulative_mstar_formed @@ -18,7 +17,13 @@ @partial(jjit, static_argnames="return_smh") def calc_sfh_singlegal( - sfh_params, mah_params, tarr, lgt0=LGT0, fb=FB, return_smh=False + sfh_params, + mah_params, + tarr, + *, + lgt0, + fb, + return_smh=False, ): """Calculate the Diffstar SFH for a single galaxy @@ -72,7 +77,15 @@ def calc_sfh_singlegal( @partial(jjit, static_argnames="return_smh") -def calc_sfh_galpop(sfh_params, mah_params, tarr, lgt0=LGT0, fb=FB, return_smh=False): +def calc_sfh_galpop( + sfh_params, + mah_params, + tarr, + *, + lgt0, + fb, + return_smh=False, +): """Calculate the Diffstar SFH for a single galaxy Parameters diff --git a/docs/source/demo_diffmahpop_diffstarpop_sfh.ipynb b/docs/source/demo_diffmahpop_diffstarpop_sfh.ipynb index 0b95dbc..0ec8969 100644 --- a/docs/source/demo_diffmahpop_diffstarpop_sfh.ipynb +++ b/docs/source/demo_diffmahpop_diffstarpop_sfh.ipynb @@ -83,7 +83,6 @@ "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", @@ -91,7 +90,8 @@ "\n", "# Some constants\n", "from dsps.constants import T_TABLE_MIN\n", - "from diffstar.defaults import LGT0, TODAY\n" + "from diffstar.defaults import LGT0 as DEFAULT_LGT0\n", + "TODAY = 10**DEFAULT_LGT0" ] }, { @@ -202,13 +202,15 @@ "source": [ "from diffmah.diffmah_kernels import mah_halopop\n", "from diffstar.diffstarpop import mc_diffstar_sfh_galpop\n", + "from diffstar.defaults import FB as DEFAULT_FB\n", + "\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", + "dmhdt_fit, log_mah_fit = mah_halopop(subcat.mah_params, tarr, DEFAULT_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", @@ -239,7 +241,7 @@ " default_sfh_q,\n", " frac_q,\n", " mc_is_q,\n", - ") = mc_diffstar_sfh_galpop(*args)\n", + ") = mc_diffstar_sfh_galpop(*args, lgt0=DEFAULT_LGT0, fb=DEFAULT_FB)\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", @@ -350,7 +352,8 @@ " 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", + " ax[0].fill_between(tarr, range_log_mah_fit[0], range_log_mah_fit[1], \n", + " 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", @@ -373,19 +376,23 @@ " default_sfh_q,\n", " frac_q,\n", " mc_is_q,\n", - " ) = mc_diffstar_sfh_galpop(*args)\n", + " ) = mc_diffstar_sfh_galpop(*args, lgt0=DEFAULT_LGT0, fb=DEFAULT_FB)\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", + " sel = (subcat.logmp0 > mpeak - 0.2) \n", + " sel = sel & (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", + " range_mean_default_sfh = np.percentile(\n", + " 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", + " ax[k+1].fill_between(\n", + " tarr, range_mean_default_sfh[0], range_mean_default_sfh[1], \n", + " color=colors[i], alpha=0.1)\n", "\n", "\n", "ax[0].set_ylim(9, 14)\n", @@ -428,7 +435,7 @@ ], "metadata": { "kernelspec": { - "display_name": "diffstuff", + "display_name": "Python 3 (ipykernel)", "language": "python", "name": "python3" }, @@ -442,7 +449,7 @@ "name": "python", "nbconvert_exporter": "python", "pygments_lexer": "ipython3", - "version": "3.11.9" + "version": "3.12.9" } }, "nbformat": 4, diff --git a/docs/source/demo_diffstar_sfh.ipynb b/docs/source/demo_diffstar_sfh.ipynb index 416b8f9..0d078d0 100644 --- a/docs/source/demo_diffstar_sfh.ipynb +++ b/docs/source/demo_diffstar_sfh.ipynb @@ -60,7 +60,9 @@ "from diffmah.defaults import DEFAULT_MAH_PARAMS\n", "from diffstar.defaults import DEFAULT_DIFFSTAR_PARAMS\n", "\n", - "today_gyr = 13.8 \n", + "today_gyr = 13.8\n", + "LGT0 = np.log10(today_gyr)\n", + "DEFAULT_FB = 0.156\n", "tarr = np.linspace(0.9, today_gyr, 100)" ] }, @@ -74,7 +76,8 @@ "from diffstar import calc_sfh_singlegal\n", "\n", "sfh_gal = calc_sfh_singlegal(\n", - " DEFAULT_DIFFSTAR_PARAMS, DEFAULT_MAH_PARAMS, tarr)" + " DEFAULT_DIFFSTAR_PARAMS, DEFAULT_MAH_PARAMS, tarr, \n", + " lgt0=LGT0, fb=DEFAULT_FB)" ] }, { @@ -163,6 +166,7 @@ " gyr_since_infall,\n", " sfh_key,\n", " tarr,\n", + " lgt0=LGT0, fb=DEFAULT_FB\n", ")\n", "\n", "print(mc_diffstar_result._fields)" @@ -183,7 +187,10 @@ "metadata": {}, "outputs": [], "source": [ - "sfh = np.where(mc_diffstar_result.mc_is_q.reshape((n_halos, 1)), mc_diffstar_result.sfh_q, mc_diffstar_result.sfh_ms)" + "sfh = np.where(\n", + " mc_diffstar_result.mc_is_q.reshape((n_halos, 1)), \n", + " mc_diffstar_result.sfh_q, mc_diffstar_result.sfh_ms\n", + ")" ] }, { @@ -205,11 +212,19 @@ "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": "de5b95cd-96e8-4082-ba86-6c1f2a55c123", + "metadata": {}, + "outputs": [], + "source": [] } ], "metadata": { "kernelspec": { - "display_name": "diffstuff", + "display_name": "Python 3 (ipykernel)", "language": "python", "name": "python3" }, @@ -223,7 +238,7 @@ "name": "python", "nbconvert_exporter": "python", "pygments_lexer": "ipython3", - "version": "3.11.9" + "version": "3.12.9" } }, "nbformat": 4, diff --git a/scripts/diffstarpop_scripts/fit_get_loss_helpers_mgash.py b/scripts/diffstarpop_scripts/fit_get_loss_helpers_mgash.py index d8c8bb8..8699aa9 100644 --- a/scripts/diffstarpop_scripts/fit_get_loss_helpers_mgash.py +++ b/scripts/diffstarpop_scripts/fit_get_loss_helpers_mgash.py @@ -4,12 +4,12 @@ import h5py import numpy as np from diffmah.diffmah_kernels import DiffmahParams, mah_halopop -from diffstar.defaults import LGT0 +from diffstar.defaults import LGT0, FB from jax import random as jran from jax import numpy as jnp -def get_loss_data_smhm(indir, nhalos): +def get_loss_data_smhm(indir, nhalos, lgt0=LGT0, fb=FB): # Load SMHM data --------------------------------------------- print("Loading SMHM data...") @@ -61,7 +61,7 @@ def get_loss_data_smhm(indir, nhalos): t_obs_targets = [] smhm_targets = [] - tarr_logm0 = np.logspace(-1, LGT0, 50) + tarr_logm0 = np.logspace(-1, lgt0, 50) for i in range(len(age_targets)): t_target = age_targets[i] @@ -81,7 +81,7 @@ def get_loss_data_smhm(indir, nhalos): 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) + 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) @@ -124,7 +124,7 @@ def get_loss_data_smhm(indir, nhalos): return loss_data, plot_data -def get_loss_data_pdfs_mstar(indir, nhalos): +def get_loss_data_pdfs_mstar(indir, nhalos, lgt0=LGT0, fb=FB): # Load PDF data --------------------------------------------- print("Loading PDF Mstar data...") @@ -195,7 +195,7 @@ def get_loss_data_pdfs_mstar(indir, nhalos): t_obs_targets = [] mstar_counts_target = [] - tarr_logm0 = np.logspace(-1, LGT0, 50) + tarr_logm0 = np.logspace(-1, lgt0, 50) for i in range(len(age_targets)): t_target = age_targets[i] @@ -216,7 +216,7 @@ def get_loss_data_pdfs_mstar(indir, nhalos): 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) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, lgt0) logmp0_data.append(log_mah_fit[:, -1]) # break @@ -239,6 +239,8 @@ def get_loss_data_pdfs_mstar(indir, nhalos): gyr_since_infall_data, ran_key_data, t_obs_targets, + lgt0, + fb, logmstar_bins_pdf, mstar_counts_target, ) @@ -261,7 +263,7 @@ def get_loss_data_pdfs_mstar(indir, nhalos): t_obs_targets = [] mstar_counts_target = [] - tarr_logm0 = np.logspace(-1, LGT0, 50) + tarr_logm0 = np.logspace(-1, lgt0, 50) for i in range(len(age_targets)): t_target = age_targets[i] @@ -282,7 +284,7 @@ def get_loss_data_pdfs_mstar(indir, nhalos): 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) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, lgt0) logmp0_data.append(log_mah_fit[:, -1]) # break @@ -305,6 +307,8 @@ def get_loss_data_pdfs_mstar(indir, nhalos): gyr_since_infall_data, ran_key_data, t_obs_targets, + lgt0, + fb, logmstar_bins_pdf, mstar_counts_target, ) @@ -344,7 +348,7 @@ def prepare_ragged(indx_pdf, nmhalo_pdf, index_mhalo): return idx, w # shapes: (nz, Mmax), (nz, Mmax) -def get_loss_data_pdfs_ssfr_central(indir, nhalos): +def get_loss_data_pdfs_ssfr_central(indir, nhalos, lgt0=LGT0, fb=FB): print("Loading PDF Mstar data...") @@ -431,7 +435,7 @@ def get_loss_data_pdfs_ssfr_central(indir, nhalos): "zmab,zm->zab", mstar_ssfr_pdfs_cent, mhalo_pdf_cen ) - tarr_logm0 = np.logspace(-1, LGT0, 50) + tarr_logm0 = np.logspace(-1, lgt0, 50) index_mhalo = [] indx_pdf = [] @@ -456,7 +460,7 @@ def get_loss_data_pdfs_ssfr_central(indir, nhalos): 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) + 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) @@ -514,6 +518,8 @@ def get_loss_data_pdfs_ssfr_central(indir, nhalos): gyr_since_infall_data, ran_key_data, t_obs_targets, + lgt0, + fb, ndbins_lo, ndbins_hi, logmstar_bins_pdf, @@ -577,7 +583,7 @@ def get_loss_data_pdfs_ssfr_central(indir, nhalos): 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) + 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) @@ -635,6 +641,8 @@ def get_loss_data_pdfs_ssfr_central(indir, nhalos): gyr_since_infall_data, ran_key_data, t_obs_targets, + lgt0, + fb, ndbins_lo, ndbins_hi, logmstar_bins_pdf, @@ -655,7 +663,7 @@ def get_loss_data_pdfs_ssfr_central(indir, nhalos): return loss_data_ssfr, plot_data -def get_loss_data_pdfs_ssfr_satellite(indir, nhalos): +def get_loss_data_pdfs_ssfr_satellite(indir, nhalos, lgt0=LGT0, fb=FB): print("Loading PDF Mstar data...") @@ -743,7 +751,7 @@ def get_loss_data_pdfs_ssfr_satellite(indir, nhalos): indx_pdf = [] _run_indx = 0 - tarr_logm0 = np.logspace(-1, LGT0, 50) + tarr_logm0 = np.logspace(-1, lgt0, 50) for i in range(len(age_targets)): t_target = age_targets[i] index_mhalo_atz = [] @@ -764,7 +772,7 @@ def get_loss_data_pdfs_ssfr_satellite(indir, nhalos): 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) + 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) @@ -823,6 +831,8 @@ def get_loss_data_pdfs_ssfr_satellite(indir, nhalos): gyr_since_infall_data, ran_key_data, t_obs_targets, + lgt0, + fb, ndbins_lo, ndbins_hi, logmstar_bins_pdf, @@ -884,7 +894,7 @@ def get_loss_data_pdfs_ssfr_satellite(indir, nhalos): 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) + 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) @@ -942,6 +952,8 @@ def get_loss_data_pdfs_ssfr_satellite(indir, nhalos): gyr_since_infall_data, ran_key_data, t_obs_targets, + lgt0, + fb, ndbins_lo, ndbins_hi, logmstar_bins_pdf, @@ -962,7 +974,7 @@ def get_loss_data_pdfs_ssfr_satellite(indir, nhalos): return loss_data_ssfr_sat, plot_data -def get_loss_data_pdfs_mstar_cen(indir, nhalos): +def get_loss_data_pdfs_mstar_cen(indir, nhalos, lgt0=LGT0, fb=FB): # Load PDF data --------------------------------------------- print("Loading PDF Mstar for centrals data...") @@ -1037,7 +1049,7 @@ def get_loss_data_pdfs_mstar_cen(indir, nhalos): t_obs_targets = [] mstar_counts_target = [] - tarr_logm0 = np.logspace(-1, LGT0, 50) + tarr_logm0 = np.logspace(-1, lgt0, 50) for i in range(len(age_targets)): t_target = age_targets[i] @@ -1060,7 +1072,7 @@ def get_loss_data_pdfs_mstar_cen(indir, nhalos): 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) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, lgt0) logmp0_data.append(log_mah_fit[:, -1]) # break @@ -1083,6 +1095,8 @@ def get_loss_data_pdfs_mstar_cen(indir, nhalos): gyr_since_infall_data, ran_key_data, t_obs_targets, + lgt0, + fb, logmstar_bins_pdf, mstar_counts_target, ) @@ -1105,7 +1119,7 @@ def get_loss_data_pdfs_mstar_cen(indir, nhalos): t_obs_targets = [] mstar_counts_target = [] - tarr_logm0 = np.logspace(-1, LGT0, 50) + tarr_logm0 = np.logspace(-1, lgt0, 50) for i in range(len(age_targets)): t_target = age_targets[i] @@ -1128,7 +1142,7 @@ def get_loss_data_pdfs_mstar_cen(indir, nhalos): 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) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, lgt0) logmp0_data.append(log_mah_fit[:, -1]) # break @@ -1151,6 +1165,8 @@ def get_loss_data_pdfs_mstar_cen(indir, nhalos): gyr_since_infall_data, ran_key_data, t_obs_targets, + lgt0, + fb, logmstar_bins_pdf, mstar_counts_target, ) @@ -1168,7 +1184,7 @@ def get_loss_data_pdfs_mstar_cen(indir, nhalos): return loss_data_mstar, plot_data -def get_loss_data_pdfs_mstar_sat(indir, nhalos): +def get_loss_data_pdfs_mstar_sat(indir, nhalos, lgt0=LGT0, fb=FB): # Load PDF data --------------------------------------------- print("Loading PDF Mstar for satellites data...") @@ -1243,7 +1259,7 @@ def get_loss_data_pdfs_mstar_sat(indir, nhalos): t_obs_targets = [] mstar_counts_target = [] - tarr_logm0 = np.logspace(-1, LGT0, 50) + tarr_logm0 = np.logspace(-1, lgt0, 50) for i in range(len(age_targets)): t_target = age_targets[i] @@ -1266,7 +1282,7 @@ def get_loss_data_pdfs_mstar_sat(indir, nhalos): 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) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, lgt0) logmp0_data.append(log_mah_fit[:, -1]) # break @@ -1289,6 +1305,8 @@ def get_loss_data_pdfs_mstar_sat(indir, nhalos): gyr_since_infall_data, ran_key_data, t_obs_targets, + lgt0, + fb, logmstar_bins_pdf, mstar_counts_target, ) @@ -1311,7 +1329,7 @@ def get_loss_data_pdfs_mstar_sat(indir, nhalos): t_obs_targets = [] mstar_counts_target = [] - tarr_logm0 = np.logspace(-1, LGT0, 50) + tarr_logm0 = np.logspace(-1, lgt0, 50) for i in range(len(age_targets)): t_target = age_targets[i] @@ -1334,7 +1352,7 @@ def get_loss_data_pdfs_mstar_sat(indir, nhalos): 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) + dmhdt_fit, log_mah_fit = mah_halopop(mah_pars_ntuple, tarr_logm0, lgt0) logmp0_data.append(log_mah_fit[:, -1]) # break @@ -1357,6 +1375,8 @@ def get_loss_data_pdfs_mstar_sat(indir, nhalos): gyr_since_infall_data, ran_key_data, t_obs_targets, + lgt0, + fb, logmstar_bins_pdf, mstar_counts_target, ) diff --git a/scripts/diffstarpop_scripts/save_loss_data.py b/scripts/diffstarpop_scripts/save_loss_data.py new file mode 100644 index 0000000..e0df86c --- /dev/null +++ b/scripts/diffstarpop_scripts/save_loss_data.py @@ -0,0 +1,319 @@ +""" +save_loss_data.py + +Usage (inside Python): + from save_loss_data import save_loss_data_h5 + save_loss_data_h5("unit_test_loss_data.h5", + loss_data_mstar, + loss_data_ssfr, + loss_data_ssfr_sat) + +This expects you already have the three tuples in memory, in the exact orders shown +in your message. +""" + +from typing import Any, Iterable, Sequence +import numpy as np +import h5py +import json + +from fit_get_loss_helpers_mgash import ( + get_loss_data_smhm, + get_loss_data_pdfs_mstar, + get_loss_data_pdfs_ssfr_central, + get_loss_data_pdfs_ssfr_satellite, +) + +# --- Names in tuple order (so we can label datasets clearly) --- + +MSTAR_FIELDS = [ + "mah_params_data", + "logmp0_data", + "upid_data", + "lgmu_infall_data", + "logmhost_infall_data", + "gyr_since_infall_data", + "ran_key_data", + "t_obs_targets", + "lgt0", + "fb", + "logmstar_bins_pdf", + "mstar_counts_target", +] + +SSFR_FIELDS = [ + "mah_params_data", + "logmp0_data", + "upid_data", + "lgmu_infall_data", + "logmhost_infall_data", + "gyr_since_infall_data", + "ran_key_data", + "t_obs_targets", + "lgt0", + "fb", + "ndbins_lo", + "ndbins_hi", + "logmstar_bins_pdf", + "logssfr_bins_pdf", + "mhalo_pdf_cen_ragged", # ragged allowed + "indx_pdf", + "target_mstar_ids", + "target_data", +] + +SSFR_SAT_FIELDS = [ + "mah_params_data", + "logmp0_data", + "upid_data", + "lgmu_infall_data", + "logmhost_infall_data", + "gyr_since_infall_data", + "ran_key_data", + "t_obs_targets", + "lgt0", + "fb", + "ndbins_lo", + "ndbins_hi", + "logmstar_bins_pdf", + "logssfr_bins_pdf", + "mhalo_pdf_sat_ragged", # ragged allowed + "indx_pdf", + "target_mstar_ids", + "target_data_sat", +] + + +def _to_numpy(x: Any): + """Convert JAX/array-like -> NumPy array without copying if possible.""" + if isinstance(x, np.ndarray): + return x + try: + return np.asarray(x) + except Exception: + return x # leave as-is (e.g., a list of arrays) + + +def _is_array_like(x: Any) -> bool: + return isinstance(x, (np.ndarray,)) or hasattr(x, "__array__") + + +def _is_sequence(x: Any) -> bool: + return isinstance(x, (list, tuple)) + + +def _is_ragged_sequence(seq: Sequence[Any]) -> bool: + """ + Heuristic: sequence of array-like where shapes are not all equal. + Also treat object-dtype arrays as ragged. + """ + if isinstance(seq, np.ndarray) and seq.dtype == object: + return True + if not _is_sequence(seq): + return False + shapes = [] + for el in seq: + if _is_array_like(el): + a = _to_numpy(el) + shapes.append(a.shape) + else: + # Non-array element -> treat as ragged to be safe + return True + return len(set(shapes)) > 1 + + +def _save_value(group: h5py.Group, name: str, value: Any): + """ + Save a single item under group/name. + - Regular ndarrays: one dataset + - Scalar: 0-D dataset + - Ragged sequences: a subgroup with datasets '0000', '0001', ... + - Sequence with equal shapes: stacked into one dataset + """ + # If object-dtype NumPy => likely ragged + if isinstance(value, np.ndarray) and value.dtype == object: + value = list(value) # treat as ragged sequence + + if _is_sequence(value) and _is_ragged_sequence(value): + # Save as subgroup with one dataset per element + sub = group.create_group(name) + for i, el in enumerate(value): + arr = _to_numpy(el) + sub.create_dataset(f"{i:04d}", data=arr) + sub.attrs["format"] = "ragged_list_of_datasets" + sub.attrs["length"] = len(value) + return + + # Non-ragged sequences of equal-shaped arrays -> stack + if _is_sequence(value): + # Convert to array if possible (will stack) + try: + arr = _to_numpy(value) + group.create_dataset(name, data=arr) + return + except Exception: + # Fallback: save as subgroup + sub = group.create_group(name) + for i, el in enumerate(value): + arr = _to_numpy(el) + sub.create_dataset(f"{i:04d}", data=arr) + sub.attrs["format"] = "list_of_datasets" + sub.attrs["length"] = len(value) + return + + # Scalar or array-like + if _is_array_like(value): + arr = _to_numpy(value) + group.create_dataset(name, data=arr) + return + + # Last resort: store JSON-serializable objects as attrs + try: + group.attrs[name] = json.dumps(value) + except Exception: + # If we end up here, user passed a very custom object + # Save a string repr so the test can still load something. + group.attrs[name] = repr(value) + + +def save_loss_data_h5( + filename: str, + loss_data_mstar: Iterable[Any], + loss_data_ssfr: Iterable[Any], + loss_data_ssfr_sat: Iterable[Any], +): + """ + Save three loss-data tuples into an HDF5 file with a clean hierarchy: + + /loss_data_mstar/... + /loss_data_ssfr/... + /loss_data_ssfr_sat/... + + Each dataset is named after the variable (e.g., 'mah_params_data'). + Ragged arrays/lists are stored as a subgroup with one dataset per element. + + Parameters + ---------- + filename : str + Output .h5 path. + loss_data_mstar, loss_data_ssfr, loss_data_ssfr_sat : tuple-like + Tuples exactly matching the field orders defined above. + """ + # Sanity checks on tuple lengths + if len(loss_data_mstar) != len(MSTAR_FIELDS): + raise ValueError( + f"loss_data_mstar length {len(loss_data_mstar)} != {len(MSTAR_FIELDS)}" + ) + if len(loss_data_ssfr) != len(SSFR_FIELDS): + raise ValueError( + f"loss_data_ssfr length {len(loss_data_ssfr)} != {len(SSFR_FIELDS)}" + ) + if len(loss_data_ssfr_sat) != len(SSFR_SAT_FIELDS): + raise ValueError( + f"loss_data_ssfr_sat length {len(loss_data_ssfr_sat)} != {len(SSFR_SAT_FIELDS)}" + ) + + with h5py.File(filename, "w") as f: + # mstar + g_m = f.create_group("loss_data_mstar") + g_m.attrs["field_order"] = json.dumps(MSTAR_FIELDS) + for name, val in zip(MSTAR_FIELDS, loss_data_mstar): + _save_value(g_m, name, val) + + # ssfr (centrals) + g_c = f.create_group("loss_data_ssfr") + g_c.attrs["field_order"] = json.dumps(SSFR_FIELDS) + for name, val in zip(SSFR_FIELDS, loss_data_ssfr): + _save_value(g_c, name, val) + + # ssfr (satellites) + g_s = f.create_group("loss_data_ssfr_sat") + g_s.attrs["field_order"] = json.dumps(SSFR_SAT_FIELDS) + for name, val in zip(SSFR_SAT_FIELDS, loss_data_ssfr_sat): + _save_value(g_s, name, val) + + # File-level note to help future you + f.attrs["description"] = ( + "Unit testing data for mstar/ssfr kernels. Ragged lists are stored as groups " + "with one dataset per element and 'format' attr." + ) + + print(f"Wrote {filename}") + + +# --- Optional: small loader helper for ragged groups (use in tests) --- + + +def load_loss_data_h5(filename: str): + """ + Load the three tuples back from disk. Ragged groups are returned as lists + of NumPy arrays. Returns (loss_data_mstar, loss_data_ssfr, loss_data_ssfr_sat). + """ + + def _load_group(g: h5py.Group, field_names): + out = [] + for name in field_names: + if name in g: + obj = g[name] + if isinstance(obj, h5py.Dataset): + out.append(obj[()]) + elif isinstance(obj, h5py.Group): + fmt = ( + obj.attrs.get("format", "").decode() + if isinstance(obj.attrs.get("format", ""), bytes) + else obj.attrs.get("format", "") + ) + if fmt in ("ragged_list_of_datasets", "list_of_datasets"): + items = [obj[k][()] for k in sorted(obj.keys())] + out.append(items) + else: + # Unknown layout—try datasets in key order + items = [obj[k][()] for k in sorted(obj.keys())] + out.append(items) + else: + raise RuntimeError(f"Unexpected HDF5 object at {g.name}/{name}") + else: + # Might be stored as attr (rare) + if name in g.attrs: + val = g.attrs[name] + try: + out.append(json.loads(val)) + except Exception: + out.append(val) + else: + raise KeyError(f"Field '{name}' not found in group {g.name}") + return tuple(out) + + with h5py.File(filename, "r") as f: + m_fields = json.loads(f["loss_data_mstar"].attrs["field_order"]) + c_fields = json.loads(f["loss_data_ssfr"].attrs["field_order"]) + s_fields = json.loads(f["loss_data_ssfr_sat"].attrs["field_order"]) + + mstar = _load_group(f["loss_data_mstar"], m_fields) + ssfr = _load_group(f["loss_data_ssfr"], c_fields) + ssfr_s = _load_group(f["loss_data_ssfr_sat"], s_fields) + + return mstar, ssfr, ssfr_s + + +indir = ( + "/Users/alarcon/Documents/diffmah_data/mgash/smdpl_dr1_nomerging_pdf_target_data/" +) +nhalos = 10 +loss_data_mstar, plot_data_pdf = get_loss_data_pdfs_mstar(indir, nhalos) +loss_data_ssfr, plot_data_pdf_ssfr_cen = get_loss_data_pdfs_ssfr_central(indir, nhalos) +loss_data_ssfr_sat, plot_data_pdf_ssfr_sat = get_loss_data_pdfs_ssfr_satellite( + indir, nhalos +) +fname = "loss_kernels_testing_data_10halos.h5" +save_loss_data_h5(fname, loss_data_mstar, loss_data_ssfr, loss_data_ssfr_sat) + + +nhalos = 100 +loss_data_mstar, plot_data_pdf = get_loss_data_pdfs_mstar(indir, nhalos) +loss_data_ssfr, plot_data_pdf_ssfr_cen = get_loss_data_pdfs_ssfr_central(indir, nhalos) +loss_data_ssfr_sat, plot_data_pdf_ssfr_sat = get_loss_data_pdfs_ssfr_satellite( + indir, nhalos +) +fname = "loss_kernels_testing_data_100halos.h5" +save_loss_data_h5(fname, loss_data_mstar, loss_data_ssfr, loss_data_ssfr_sat)