Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 2 additions & 3 deletions diffstar/diffstarpop/kernels/satquenchpop_model.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,4 @@
"""
"""
""" """

from collections import OrderedDict, namedtuple

Expand All @@ -11,7 +10,7 @@

DEFAULT_Q_SPEED = 5.0
DEFAULT_LGMH_K = 5.0
LGMU_SPEED = 10.0
LGMU_SPEED = 2.0

DEFAULT_SATQUENCH_PDICT = OrderedDict(t_delay=1.0, qprob_hi=0.75)
SatQuenchParams = namedtuple("SatQuenchParams", DEFAULT_SATQUENCH_PDICT.keys())
Expand Down
157 changes: 83 additions & 74 deletions diffstar/diffstarpop/kernels/sfh_pdf_mgash.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,6 @@
smoothly_clipped_line,
)


TODAY = 13.8
LGT0 = jnp.log10(TODAY)

Expand All @@ -24,34 +23,36 @@
BOUNDING_K = 0.1

SFH_PDF_QUENCH_MU_PDICT = OrderedDict(
mean_ulgm_mseq_xtp=12.027,
mean_ulgm_mseq_ytp=12.030,
mean_ulgm_mseq_lo=0.901,
mean_ulgm_mseq_hi=0.104,
mean_ulgy_mseq_int=-9.50,
mean_ulgy_mseq_slp=0.43,
mean_ul_mseq_int=-0.75,
mean_ul_mseq_slp=0.80,
mean_uh_mseq_int=-2.04,
mean_uh_mseq_slp=-3.04,
mean_ulgm_qseq_xtp=12.246,
mean_ulgm_qseq_ytp=12.200,
mean_ulgm_qseq_lo=0.812,
mean_ulgm_qseq_hi=0.094,
mean_ulgy_qseq_int=-9.50,
mean_ulgy_qseq_slp=0.43,
mean_ul_qseq_int=-0.75,
mean_ul_qseq_slp=0.80,
mean_uh_qseq_int=-2.04,
mean_uh_qseq_slp=-3.04,
mean_uqt_int=0.96,
mean_uqt_slp=-0.20,
mean_uqs_int=-0.16,
mean_uqs_slp=0.47,
mean_udrop_int=-2.05,
mean_udrop_slp=0.18,
mean_urej_int=-0.97,
mean_urej_slp=-0.06,
[
("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_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_udrop_int", -2.997),
("mean_udrop_slp", 0.431),
("mean_urej_int", -3.370),
("mean_urej_slp", 1.119),
]
)
SFH_PDF_QUENCH_MU_BOUNDS_PDICT = OrderedDict(
mean_ulgm_mseq_xtp=(11.0, 14.0),
Expand Down Expand Up @@ -85,22 +86,24 @@
)

SFH_PDF_QUENCH_COV_MS_BLOCK_PDICT = OrderedDict(
std_ulgm_mseq_int=0.325,
std_ulgm_mseq_slp=-0.028,
std_ulgy_mseq_int=0.238,
std_ulgy_mseq_slp=-0.004,
std_ul_mseq_int=0.345,
std_ul_mseq_slp=0.008,
std_uh_mseq_int=0.345,
std_uh_mseq_slp=0.008,
std_ulgm_qseq_int=0.243,
std_ulgm_qseq_slp=-0.037,
std_ulgy_qseq_int=0.327,
std_ulgy_qseq_slp=-0.082,
std_ul_qseq_int=0.210,
std_ul_qseq_slp=0.271,
std_uh_qseq_int=0.210,
std_uh_qseq_slp=0.271,
[
("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),
]
)
SFH_PDF_QUENCH_COV_MS_BLOCK_BOUNDS_PDICT = OrderedDict(
std_ulgm_mseq_int=(0.01, 1.0),
Expand All @@ -122,14 +125,16 @@
)

SFH_PDF_QUENCH_COV_Q_BLOCK_PDICT = OrderedDict(
std_uqt_int=0.070,
std_uqt_slp=-0.045,
std_uqs_int=0.444,
std_uqs_slp=-0.300,
std_udrop_int=0.779,
std_udrop_slp=-0.166,
std_urej_int=1.538,
std_urej_slp=-0.018,
[
("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),
]
)
SFH_PDF_QUENCH_COV_Q_BLOCK_BOUNDS_PDICT = OrderedDict(
std_uqt_int=(0.01, 0.5),
Expand All @@ -143,22 +148,24 @@
)

SFH_PDF_FRAC_QUENCH_PDICT = OrderedDict(
frac_quench_cen_x0_tpeak=7.0,
frac_quench_cen_k_tpeak=2.0,
frac_quench_cen_x0_ylotpeak=13.0,
frac_quench_cen_x0_yhitpeak=12.0,
frac_quench_cen_ylo_ylotpeak=0.65,
frac_quench_cen_ylo_yhitpeak=0.05,
frac_quench_cen_k=3.848,
frac_quench_cen_yhi=0.971,
frac_quench_sat_x0_tpeak=7.0,
frac_quench_sat_k_tpeak=2.0,
frac_quench_sat_x0_ylotpeak=13.0,
frac_quench_sat_x0_yhitpeak=12.0,
frac_quench_sat_ylo_ylotpeak=0.65,
frac_quench_sat_ylo_yhitpeak=0.05,
frac_quench_sat_k=3.848,
frac_quench_sat_yhi=0.971,
[
("frac_quench_cen_x0_tpeak", 10.225),
("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_sat_k", 4.995),
("frac_quench_sat_yhi", 0.887),
]
)
SFH_PDF_FRAC_QUENCH_BOUNDS_PDICT = OrderedDict(
frac_quench_cen_x0_tpeak=(1.0, 14.0),
Expand Down Expand Up @@ -201,11 +208,13 @@
)

DELTA_UQT_PDICT = OrderedDict(
delta_uqt_x0=5.0,
delta_uqt_k=1.0,
delta_uqt_ylo=-0.4,
delta_uqt_yhi=0.1,
delta_uqt_slope=0.01,
[
("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_BOUNDS_PDICT = OrderedDict(
delta_uqt_x0=(1.0, 14.0),
Expand Down
31 changes: 31 additions & 0 deletions diffstar/diffstarpop/tests/test_defaults.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,31 @@
""""""

from .. import DEFAULT_DIFFSTARPOP_PARAMS
from ..kernels.params.params_diffstarpopfits_mgash import (
DiffstarPop_Params_Diffstarpopfits_mgash,
)

DEFAULT_MODELNAME = "smdpl_dr1_nomerging"


def test_default_params_has_same_fields_as_universemachine_dr1():
params = DiffstarPop_Params_Diffstarpopfits_mgash[DEFAULT_MODELNAME]
gen = zip(
DEFAULT_DIFFSTARPOP_PARAMS._fields,
params._fields,
)
for default_key, um_dr1_key in gen:
assert default_key == um_dr1_key


def test_default_params_has_same_values_as_universemachine_dr1():
params = DiffstarPop_Params_Diffstarpopfits_mgash[DEFAULT_MODELNAME]
gen = zip(
DEFAULT_DIFFSTARPOP_PARAMS._fields,
params._fields,
)
for default_key, um_dr1_key in gen:
assert default_key == um_dr1_key
default_val = getattr(DEFAULT_DIFFSTARPOP_PARAMS, default_key)
um_dr1_val = getattr(params, um_dr1_key)
assert default_val == um_dr1_val, default_key
31 changes: 21 additions & 10 deletions diffstar/diffstarpop/tests/test_gradients.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,22 +3,33 @@
import numpy as np
from diffmah.diffmah_kernels import mah_halopop
from diffsky.mass_functions.mc_diffmah_tpeak import mc_subhalos
from diffstar.defaults import LGT0
from dsps.constants import T_TABLE_MIN
from jax import jit as jjit
from jax import numpy as jnp
from jax import random as jran
from jax import value_and_grad

from diffstar.defaults import LGT0

from .. import get_bounded_diffstarpop_params, mc_diffstar_sfh_galpop
from ..defaults import (
DEFAULT_DIFFSTARPOP_PARAMS,
DEFAULT_DIFFSTARPOP_U_PARAMS,
)
from ..defaults import DEFAULT_DIFFSTARPOP_PARAMS, DEFAULT_DIFFSTARPOP_U_PARAMS
from ..kernels.diffstarpop_mgash import _diffstarpop_means_covs
from ..kernels.params.params_diffstarpopfits_mgash import (
DiffstarPop_Params_Diffstarpopfits_mgash,
DiffstarPop_UParams_Diffstarpopfits_mgash,
)
from ..kernels.satquenchpop_model import (
SatQuenchPopUParams,
DEFAULT_SATQUENCHPOP_U_PARAMS,
SatQuenchPopUParams,
)

# sim_name_list = ["smdpl_dr1_nomerging", "smdpl_dr1", "tng", "galacticus_in_situ", "galacticus_in_plus_ex_situ"]
MODEL_NAME = "smdpl_dr1_nomerging"
TESTING_PARAMS = DEFAULT_DIFFSTARPOP_PARAMS._replace(
**DiffstarPop_Params_Diffstarpopfits_mgash[MODEL_NAME]._asdict()
)
TESTING_U_PARAMS = DEFAULT_DIFFSTARPOP_U_PARAMS._replace(
**DiffstarPop_UParams_Diffstarpopfits_mgash[MODEL_NAME]._asdict()
)


Expand All @@ -29,10 +40,10 @@ def _mse(pred, target):


def get_random_dpp_params(ran_key, dp=0.1):
u_params = jnp.array(DEFAULT_DIFFSTARPOP_U_PARAMS)
u_params = jnp.array(TESTING_U_PARAMS)
u = jran.uniform(ran_key, minval=-dp, maxval=dp, shape=(len(u_params),))
ran_u_params = np.array(u_params) + u
dpp_u_params = DEFAULT_DIFFSTARPOP_U_PARAMS._make(ran_u_params)
dpp_u_params = TESTING_U_PARAMS._make(ran_u_params)
dpp_params = get_bounded_diffstarpop_params(dpp_u_params)
return dpp_params, dpp_u_params

Expand Down Expand Up @@ -76,7 +87,7 @@ def test_all_diffstarpop_u_param_gradients_are_nonzero():

# compute SFHs for the default galaxy population
args = (
DEFAULT_DIFFSTARPOP_PARAMS,
TESTING_PARAMS,
subcat.mah_params,
subcat.logmp0,
subcat.upids,
Expand Down Expand Up @@ -180,7 +191,7 @@ def test_gradients_of_diffstarpop_pdf_satquench_params_are_nonzero():
logmhost = 13.5
gyr_since_infall = 1.0
args = (
DEFAULT_DIFFSTARPOP_PARAMS,
TESTING_PARAMS,
logmp0,
tpeak,
lgmu_infall,
Expand Down
Loading