diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml index 5bc36fa..8b976a5 100644 --- a/.github/workflows/tests.yml +++ b/.github/workflows/tests.yml @@ -45,7 +45,12 @@ jobs: setuptools \ "setuptools_scm>=7,<8" \ scipy \ - python-build + python-build \ + jaxopt \ + tqdm \ + optax \ + astropy + pip install --no-deps git+https://github.com/AlanPearl/kdescent.git python -m pip install --no-build-isolation --no-deps -e . - name: test diff --git a/diffmah/diffmahpop_kernels/kdescent_testing/__init__.py b/diffmah/diffmahpop_kernels/kdescent_testing/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/diffmah/diffmahpop_kernels/kdescent_testing/dmp_wrappers.py b/diffmah/diffmahpop_kernels/kdescent_testing/dmp_wrappers.py new file mode 100644 index 0000000..9023a32 --- /dev/null +++ b/diffmah/diffmahpop_kernels/kdescent_testing/dmp_wrappers.py @@ -0,0 +1,109 @@ +""" +""" + +from jax import value_and_grad, vmap + +try: + import kdescent +except ImportError: + pass +from jax import jit as jjit +from jax import numpy as jnp + +from .. import diffmahpop_params as dpp +from .. import mc_diffmahpop_kernels as mdk + + +@jjit +def mc_diffmah_preds(diffmahpop_u_params, pred_data): + diffmahpop_params = dpp.get_diffmahpop_params_from_u_params(diffmahpop_u_params) + tarr, lgm_obs, t_obs, ran_key, lgt0 = pred_data + _res = mdk._mc_diffmah_halo_sample( + diffmahpop_params, tarr, lgm_obs, t_obs, ran_key, lgt0 + ) + ftpt0 = _res[3] + log_mah_tpt0 = _res[6] + log_mah_tp = _res[8] + return log_mah_tpt0, log_mah_tp, ftpt0 + + +@jjit +def single_sample_kde_loss_self_fit( + diffmahpop_u_params, + tarr, + lgm_obs, + t_obs, + ran_key, + lgt0, + log_mahs_target, + weights_target, +): + pred_data = tarr, lgm_obs, t_obs, ran_key, lgt0 + _res = mc_diffmah_preds(diffmahpop_u_params, pred_data) + log_mah_tpt0, log_mah_tp, ftpt0 = _res + log_mahs_pred = jnp.concatenate((log_mah_tpt0, log_mah_tp)) + weights_pred = jnp.concatenate((ftpt0, 1 - ftpt0)) + + kcalc = kdescent.KCalc(log_mahs_target, weights_target) + model_counts, truth_counts = kcalc.compare_kde_counts( + ran_key, log_mahs_pred, weights_pred + ) + diff = model_counts - truth_counts + loss = jnp.mean(diff**2) + return loss + + +single_sample_kde_loss_and_grad_self_fit = jjit( + value_and_grad(single_sample_kde_loss_self_fit) +) + + +@jjit +def get_single_sample_self_fit_target_data( + u_params, tarr, lgm_obs, t_obs, ran_key, lgt0 +): + pred_data = tarr, lgm_obs, t_obs, ran_key, lgt0 + _res = mc_diffmah_preds(u_params, pred_data) + log_mah_tpt0, log_mah_tp, ftpt0 = _res + log_mahs_target = jnp.concatenate((log_mah_tpt0, log_mah_tp)) + weights_target = jnp.concatenate((ftpt0, 1 - ftpt0)) + return log_mahs_target, weights_target + + +_A = (None, 0, 0, 0, 0, None) +get_multisample_self_fit_target_data = jjit( + vmap(get_single_sample_self_fit_target_data, in_axes=_A) +) + +_L = (None, 0, 0, 0, 0, None, 0, 0) +_multisample_kde_loss_self_fit = jjit(vmap(single_sample_kde_loss_self_fit, in_axes=_L)) + + +@jjit +def multisample_kde_loss_self_fit( + diffmahpop_u_params, + tarr_matrix, + lgmobsarr, + tobsarr, + ran_keys, + lgt0, + log_mahs_targets, + weights_targets, +): + losses = _multisample_kde_loss_self_fit( + diffmahpop_u_params, + tarr_matrix, + lgmobsarr, + tobsarr, + ran_keys, + lgt0, + log_mahs_targets, + weights_targets, + ) + loss = jnp.mean(losses) + return loss + + +multisample_kde_loss_and_grad_self_fit = jjit( + value_and_grad(multisample_kde_loss_self_fit) +) diff --git a/diffmah/diffmahpop_kernels/kdescent_testing/kde1d_wrappers.py b/diffmah/diffmahpop_kernels/kdescent_testing/kde1d_wrappers.py new file mode 100644 index 0000000..4ddcc1f --- /dev/null +++ b/diffmah/diffmahpop_kernels/kdescent_testing/kde1d_wrappers.py @@ -0,0 +1,146 @@ +""" +""" + +from jax import random as jran +from jax import value_and_grad, vmap + +try: + import kdescent +except ImportError: + pass +from jax import jit as jjit +from jax import numpy as jnp + +from .. import diffmahpop_params as dpp +from .. import mc_diffmahpop_kernels as mdk + +N_T_PER_BIN = 5 + + +@jjit +def mc_diffmah_preds(diffmahpop_u_params, pred_data): + diffmahpop_params = dpp.get_diffmahpop_params_from_u_params(diffmahpop_u_params) + tarr, lgm_obs, t_obs, ran_key, lgt0 = pred_data + _res = mdk._mc_diffmah_halo_sample( + diffmahpop_params, tarr, lgm_obs, t_obs, ran_key, lgt0 + ) + ftpt0 = _res[3] + log_mah_tpt0 = _res[6] + log_mah_tp = _res[8] + return log_mah_tpt0, log_mah_tp, ftpt0 + + +@jjit +def single_sample_kde_loss_self_fit( + diffmahpop_u_params, + tarr, + lgm_obs, + t_obs, + ran_key, + lgt0, + log_mahs_target, + weights_target, +): + s = (-1, 1) + kcalc0 = kdescent.KCalc(log_mahs_target[:, 0].reshape(s), weights_target) + kcalc1 = kdescent.KCalc(log_mahs_target[:, 1].reshape(s), weights_target) + kcalc2 = kdescent.KCalc(log_mahs_target[:, 2].reshape(s), weights_target) + kcalc3 = kdescent.KCalc(log_mahs_target[:, 3].reshape(s), weights_target) + kcalc4 = kdescent.KCalc(log_mahs_target[:, 4].reshape(s), weights_target) + + ran_key, pred_key = jran.split(ran_key, 2) + pred_data = tarr, lgm_obs, t_obs, pred_key, lgt0 + _res = mc_diffmah_preds(diffmahpop_u_params, pred_data) + log_mah_tpt0, log_mah_tp, ftpt0 = _res + + log_mahs_pred = jnp.concatenate((log_mah_tpt0, log_mah_tp)) + weights_pred = jnp.concatenate((ftpt0, 1 - ftpt0)) + + kcalc_keys = jran.split(ran_key, N_T_PER_BIN) + + model_counts0, truth_counts0 = kcalc0.compare_kde_counts( + kcalc_keys[0], log_mahs_pred[:, 0].reshape(s), weights_pred + ) + model_counts1, truth_counts1 = kcalc1.compare_kde_counts( + kcalc_keys[1], log_mahs_pred[:, 1].reshape(s), weights_pred + ) + model_counts2, truth_counts2 = kcalc2.compare_kde_counts( + kcalc_keys[2], log_mahs_pred[:, 2].reshape(s), weights_pred + ) + model_counts3, truth_counts3 = kcalc3.compare_kde_counts( + kcalc_keys[3], log_mahs_pred[:, 3].reshape(s), weights_pred + ) + model_counts4, truth_counts4 = kcalc4.compare_kde_counts( + kcalc_keys[4], log_mahs_pred[:, 4].reshape(s), weights_pred + ) + + diff0 = model_counts0 - truth_counts0 + diff1 = model_counts1 - truth_counts1 + diff2 = model_counts2 - truth_counts2 + diff3 = model_counts3 - truth_counts3 + diff4 = model_counts4 - truth_counts4 + + loss0 = jnp.mean(diff0**2) + loss1 = jnp.mean(diff1**2) + loss2 = jnp.mean(diff2**2) + loss3 = jnp.mean(diff3**2) + loss4 = jnp.mean(diff4**2) + + loss = loss0 + loss1 + loss2 + loss3 + loss4 + return loss + + +single_sample_kde_loss_and_grad_self_fit = jjit( + value_and_grad(single_sample_kde_loss_self_fit) +) + + +@jjit +def get_single_sample_self_fit_target_data( + u_params, tarr, lgm_obs, t_obs, ran_key, lgt0 +): + pred_data = tarr, lgm_obs, t_obs, ran_key, lgt0 + _res = mc_diffmah_preds(u_params, pred_data) + log_mah_tpt0, log_mah_tp, ftpt0 = _res + log_mahs_target = jnp.concatenate((log_mah_tpt0, log_mah_tp)) + weights_target = jnp.concatenate((ftpt0, 1 - ftpt0)) + return log_mahs_target, weights_target + + +_A = (None, 0, 0, 0, 0, None) +get_multisample_self_fit_target_data = jjit( + vmap(get_single_sample_self_fit_target_data, in_axes=_A) +) + +_L = (None, 0, 0, 0, 0, None, 0, 0) +_multisample_kde_loss_self_fit = jjit(vmap(single_sample_kde_loss_self_fit, in_axes=_L)) + + +@jjit +def multisample_kde_loss_self_fit( + diffmahpop_u_params, + tarr_matrix, + lgmobsarr, + tobsarr, + ran_keys, + lgt0, + log_mahs_targets, + weights_targets, +): + losses = _multisample_kde_loss_self_fit( + diffmahpop_u_params, + tarr_matrix, + lgmobsarr, + tobsarr, + ran_keys, + lgt0, + log_mahs_targets, + weights_targets, + ) + loss = jnp.mean(losses) + return loss + + +multisample_kde_loss_and_grad_self_fit = jjit( + value_and_grad(multisample_kde_loss_self_fit) +) diff --git a/diffmah/diffmahpop_kernels/kdescent_testing/kde2d_wrappers.py b/diffmah/diffmahpop_kernels/kdescent_testing/kde2d_wrappers.py new file mode 100644 index 0000000..9ecaab1 --- /dev/null +++ b/diffmah/diffmahpop_kernels/kdescent_testing/kde2d_wrappers.py @@ -0,0 +1,386 @@ +""" +""" + +import os +from glob import glob + +import numpy as np +from astropy.cosmology import Planck15 +from astropy.table import Table +from jax import random as jran +from jax import value_and_grad, vmap + +from ...diffmah_kernels import mah_halopop + +try: + import kdescent +except ImportError: + pass +from jax import jit as jjit +from jax import numpy as jnp + +from ...diffmah_kernels import DEFAULT_MAH_PARAMS +from .. import diffmahpop_params as dpp +from .. import mc_diffmahpop_kernels as mdk + +N_T_PER_BIN = 5 +LGSMAH_MIN = -15 +EPS = 1e-3 + + +@jjit +def mc_diffmah_preds(diffmahpop_u_params, pred_data): + diffmahpop_params = dpp.get_diffmahpop_params_from_u_params(diffmahpop_u_params) + tarr, lgm_obs, t_obs, ran_key, lgt0 = pred_data + _res = mdk._mc_diffmah_halo_sample( + diffmahpop_params, tarr, lgm_obs, t_obs, ran_key, lgt0 + ) + ftpt0 = _res[3] + dmhdt_tpt0 = _res[5] + log_mah_tpt0 = _res[6] + dmhdt_tp = _res[7] + log_mah_tp = _res[8] + + dmhdt_tpt0 = jnp.clip(dmhdt_tpt0, 10**LGSMAH_MIN) # make log-safe + dmhdt_tp = jnp.clip(dmhdt_tp, 10**LGSMAH_MIN) # make log-safe + + lgsmah_tpt0 = jnp.log10(dmhdt_tpt0) - log_mah_tpt0 # compute lgsmah + lgsmah_tpt0 = jnp.clip(lgsmah_tpt0, LGSMAH_MIN) # impose lgsMAH clip + + lgsmah_tp = jnp.log10(dmhdt_tp) - log_mah_tpt0 # compute lgsmah + lgsmah_tp = jnp.clip(lgsmah_tp, LGSMAH_MIN) # impose lgsMAH clip + + return lgsmah_tpt0, log_mah_tpt0, lgsmah_tp, log_mah_tp, ftpt0 + + +@jjit +def get_single_sample_self_fit_target_data( + u_params, tarr, lgm_obs, t_obs, ran_key, lgt0 +): + pred_data = tarr, lgm_obs, t_obs, ran_key, lgt0 + _res = mc_diffmah_preds(u_params, pred_data) + lgsmah_tpt0, log_mah_tpt0, lgsmah_tp, log_mah_tp, ftpt0 = _res + weights_target = jnp.concatenate((ftpt0, 1 - ftpt0)) + lgsmah_target = jnp.concatenate((lgsmah_tpt0, lgsmah_tp)) + log_mahs_target = jnp.concatenate((log_mah_tpt0, log_mah_tp)) + X_target = jnp.array((lgsmah_target, log_mahs_target)).swapaxes(0, 1) + return X_target, weights_target + + +@jjit +def get_single_cen_sample_target_data(mah_params, t_peak, tarr, lgm_obs, lgt0): + dmhdt, log_mah = mah_halopop(mah_params, tarr, t_peak, lgt0) + + # renormalize MAHs to zero to at lgm_obs + delta_log_mahs_target = log_mah - lgm_obs + + frac_peaked_target = jnp.mean(dmhdt == 0, axis=0) + weights_target = jnp.where(dmhdt == 0, 0.0, 1.0) + + log_dmhdt = jnp.log10(jnp.clip(dmhdt, 10**LGSMAH_MIN)) + + lgsmah_target = log_dmhdt - log_mah # use log_mah since dmhdt was never rescaled + lgsmah_target = jnp.clip(lgsmah_target, LGSMAH_MIN) + + X_target = jnp.array((lgsmah_target, delta_log_mahs_target)).swapaxes(0, 1) + + return X_target, weights_target, frac_peaked_target + + +@jjit +def single_sample_kde_loss_self_fit( + diffmahpop_u_params, + tarr, + lgm_obs, + t_obs, + ran_key, + lgt0, + X_target, + weights_target, +): + kcalc0 = kdescent.KCalc(X_target[:, :, 0], weights_target) + kcalc1 = kdescent.KCalc(X_target[:, :, 1], weights_target) + kcalc2 = kdescent.KCalc(X_target[:, :, 2], weights_target) + kcalc3 = kdescent.KCalc(X_target[:, :, 3], weights_target) + kcalc4 = kdescent.KCalc(X_target[:, :, 4], weights_target) + + ran_key, pred_key = jran.split(ran_key, 2) + pred_data = tarr, lgm_obs, t_obs, pred_key, lgt0 + _res = mc_diffmah_preds(diffmahpop_u_params, pred_data) + dmhdt_tpt0, log_mah_tpt0, dmhdt_tp, log_mah_tp, ftpt0 = _res + + weights_pred = jnp.concatenate((ftpt0, 1 - ftpt0)) + dmhdts_pred = jnp.concatenate((dmhdt_tpt0, dmhdt_tp)) + log_mahs_pred = jnp.concatenate((log_mah_tpt0, log_mah_tp)) + X_preds = jnp.array((dmhdts_pred, log_mahs_pred)).swapaxes(0, 1) + + kcalc_keys = jran.split(ran_key, N_T_PER_BIN) + + model_counts0, truth_counts0 = kcalc0.compare_kde_counts( + kcalc_keys[0], X_preds[:, :, 0], weights_pred + ) + model_counts1, truth_counts1 = kcalc1.compare_kde_counts( + kcalc_keys[1], X_preds[:, :, 1], weights_pred + ) + model_counts2, truth_counts2 = kcalc2.compare_kde_counts( + kcalc_keys[2], X_preds[:, :, 2], weights_pred + ) + model_counts3, truth_counts3 = kcalc3.compare_kde_counts( + kcalc_keys[3], X_preds[:, :, 3], weights_pred + ) + model_counts4, truth_counts4 = kcalc4.compare_kde_counts( + kcalc_keys[4], X_preds[:, :, 4], weights_pred + ) + + diff0 = model_counts0 - truth_counts0 + diff1 = model_counts1 - truth_counts1 + diff2 = model_counts2 - truth_counts2 + diff3 = model_counts3 - truth_counts3 + diff4 = model_counts4 - truth_counts4 + + loss0 = jnp.mean(diff0**2) + loss1 = jnp.mean(diff1**2) + loss2 = jnp.mean(diff2**2) + loss3 = jnp.mean(diff3**2) + loss4 = jnp.mean(diff4**2) + + loss = loss0 + loss1 + loss2 + loss3 + loss4 + return loss + + +@jjit +def single_sample_kde_loss_kern( + diffmahpop_u_params, + tarr, + lgm_obs, + t_obs, + ran_key, + lgt0, + X_target, + weights_target, + frac_peaked_target, +): + kcalc0 = kdescent.KCalc(X_target[:, :, 0], weights_target[:, 0]) + kcalc1 = kdescent.KCalc(X_target[:, :, 1], weights_target[:, 1]) + kcalc2 = kdescent.KCalc(X_target[:, :, 2], weights_target[:, 2]) + kcalc3 = kdescent.KCalc(X_target[:, :, 3], weights_target[:, 3]) + + lgsmah_target_t_obs = X_target[:, 0, 4].reshape((-1, 1)) + kcalc_t_obs = kdescent.KCalc(lgsmah_target_t_obs) + + ran_key, pred_key = jran.split(ran_key, 2) + pred_data = tarr, lgm_obs, t_obs, pred_key, lgt0 + _res = mc_diffmah_preds(diffmahpop_u_params, pred_data) + lgsmah_tpt0, log_mah_tpt0, lgsmah_tp, log_mah_tp, ftpt0 = _res + + weights_ftpt0 = jnp.concatenate((ftpt0, 1 - ftpt0)) + lgsmah_pred = jnp.concatenate((lgsmah_tpt0, lgsmah_tp)) + log_mahs_pred = jnp.concatenate((log_mah_tpt0, log_mah_tp)) + delta_log_mahs_pred = log_mahs_pred - lgm_obs + + frac_peaked_pred = jnp.average( + lgsmah_pred == LGSMAH_MIN, axis=0, weights=weights_ftpt0 + ) + + weights_tp = jnp.where(lgsmah_tp == LGSMAH_MIN, 0.0, 1.0) + weights_tpt0 = jnp.ones_like(weights_tp) + weights = jnp.concatenate((weights_tpt0, weights_tpt0)) + weights = weights * weights_ftpt0.reshape((-1, 1)) + + X_preds = jnp.array((lgsmah_pred, delta_log_mahs_pred)).swapaxes(0, 1) + + kcalc_keys = jran.split(ran_key, N_T_PER_BIN) + + model_counts0, truth_counts0 = kcalc0.compare_kde_counts( + kcalc_keys[0], X_preds[:, :, 0], weights[:, 0] + ) + model_counts1, truth_counts1 = kcalc1.compare_kde_counts( + kcalc_keys[1], X_preds[:, :, 1], weights[:, 1] + ) + model_counts2, truth_counts2 = kcalc2.compare_kde_counts( + kcalc_keys[2], X_preds[:, :, 2], weights[:, 2] + ) + model_counts3, truth_counts3 = kcalc3.compare_kde_counts( + kcalc_keys[3], X_preds[:, :, 3], weights[:, 3] + ) + + lgsmah_pred_t_obs = lgsmah_pred[:, -1].reshape((-1, 1)) + model_counts4, truth_counts4 = kcalc_t_obs.compare_kde_counts( + kcalc_keys[4], lgsmah_pred_t_obs, weights[:, 4] + ) + + delta_lgm_obs = delta_log_mahs_pred[:, -1] + + diff0 = model_counts0 - truth_counts0 + diff1 = model_counts1 - truth_counts1 + diff2 = model_counts2 - truth_counts2 + diff3 = model_counts3 - truth_counts3 + diff4 = model_counts4 - truth_counts4 + + fracdiff0 = diff0 / truth_counts0 + fracdiff1 = diff1 / truth_counts1 + fracdiff2 = diff2 / truth_counts2 + fracdiff3 = diff3 / truth_counts3 + fracdiff4 = diff4 / truth_counts4 + + loss0 = jnp.mean(jnp.abs(fracdiff0)) + loss1 = jnp.mean(jnp.abs(fracdiff1)) + loss2 = jnp.mean(jnp.abs(fracdiff2)) + loss3 = jnp.mean(jnp.abs(fracdiff3)) + loss4 = jnp.mean(jnp.abs(fracdiff4)) + # loss1 = jnp.mean(fracdiff1**2) + # loss2 = jnp.mean(fracdiff2**2) + # loss3 = jnp.mean(fracdiff3**2) + # loss4 = jnp.mean(fracdiff4**2) + + # loss_lgm_obs = jnp.mean(delta_lgm_obs**2) + loss_lgm_obs = jnp.mean(jnp.abs(delta_lgm_obs)) + + frac_peaked_diff = frac_peaked_pred - frac_peaked_target + # loss_frac_peaked = jnp.mean(frac_peaked_diff**2) + loss_frac_peaked = jnp.mean(jnp.abs(frac_peaked_diff)) + + loss = loss0 + loss1 + loss2 + loss3 + loss4 + loss_frac_peaked + loss_lgm_obs + # return (loss0, loss1, loss2, loss3, loss4, loss_frac_peaked, loss_lgm_obs) + return loss + + +single_sample_kde_loss_and_grad_kern = jjit(value_and_grad(single_sample_kde_loss_kern)) + +single_sample_kde_loss_and_grad_self_fit = jjit( + value_and_grad(single_sample_kde_loss_self_fit) +) + +_A = (None, 0, 0, 0, 0, None) +get_multisample_self_fit_target_data = jjit( + vmap(get_single_sample_self_fit_target_data, in_axes=_A) +) + +_L = (None, 0, 0, 0, 0, None, 0, 0) +_multisample_kde_loss_self_fit = jjit(vmap(single_sample_kde_loss_self_fit, in_axes=_L)) + +_L2 = (None, 0, 0, 0, 0, None, 0, 0, 0) +_multi_sample_kde_loss_kern = jjit(vmap(single_sample_kde_loss_kern, in_axes=_L2)) + + +@jjit +def multi_sample_kde_loss_kern( + diffmahpop_u_params, + tarr_matrix, + lgm_obs_arr, + t_obs_arr, + ran_keys, + lgt0, + X_targets, + weights_targets, + frac_peaked_targets, +): + losses = _multi_sample_kde_loss_kern( + diffmahpop_u_params, + tarr_matrix, + lgm_obs_arr, + t_obs_arr, + ran_keys, + lgt0, + X_targets, + weights_targets, + frac_peaked_targets, + ) + loss = jnp.mean(losses) + return loss + + +multi_sample_kde_loss_and_grad_kern = jjit(value_and_grad(multi_sample_kde_loss_kern)) + + +@jjit +def multisample_kde_loss_self_fit( + diffmahpop_u_params, + tarr_matrix, + lgmobsarr, + tobsarr, + ran_keys, + lgt0, + X_targets, + weights_targets, +): + losses = _multisample_kde_loss_self_fit( + diffmahpop_u_params, + tarr_matrix, + lgmobsarr, + tobsarr, + ran_keys, + lgt0, + X_targets, + weights_targets, + ) + loss = jnp.mean(losses) + return loss + + +multisample_kde_loss_and_grad_self_fit = jjit( + value_and_grad(multisample_kde_loss_self_fit) +) + + +def get_cens_target_data(drn, ran_key, lgt0): + cen_target_fnames = sorted(glob(os.path.join(drn, "*cen_mah*.h5"))) + cen_bnames = [os.path.basename(fn) for fn in cen_target_fnames] + + cen_scale_factors = np.array([float(bn.split("_")[-1][:-3]) for bn in cen_bnames]) + cen_redshifts = 1 / cen_scale_factors - 1.0 + + cen_t_obs = Planck15.age(cen_redshifts).value + + t_obs_collector = [] + lgm_obs_collector = [] + tarr_collector = [] + X_target_collector = [] + weights_target_collector = [] + frac_peaked_target_collector = [] + for it_obs, t_obs in enumerate(cen_t_obs): + cens = Table.read(cen_target_fnames[it_obs], path="data") + + lgm_obs_arr = np.sort(np.unique(cens["lgm_obs"])) + mah_keys = ("logm0", "logtc", "early_index", "late_index") + mah_params = DEFAULT_MAH_PARAMS._make([cens[key] for key in mah_keys]) + + for im_obs, lgm_obs in enumerate(lgm_obs_arr): + mmsk = cens["lgm_obs"] == lgm_obs + mah_params_target = DEFAULT_MAH_PARAMS._make([x[mmsk] for x in mah_params]) + t_peak_target = cens["t_peak"][mmsk] + + tarr = np.linspace(0.5, t_obs - EPS, N_T_PER_BIN) + + _res = get_single_cen_sample_target_data( + mah_params_target, t_peak_target, tarr, lgm_obs, lgt0 + ) + X_target, weights_target, frac_peaked_target = _res + tarr_collector.append(tarr) + lgm_obs_collector.append(lgm_obs) + t_obs_collector.append(t_obs) + X_target_collector.append(X_target) + weights_target_collector.append(weights_target) + frac_peaked_target_collector.append(frac_peaked_target) + + n_samples = len(t_obs_collector) + + tarr_collector = jnp.array(tarr_collector) + lgm_obs_collector = jnp.array(lgm_obs_collector) + t_obs_collector = jnp.array(t_obs_collector) + ran_keys = jran.split(ran_key, n_samples) + X_target_collector = jnp.array(X_target_collector) + weights_target_collector = jnp.array(weights_target_collector) + frac_peaked_target_collector = jnp.array(frac_peaked_target_collector) + + loss_data = ( + tarr_collector, + lgm_obs_collector, + t_obs_collector, + ran_keys, + lgt0, + X_target_collector, + weights_target_collector, + frac_peaked_target_collector, + ) + return loss_data diff --git a/diffmah/diffmahpop_kernels/kdescent_testing/tests/__init__.py b/diffmah/diffmahpop_kernels/kdescent_testing/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/diffmah/diffmahpop_kernels/kdescent_testing/tests/test_kde1d_wrappers.py b/diffmah/diffmahpop_kernels/kdescent_testing/tests/test_kde1d_wrappers.py new file mode 100644 index 0000000..15df730 --- /dev/null +++ b/diffmah/diffmahpop_kernels/kdescent_testing/tests/test_kde1d_wrappers.py @@ -0,0 +1,192 @@ +""" +""" + +import numpy as np +import pytest +from jax import jit as jjit +from jax import numpy as jnp +from jax import random as jran + +from ... import diffmahpop_params as dpp +from .. import kde1d_wrappers as k1w + +try: + import kdescent # noqa + + HAS_KDESCENT = True +except ImportError: + HAS_KDESCENT = False + +T_MIN_FIT = 0.5 + + +def test_mc_diffmah_preds(): + ran_key = jran.key(0) + t_0 = 13.8 + lgt0 = np.log10(t_0) + n_times = 5 + + n_tests = 20 + for __ in range(n_tests): + ran_key, m_key, t_key = jran.split(ran_key, 3) + lgm_obs = jran.uniform(m_key, minval=9, maxval=16, shape=()) + t_obs = jran.uniform(t_key, minval=3, maxval=t_0, shape=()) + tarr = np.linspace(T_MIN_FIT, t_obs, n_times) + pred_data = tarr, lgm_obs, t_obs, ran_key, lgt0 + _preds = k1w.mc_diffmah_preds(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS, pred_data) + for _x in _preds: + assert np.all(np.isfinite(_x)) + log_mah_tpt0, log_mah_tp, ftpt0 = _preds + assert np.all(log_mah_tpt0 < 20) + assert np.all(log_mah_tp < 20) + assert np.all(ftpt0 <= 1) + assert np.all(ftpt0 >= 0) + + +@pytest.mark.skipif("not HAS_KDESCENT") +def test_single_sample_kde_loss_self_fit(): + """Enforce that single-sample loss has finite grads""" + ran_key = jran.key(0) + + t_0 = 13.8 + lgt0 = np.log10(t_0) + DP = 0.5 + + n_tests = 100 + for __ in range(n_tests): + + # Use random diffmahpop parameter to generate fiducial data + u_p_fid_key, ran_key = jran.split(ran_key, 2) + n_params = len(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + uran = jran.uniform(u_p_fid_key, minval=-DP, maxval=DP, shape=(n_params,)) + _u_p_list = [x + u for x, u in zip(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS, uran)] + u_p_fid = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(jnp.array(_u_p_list)) + + ran_key, lgm_key, t_obs_key = jran.split(ran_key, 3) + lgm_obs = jran.uniform(lgm_key, minval=11, maxval=15, shape=()) + t_obs = jran.uniform(t_obs_key, minval=4, maxval=t_0, shape=()) + + tarr = np.linspace(T_MIN_FIT, t_obs, k1w.N_T_PER_BIN) + + _res = k1w.get_single_sample_self_fit_target_data( + u_p_fid, tarr, lgm_obs, t_obs, ran_key, lgt0 + ) + log_mahs_target, weights_target = _res + + # use default params as the initial guess + u_p_init = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make( + dpp.DEFAULT_DIFFMAHPOP_U_PARAMS + ) + + loss_data = tarr, lgm_obs, t_obs, ran_key, lgt0, log_mahs_target, weights_target + loss = k1w.single_sample_kde_loss_self_fit(u_p_init, *loss_data) + assert np.all(np.isfinite(loss)), (lgm_obs, t_obs) + assert loss > 0 + + loss, grads = k1w.single_sample_kde_loss_and_grad_self_fit(u_p_init, *loss_data) + assert np.all(np.isfinite(grads)), (lgm_obs, t_obs) + + +@pytest.mark.skipif("not HAS_KDESCENT") +def test_multisample_kde_loss_self_fit(): + """Enforce that multi-sample loss has finite grads""" + ran_key = jran.key(0) + + # Use random diffmahpop parameter to generate fiducial data + DP = 0.1 + u_p_fid_key, ran_key = jran.split(ran_key, 2) + n_params = len(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + uran = jran.uniform(u_p_fid_key, minval=-DP, maxval=DP, shape=(n_params,)) + _u_p_list = [x + u for x, u in zip(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS, uran)] + u_p_fid = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(jnp.array(_u_p_list)) + + lgmobs_key, tobs_key, ran_key = jran.split(ran_key, 3) + num_samples = 5 + t_0 = 13.8 + lgt0 = np.log10(t_0) + lgmobsarr = jran.uniform(lgmobs_key, minval=11, maxval=15, shape=(num_samples,)) + tobsarr = jran.uniform(tobs_key, minval=4, maxval=13, shape=(num_samples,)) + num_target_redshifts_per_t_obs = 10 + + tarr_matrix = jnp.array( + [jnp.linspace(T_MIN_FIT, t, num_target_redshifts_per_t_obs) for t in tobsarr] + ) + _keys = jran.split(ran_key, num_samples * 2) + _res = k1w.get_multisample_self_fit_target_data( + u_p_fid, tarr_matrix, lgmobsarr, tobsarr, _keys[:num_samples], lgt0 + ) + log_mahs_targets, weights_targets = _res + loss_data = ( + tarr_matrix, + lgmobsarr, + tobsarr, + _keys[num_samples:], + lgt0, + log_mahs_targets, + weights_targets, + ) + # use default params as the initial guess + u_p_init = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + loss = k1w.multisample_kde_loss_self_fit(u_p_init, *loss_data) + assert np.all(np.isfinite(loss)) + assert loss > 0 + + loss, grads = k1w.multisample_kde_loss_and_grad_self_fit(u_p_init, *loss_data) + assert np.all(np.isfinite(grads)) + + +@pytest.mark.skipif("not HAS_KDESCENT") +def test_kdescent_adam_self_fit(): + """Enforce that kdescent.adam terminates without NaNs""" + ran_key = jran.key(0) + + # Use random diffmahpop parameter to generate fiducial data + DP = 0.5 + u_p_fid_key, ran_key = jran.split(ran_key, 2) + n_params = len(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + uran = jran.uniform(u_p_fid_key, minval=-DP, maxval=DP, shape=(n_params,)) + _u_p_list = [x + u for x, u in zip(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS, uran)] + u_p_fid = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(jnp.array(_u_p_list)) + + lgmobs_key, tobs_key, ran_key = jran.split(ran_key, 3) + num_samples = 5 + t_0 = 13.8 + lgt0 = np.log10(t_0) + lgmobsarr = jran.uniform(lgmobs_key, minval=11, maxval=15, shape=(num_samples,)) + tobsarr = jran.uniform(tobs_key, minval=4, maxval=13, shape=(num_samples,)) + num_target_redshifts_per_t_obs = 10 + + tarr_matrix = jnp.array( + [jnp.linspace(T_MIN_FIT, t, num_target_redshifts_per_t_obs) for t in tobsarr] + ) + + @jjit + def kde_loss(u_p, randkey): + _keys = jran.split(randkey, num_samples * 2) + _res = k1w.get_multisample_self_fit_target_data( + u_p_fid, tarr_matrix, lgmobsarr, tobsarr, _keys[:num_samples], lgt0 + ) + log_mahs_targets, weights_targets = _res + loss_data = ( + tarr_matrix, + lgmobsarr, + tobsarr, + _keys[num_samples:], + lgt0, + log_mahs_targets, + weights_targets, + ) + return k1w.multisample_kde_loss_self_fit(u_p, *loss_data) + + # use default params as the initial guess + u_p_init = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + + adam_results = kdescent.adam( + kde_loss, + u_p_init, + nsteps=5, + learning_rate=0.1, + randkey=12345, + ) + u_p_best = adam_results[-1] + assert np.all(np.isfinite(u_p_best)) diff --git a/diffmah/diffmahpop_kernels/kdescent_testing/tests/test_kde2d_wrappers.py b/diffmah/diffmahpop_kernels/kdescent_testing/tests/test_kde2d_wrappers.py new file mode 100644 index 0000000..69b53a2 --- /dev/null +++ b/diffmah/diffmahpop_kernels/kdescent_testing/tests/test_kde2d_wrappers.py @@ -0,0 +1,327 @@ +""" +""" + +import os +from glob import glob + +import numpy as np +import pytest +from jax import jit as jjit +from jax import numpy as jnp +from jax import random as jran + +from ....diffmah_kernels import DEFAULT_MAH_PARAMS +from ... import diffmahpop_params as dpp +from ... import mc_diffmahpop_kernels as mdk +from .. import kde2d_wrappers as k2w + +try: + import kdescent # noqa + + HAS_KDESCENT = True +except ImportError: + HAS_KDESCENT = False + +T_MIN_FIT = 0.5 +EPS = 1e-3 + +DATA_DRN = "/Users/aphearin/work/DATA/diffmahpop_data" +CEN_TARGET_FNAMES = sorted(glob(os.path.join(DATA_DRN, "*cen_mah*.h5"))) +try: + assert len(CEN_TARGET_FNAMES) > 0 + from astropy.table import Table + + cens = Table.read(CEN_TARGET_FNAMES[-1]) + + HAS_TARGET_DATA = True +except (AssertionError, ImportError): + HAS_TARGET_DATA = False + + +@pytest.mark.skip +def test_mc_diffmah_preds(): + ran_key = jran.key(0) + t_0 = 13.8 + lgt0 = np.log10(t_0) + n_times = 5 + + n_tests = 20 + for __ in range(n_tests): + ran_key, m_key, t_key = jran.split(ran_key, 3) + lgm_obs = jran.uniform(m_key, minval=9, maxval=16, shape=()) + t_obs = jran.uniform(t_key, minval=3, maxval=t_0, shape=()) + tarr = np.linspace(T_MIN_FIT, t_obs, n_times) + pred_data = tarr, lgm_obs, t_obs, ran_key, lgt0 + _preds = k2w.mc_diffmah_preds(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS, pred_data) + for _x in _preds: + assert np.all(np.isfinite(_x)) + lgsmar_tpt0, log_mah_tpt0, lgsmar_tp, log_mah_tp, ftpt0 = _preds + assert np.all(lgsmar_tpt0 < 20) + assert np.all(lgsmar_tp < 20) + assert np.all(log_mah_tpt0 < 20) + assert np.all(log_mah_tp < 20) + assert np.all(ftpt0 <= 1) + assert np.all(ftpt0 >= 0) + + +@pytest.mark.skip +@pytest.mark.skipif("not HAS_KDESCENT") +def test_single_sample_kde_loss_self_fit(): + """Enforce that single-sample loss has finite grads""" + ran_key = jran.key(0) + + t_0 = 13.8 + lgt0 = np.log10(t_0) + DP = 0.5 + + n_tests = 100 + for __ in range(n_tests): + + # Use random diffmahpop parameter to generate fiducial data + u_p_fid_key, ran_key = jran.split(ran_key, 2) + n_params = len(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + uran = jran.uniform(u_p_fid_key, minval=-DP, maxval=DP, shape=(n_params,)) + _u_p_list = [x + u for x, u in zip(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS, uran)] + u_p_fid = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(jnp.array(_u_p_list)) + + ran_key, lgm_key, t_obs_key = jran.split(ran_key, 3) + lgm_obs = jran.uniform(lgm_key, minval=11, maxval=15, shape=()) + t_obs = jran.uniform(t_obs_key, minval=4, maxval=t_0, shape=()) + + tarr = np.linspace(T_MIN_FIT, t_obs, k2w.N_T_PER_BIN) + + _res = k2w.get_single_sample_self_fit_target_data( + u_p_fid, tarr, lgm_obs, t_obs, ran_key, lgt0 + ) + X_target, weights_target = _res + + # use default params as the initial guess + u_p_init = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make( + dpp.DEFAULT_DIFFMAHPOP_U_PARAMS + ) + + loss_data = tarr, lgm_obs, t_obs, ran_key, lgt0, X_target, weights_target + loss = k2w.single_sample_kde_loss_self_fit(u_p_init, *loss_data) + assert np.all(np.isfinite(loss)), (lgm_obs, t_obs) + assert loss > 0 + + loss, grads = k2w.single_sample_kde_loss_and_grad_self_fit(u_p_init, *loss_data) + assert np.all(np.isfinite(grads)), (lgm_obs, t_obs) + + +@pytest.mark.skip +@pytest.mark.skipif("not HAS_KDESCENT") +def test_multisample_kde_loss_self_fit(): + """Enforce that multi-sample loss has finite grads""" + ran_key = jran.key(0) + + # Use random diffmahpop parameter to generate fiducial data + DP = 0.1 + u_p_fid_key, ran_key = jran.split(ran_key, 2) + n_params = len(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + uran = jran.uniform(u_p_fid_key, minval=-DP, maxval=DP, shape=(n_params,)) + _u_p_list = [x + u for x, u in zip(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS, uran)] + u_p_fid = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(jnp.array(_u_p_list)) + + lgmobs_key, tobs_key, ran_key = jran.split(ran_key, 3) + num_samples = 5 + t_0 = 13.8 + lgt0 = np.log10(t_0) + lgmobsarr = jran.uniform(lgmobs_key, minval=11, maxval=15, shape=(num_samples,)) + tobsarr = jran.uniform(tobs_key, minval=4, maxval=13, shape=(num_samples,)) + num_target_redshifts_per_t_obs = 10 + + tarr_matrix = jnp.array( + [jnp.linspace(T_MIN_FIT, t, num_target_redshifts_per_t_obs) for t in tobsarr] + ) + _keys = jran.split(ran_key, num_samples * 2) + _res = k2w.get_multisample_self_fit_target_data( + u_p_fid, tarr_matrix, lgmobsarr, tobsarr, _keys[:num_samples], lgt0 + ) + X_targets, weights_targets = _res + loss_data = ( + tarr_matrix, + lgmobsarr, + tobsarr, + _keys[num_samples:], + lgt0, + X_targets, + weights_targets, + ) + # use default params as the initial guess + u_p_init = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + loss = k2w.multisample_kde_loss_self_fit(u_p_init, *loss_data) + assert np.all(np.isfinite(loss)) + assert loss > 0 + + loss, grads = k2w.multisample_kde_loss_and_grad_self_fit(u_p_init, *loss_data) + assert np.all(np.isfinite(grads)) + + +@pytest.mark.skip +@pytest.mark.skipif("not HAS_KDESCENT") +def test_kdescent_adam_self_fit(): + """Enforce that kdescent.adam terminates without NaNs""" + ran_key = jran.key(0) + + # Use random diffmahpop parameter to generate fiducial data + DP = 0.5 + u_p_fid_key, ran_key = jran.split(ran_key, 2) + n_params = len(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + uran = jran.uniform(u_p_fid_key, minval=-DP, maxval=DP, shape=(n_params,)) + _u_p_list = [x + u for x, u in zip(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS, uran)] + u_p_fid = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(jnp.array(_u_p_list)) + + lgmobs_key, tobs_key, ran_key = jran.split(ran_key, 3) + num_samples = 5 + t_0 = 13.8 + lgt0 = np.log10(t_0) + lgmobsarr = jran.uniform(lgmobs_key, minval=11, maxval=15, shape=(num_samples,)) + tobsarr = jran.uniform(tobs_key, minval=4, maxval=13, shape=(num_samples,)) + num_target_redshifts_per_t_obs = 10 + + tarr_matrix = jnp.array( + [jnp.linspace(T_MIN_FIT, t, num_target_redshifts_per_t_obs) for t in tobsarr] + ) + + @jjit + def kde_loss(u_p, randkey): + _keys = jran.split(randkey, num_samples * 2) + _res = k2w.get_multisample_self_fit_target_data( + u_p_fid, tarr_matrix, lgmobsarr, tobsarr, _keys[:num_samples], lgt0 + ) + X_targets, weights_targets = _res + loss_data = ( + tarr_matrix, + lgmobsarr, + tobsarr, + _keys[num_samples:], + lgt0, + X_targets, + weights_targets, + ) + return k2w.multisample_kde_loss_self_fit(u_p, *loss_data) + + # use default params as the initial guess + u_p_init = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + + adam_results = kdescent.adam( + kde_loss, + u_p_init, + nsteps=5, + learning_rate=0.1, + randkey=12345, + ) + u_p_best = adam_results[-1] + assert np.all(np.isfinite(u_p_best)) + + +def test_get_single_cen_sample_target_data(): + nhalos = 500 + zz = np.zeros(nhalos) + mah_params = DEFAULT_MAH_PARAMS._make([zz + p for p in DEFAULT_MAH_PARAMS]) + t_obs = 10.0 + lgm_obs = 11.5 + t_0 = 13.0 + lgt0 = np.log10(t_0) + n_t = 5 + tarr = np.linspace(0.5, t_obs - 0.001, n_t) + t_peak = np.random.uniform(2, t_0, nhalos) + + _res = k2w.get_single_cen_sample_target_data( + mah_params, t_peak, tarr, lgm_obs, lgt0 + ) + for _x in _res: + assert np.all(np.isfinite(_x)) + X_target, weights_target, frac_peaked = _res + assert frac_peaked.shape == (n_t,) + assert X_target.shape == (nhalos, 2, n_t) + assert weights_target.shape == (nhalos, n_t) + + assert np.all(frac_peaked >= 0) + assert np.all(frac_peaked <= 1) + + +def test_single_sample_kde_loss_kern(): + ran_key = jran.key(0) + t_obs = 10.0 + lgm_obs = 11.5 + t_0 = 13.0 + lgt0 = np.log10(t_0) + + n_t = 5 + tarr = np.linspace(0.5, t_obs - 0.01, n_t) + + target_key, pred_key = jran.split(ran_key, 2) + _res = mdk._mc_diffmah_halo_sample( + dpp.DEFAULT_DIFFMAHPOP_PARAMS, tarr, lgm_obs, t_obs, target_key, lgt0 + ) + ( + mah_params_tpt0, + mah_params_tp, + t_peak, + ftpt0, + mc_tpt0, + dmhdt_tpt0, + log_mah_tpt0, + dmhdt_tp, + log_mah_tp, + ) = _res + mah_params = [ + np.where(mc_tpt0, x, y) for x, y in zip(mah_params_tpt0, mah_params_tp) + ] + mah_params_target = DEFAULT_MAH_PARAMS._make(mah_params) + t_peak_target = np.where(mc_tpt0, t_0, t_peak) + + _res = k2w.get_single_cen_sample_target_data( + mah_params_target, t_peak_target, tarr, lgm_obs, lgt0 + ) + for _x in _res: + assert np.all(np.isfinite(_x)) + X_target, weights_target, frac_peaked = _res + + ntests = 100 + n_pars = len(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + for __ in range(ntests): + pred_key, kde_test_key, param_key = jran.split(pred_key, 3) + uran = jran.uniform(param_key, minval=-10, maxval=10, shape=(n_pars,)) + u_p = [x + u for x, u in zip(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS, uran)] + u_params = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(u_p) + + args = ( + u_params, + tarr, + lgm_obs, + t_obs, + kde_test_key, + lgt0, + X_target, + weights_target, + frac_peaked, + ) + + loss, grads = k2w.single_sample_kde_loss_and_grad_kern(*args) + assert np.all(np.isfinite(loss)) + assert loss > 0 + for grad in grads: + assert np.all(np.isfinite(grad)) + + +@pytest.mark.skipif("not HAS_TARGET_DATA") +def test_get_cens_target_data(): + ran_key = jran.key(0) + lgt0 = np.log10(13.8) + ran_key, target_key = jran.split(ran_key, 2) + loss_data = k2w.get_cens_target_data(DATA_DRN, target_key, lgt0) + + n_tests = 10 + n_pars = len(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + for __ in range(n_tests): + ran_key, param_key = jran.split(ran_key, 2) + uran = jran.uniform(param_key, minval=-10, maxval=10, shape=(n_pars,)) + u_p = [x + u for x, u in zip(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS, uran)] + u_params = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(u_p) + loss, grads = k2w.multi_sample_kde_loss_and_grad_kern(u_params, *loss_data) + assert np.all(np.isfinite(loss)) + assert loss > 0 + assert np.all(np.isfinite(grads)) diff --git a/diffmah/diffmahpop_kernels/kdescent_testing/tests/test_kdescent_application.py b/diffmah/diffmahpop_kernels/kdescent_testing/tests/test_kdescent_application.py new file mode 100644 index 0000000..3860514 --- /dev/null +++ b/diffmah/diffmahpop_kernels/kdescent_testing/tests/test_kdescent_application.py @@ -0,0 +1,178 @@ +""" +""" + +import numpy as np +import pytest +from jax import jit as jjit +from jax import numpy as jnp +from jax import random as jran + +from ... import diffmahpop_params as dpp +from .. import dmp_wrappers as dmpw + +try: + import kdescent # noqa + + HAS_KDESCENT = True +except ImportError: + HAS_KDESCENT = False + +T_MIN_FIT = 0.5 + + +@pytest.mark.skip +@pytest.mark.xfail +@pytest.mark.skipif("not HAS_KDESCENT") +def test_single_sample_kde_loss_self_fit(): + """Enforce that single-sample loss has finite grads""" + ran_key = jran.key(0) + + t_0 = 13.8 + lgt0 = np.log10(t_0) + DP = 1.0 + num_target_redshifts_per_t_obs = 3 + + n_tests = 10 + for __ in range(n_tests): + + # Use random diffmahpop parameter to generate fiducial data + u_p_fid_key, ran_key = jran.split(ran_key, 2) + n_params = len(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + uran = jran.uniform(u_p_fid_key, minval=-DP, maxval=DP, shape=(n_params,)) + _u_p_list = [x + u for x, u in zip(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS, uran)] + u_p_fid = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(jnp.array(_u_p_list)) + + ran_key, lgm_key, t_obs_key = jran.split(ran_key, 3) + lgm_obs = jran.uniform(lgm_key, minval=10, maxval=16, shape=()) + t_obs = jran.uniform(t_obs_key, minval=3, maxval=t_0, shape=()) + + tarr = np.linspace(T_MIN_FIT, t_obs, num_target_redshifts_per_t_obs) + + _res = dmpw.get_single_sample_self_fit_target_data( + u_p_fid, tarr, lgm_obs, t_obs, ran_key, lgt0 + ) + log_mahs_target, weights_target = _res + + # use default params as the initial guess + u_p_init = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make( + dpp.DEFAULT_DIFFMAHPOP_U_PARAMS + ) + + loss_data = tarr, lgm_obs, t_obs, ran_key, lgt0, log_mahs_target, weights_target + loss = dmpw.single_sample_kde_loss_self_fit(u_p_init, *loss_data) + assert np.all(np.isfinite(loss)), (lgm_obs, t_obs) + assert loss > 0 + + loss, grads = dmpw.single_sample_kde_loss_and_grad_self_fit( + u_p_init, *loss_data + ) + assert np.all(np.isfinite(grads)), (lgm_obs, t_obs) + + +@pytest.mark.skip +@pytest.mark.xfail +@pytest.mark.skipif("not HAS_KDESCENT") +def test_multisample_kde_loss_self_fit(): + """Enforce that multi-sample loss has finite grads""" + ran_key = jran.key(0) + + # Use random diffmahpop parameter to generate fiducial data + DP = 0.1 + u_p_fid_key, ran_key = jran.split(ran_key, 2) + n_params = len(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + uran = jran.uniform(u_p_fid_key, minval=-DP, maxval=DP, shape=(n_params,)) + _u_p_list = [x + u for x, u in zip(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS, uran)] + u_p_fid = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(jnp.array(_u_p_list)) + + lgmobs_key, tobs_key, ran_key = jran.split(ran_key, 3) + num_samples = 5 + t_0 = 13.8 + lgt0 = np.log10(t_0) + lgmobsarr = jran.uniform(lgmobs_key, minval=11, maxval=15, shape=(num_samples,)) + tobsarr = jran.uniform(tobs_key, minval=4, maxval=13, shape=(num_samples,)) + num_target_redshifts_per_t_obs = 10 + + tarr_matrix = jnp.array( + [jnp.linspace(T_MIN_FIT, t, num_target_redshifts_per_t_obs) for t in tobsarr] + ) + _keys = jran.split(ran_key, num_samples * 2) + _res = dmpw.get_multisample_self_fit_target_data( + u_p_fid, tarr_matrix, lgmobsarr, tobsarr, _keys[:num_samples], lgt0 + ) + log_mahs_targets, weights_targets = _res + loss_data = ( + tarr_matrix, + lgmobsarr, + tobsarr, + _keys[num_samples:], + lgt0, + log_mahs_targets, + weights_targets, + ) + # use default params as the initial guess + u_p_init = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + loss = dmpw.multisample_kde_loss_self_fit(u_p_init, *loss_data) + assert np.all(np.isfinite(loss)) + assert loss > 0 + + loss, grads = dmpw.multisample_kde_loss_and_grad_self_fit(u_p_init, *loss_data) + assert np.all(np.isfinite(grads)) + + +@pytest.mark.skip +@pytest.mark.xfail +@pytest.mark.skipif("not HAS_KDESCENT") +def test_kdescent_adam_self_fit(): + """Enforce that kdescent.adam terminates without NaNs""" + ran_key = jran.key(0) + + # Use random diffmahpop parameter to generate fiducial data + DP = 0.01 + u_p_fid_key, ran_key = jran.split(ran_key, 2) + n_params = len(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + uran = jran.uniform(u_p_fid_key, minval=-DP, maxval=DP, shape=(n_params,)) + _u_p_list = [x + u for x, u in zip(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS, uran)] + u_p_fid = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(jnp.array(_u_p_list)) + + lgmobs_key, tobs_key, ran_key = jran.split(ran_key, 3) + num_samples = 5 + t_0 = 13.8 + lgt0 = np.log10(t_0) + lgmobsarr = jran.uniform(lgmobs_key, minval=11, maxval=15, shape=(num_samples,)) + tobsarr = jran.uniform(tobs_key, minval=4, maxval=13, shape=(num_samples,)) + num_target_redshifts_per_t_obs = 10 + + tarr_matrix = jnp.array( + [jnp.linspace(T_MIN_FIT, t, num_target_redshifts_per_t_obs) for t in tobsarr] + ) + + @jjit + def kde_loss(u_p, randkey): + _keys = jran.split(randkey, num_samples * 2) + _res = dmpw.get_multisample_self_fit_target_data( + u_p_fid, tarr_matrix, lgmobsarr, tobsarr, _keys[:num_samples], lgt0 + ) + log_mahs_targets, weights_targets = _res + loss_data = ( + tarr_matrix, + lgmobsarr, + tobsarr, + _keys[num_samples:], + lgt0, + log_mahs_targets, + weights_targets, + ) + return dmpw.multisample_kde_loss_self_fit(u_p, *loss_data) + + # use default params as the initial guess + u_p_init = dpp.DEFAULT_DIFFMAHPOP_U_PARAMS._make(dpp.DEFAULT_DIFFMAHPOP_U_PARAMS) + + adam_results = kdescent.adam( + kde_loss, + u_p_init, + nsteps=5, + learning_rate=0.1, + randkey=12345, + ) + u_p_best = adam_results[-1] + assert np.all(np.isfinite(u_p_best)) diff --git a/diffmah/diffmahpop_kernels/mc_diffmahpop_kernels.py b/diffmah/diffmahpop_kernels/mc_diffmahpop_kernels.py index 0de8551..8fc72d6 100644 --- a/diffmah/diffmahpop_kernels/mc_diffmahpop_kernels.py +++ b/diffmah/diffmahpop_kernels/mc_diffmahpop_kernels.py @@ -1,6 +1,8 @@ """ """ +from functools import partial + from jax import jit as jjit from jax import numpy as jnp from jax import random as jran @@ -132,10 +134,12 @@ def _mc_diffmah_singlecen(diffmahpop_params, tarr, lgm_obs, t_obs, ran_key, lgt0 _mc_diffmah_singlecen_vmap_kern = jjit(vmap(_mc_diffmah_singlecen, in_axes=_V)) -@jjit -def _mc_diffmah_halo_sample(diffmahpop_params, tarr, lgm_obs, t_obs, ran_key, lgt0): - zz = jnp.zeros(NH_PER_M0BIN) - ran_keys = jran.split(ran_key, NH_PER_M0BIN) +@partial(jjit, static_argnames=["n_mc"]) +def _mc_diffmah_halo_sample( + diffmahpop_params, tarr, lgm_obs, t_obs, ran_key, lgt0, n_mc=NH_PER_M0BIN +): + zz = jnp.zeros(n_mc) + ran_keys = jran.split(ran_key, n_mc) return _mc_diffmah_singlecen_vmap_kern( diffmahpop_params, tarr, lgm_obs + zz, t_obs + zz, ran_keys, lgt0 )