From 5e5d5298f9d01dd16db3db6b9fc62661bce49ab7 Mon Sep 17 00:00:00 2001 From: Hans-Martin von Gaudecker Date: Fri, 14 Aug 2026 21:16:10 +0200 Subject: [PATCH 01/16] Let the NBEGM regime declare every choice its structure affords MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `build_actions` and the nongroup builder narrowed the M1 regime for NBEGM, dropping `buy_private` and `labor_supply` as actions and fixing their former outputs to constants. That made the solved model a different model from the one brute force solves: the household lost its coverage and hours choices. The regime now declares whichever choices its structure affords under every solver. pylcm's ride-along discrete envelope is written over a single action's grid and refuses a regime declaring several, so model build under NBEGM raises until that arity widens. The refusal is the honest outcome — a solver that cannot carry a choice refuses the regime rather than being handed a narrower one. Tests split accordingly: the regime-level and solver-config assertions stay green, model-build assertions become strict xfails naming the arity, and one green test pins the refusal itself so the xfails flip when the arity widens. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_018yHhuFzqsDw1MhdB1i2Ljm --- src/aca_model/baseline/regimes/_common.py | 23 ++- src/aca_model/baseline/regimes/_nongroup.py | 55 +------ src/aca_model/config.py | 17 +-- tests/test_nbegm_labor_live_validation.py | 68 +++------ tests/test_nbegm_model_creation.py | 156 ++++++++++---------- tests/test_nbegm_solve_validation.py | 48 ++---- 6 files changed, 122 insertions(+), 245 deletions(-) diff --git a/src/aca_model/baseline/regimes/_common.py b/src/aca_model/baseline/regimes/_common.py index 750bce4..eddfa52 100644 --- a/src/aca_model/baseline/regimes/_common.py +++ b/src/aca_model/baseline/regimes/_common.py @@ -454,27 +454,22 @@ def build_states(spec: RegimeSpec, grids: Grids) -> dict: return states -def build_actions( - spec: RegimeSpec, - grids: Grids, - *, - drop_buy_private: bool = False, - drop_labor_supply: bool = False, -) -> dict: +def build_actions(spec: RegimeSpec, grids: Grids) -> dict: """Build the action dict for a non-dead regime. - The `drop_*` flags fix a discrete action to a single level for the NBEGM - M1 vertical slice (its case-piece envelope handles at most one discrete - action). The dropped action's former consumers are rebound to the fixed - level at the regime builder, so removing it here is the action side of the - dags remove-and-fix. + Every choice the regime's structure affords is a live action: whether to + claim Social Security where claiming is neither impossible nor already + forced, how many hours to work where work is still available, and whether + to buy non-group coverage before Medicare. Which solver runs the regime + does not enter — a solver that cannot carry a choice refuses the regime + rather than being handed a narrower one. """ actions: dict = {} if spec["ss"] == "choose": actions["claim_ss"] = DiscreteGrid(ClaimedSS) - if spec["canwork"] == "canwork" and not drop_labor_supply: + if spec["canwork"] == "canwork": actions["labor_supply"] = DiscreteGrid(LaborSupply) - if spec["his"] == "nongroup" and spec["mc"] == "nomc" and not drop_buy_private: + if spec["his"] == "nongroup" and spec["mc"] == "nomc": actions["buy_private"] = DiscreteGrid(BuyPrivate) actions["consumption_dollars"] = grids.consumption_dollars return actions diff --git a/src/aca_model/baseline/regimes/_nongroup.py b/src/aca_model/baseline/regimes/_nongroup.py index 8d329ea..5189cf9 100644 --- a/src/aca_model/baseline/regimes/_nongroup.py +++ b/src/aca_model/baseline/regimes/_nongroup.py @@ -4,7 +4,6 @@ Already nongroup, so no SSI/Medicaid override needed for HIS transitions. """ -import functools from collections.abc import Callable from lcm import Regime @@ -13,7 +12,6 @@ from aca_model.agent.labor_market import LaborSupply from aca_model.baseline import health_insurance -from aca_model.baseline.health_insurance import BuyPrivate from aca_model.baseline.regimes._common import ( REGIME_SPECS, Grids, @@ -78,32 +76,11 @@ def transition( return transition -def _fixed_full_time_labor_supply() -> DiscreteAction: - """Labor supply fixed to full-time work for the NBEGM M1 slice.""" - return LaborSupply.h2000 - - -def _build_functions( - spec: RegimeSpec, *, fix_buy_private: bool = False, fix_labor_supply: bool = False -) -> dict: - """Build functions dict for a nongroup regime. - - The NBEGM M1 slice fixes both discrete actions to a single level so the - only choice is continuous consumption: - - - `fix_buy_private` binds `buy_private` to `BuyPrivate.yes` in its consumers - (premium, OOP) — the `buy_private == BuyPrivate.yes` arm — leaving the - remaining budget structure untouched. - - `fix_labor_supply` supplies `labor_supply` as a fixed full-time node read - by labor income, AIME accrual, and the lagged-supply transition (which - stays a state, so the cross-regime continuation space is unchanged). - """ +def _build_functions(spec: RegimeSpec) -> dict: + """Build functions dict for a nongroup regime.""" can_work = spec["canwork"] == "canwork" functions = build_common_functions(spec) - if can_work and fix_labor_supply: - functions["labor_supply"] = _fixed_full_time_labor_supply - functions["ss_benefit"] = select_ss_benefit(spec) # his and crossed_oamc_threshold are fixed params (constants per regime), @@ -117,14 +94,6 @@ def _build_functions( else: functions["hic_premium"] = health_insurance.premium_retired - if has_buy_private and fix_buy_private: - functions["hic_premium"] = functools.partial( - health_insurance.premium, buy_private=BuyPrivate.yes - ) - functions["primary_oop"] = functools.partial( - health_insurance.primary_oop, buy_private=BuyPrivate.yes - ) - functions.update(build_pension_functions(spec)) return functions @@ -156,18 +125,9 @@ def build_regime( if egm_solver is None else ("nbegm" if nbegm_solver is not None else "dcegm") ) - # Under NBEGM the M1 slice fixes `buy_private` (a second discrete action the - # branch compiler does not yet solve) to a single level. `labor_supply` is fixed - # too by default, leaving only continuous consumption; with - # `nbegm_live_labor_supply` it stays a live action and the branch compiler solves - # each labor level against the cliffed budget. - fix_for_nbegm = nbegm_solver is not None - fix_labor = fix_for_nbegm and not grids.grid_config.nbegm_live_labor_supply - functions = _build_functions( - spec, fix_buy_private=fix_for_nbegm, fix_labor_supply=fix_labor - ) + functions = _build_functions(spec) constraints: dict = {} - if fix_for_nbegm: + if nbegm_solver is not None: # NBEGM solves only this regime, so its solver-contract functions are # regime-level here rather than broadcast model-wide. The broadcast # borrowing constraint stays: the EGM solve enforces the limit through @@ -182,12 +142,7 @@ def build_regime( active=make_active_func(spec), states=states, state_transitions=build_state_transitions(spec, solver=state_solver), - actions=build_actions( - spec, - grids, - drop_buy_private=fix_for_nbegm, - drop_labor_supply=fix_labor, - ), + actions=build_actions(spec, grids), functions=functions, constraints=constraints, **solver_kwargs, diff --git a/src/aca_model/config.py b/src/aca_model/config.py index 47390c2..db5ed3e 100644 --- a/src/aca_model/config.py +++ b/src/aca_model/config.py @@ -151,18 +151,11 @@ class GridConfig: # value scans it in blocks of that many branches (identical result, per-branch # intermediates bounded by one block). Only consulted under `solver="nbegm"`. n_nbegm_branch_batch_size: int = 0 - # Keep `labor_supply` a live discrete action on the M1 regime under NBEGM (the - # branch compiler solves each labor level against its own continuation, utility, - # and breakpoint partition); `buy_private` stays fixed. `False` fixes both actions - # to a single level so the only choice is continuous consumption. Only consulted - # under `solver="nbegm"`. - # - # Requires `nbegm_jump_read="bridged"`. `labor_supply` enters `countable_income`, - # which carries the SSI income test, so each labor level puts that breakpoint at a - # different liquid level; the one-sided read publishes its cliff limits on a single - # query grid shared across branches, which the two cannot both satisfy. Building a - # model with live labor under `"one_sided"` raises `RegimeInitializationError`. - nbegm_live_labor_supply: bool = False + # `labor_supply` enters `countable_income`, which carries the SSI income test, so + # each labor level puts that breakpoint at a different liquid level. The one-sided + # read publishes its cliff limits on a single query grid shared across branches, + # which the two cannot both satisfy, so a regime carrying `labor_supply` builds + # only under `nbegm_jump_read="bridged"`. MODEL_CONFIG = ModelConfig() diff --git a/tests/test_nbegm_labor_live_validation.py b/tests/test_nbegm_labor_live_validation.py index 81eccf4..c9fbc74 100644 --- a/tests/test_nbegm_labor_live_validation.py +++ b/tests/test_nbegm_labor_live_validation.py @@ -1,15 +1,15 @@ -"""NBEGM solves the M1 regime with a live labor-supply choice, matching brute. - -The M1 regime `nongroup_nomc_inelig_canwork` carries `labor_supply` (5 levels) as a -genuine discrete action while `buy_private` is fixed. Labor supply feeds the `aime` -co-state (earnings accrual), the `lagged_labor_supply` co-state, and the leisure term -in period utility — every branch-dependent channel at once. NBEGM's ride-along -discrete envelope solves each labor branch against its own continuation and utility; -the value function must match a brute-force solve on the same state-action space. - -Both solvers run the full 18-regime model at the benchmark grid with `buy_private` -fixed and `labor_supply` live, in NBEGM's `"bridged"` cliff-read mode so the -comparison isolates the solver machinery from the asset-grid read convention. +"""NBEGM solves the M1 regime with its discrete choices live, matching brute. + +The M1 regime `nongroup_nomc_inelig_canwork` carries `labor_supply` (5 levels) and +`buy_private` as genuine discrete actions. Labor supply feeds the `aime` co-state +(earnings accrual), the `lagged_labor_supply` co-state, and the leisure term in period +utility — every branch-dependent channel at once. NBEGM's ride-along discrete envelope +solves each branch against its own continuation and utility; the value function must +match a brute-force solve on the same state-action space. + +Both solvers run the full 18-regime model at the benchmark grid in NBEGM's `"bridged"` +cliff-read mode, so the comparison isolates the solver machinery from the asset-grid +read convention. """ import dataclasses @@ -19,51 +19,19 @@ from helpers.model import _DERIVED_CATEGORICALS # ty: ignore[unresolved-import] from lcm import DiscreteGrid -import aca_model.baseline.regimes._nongroup as nongroup_mod from aca_model.agent.preferences import BenchmarkPrefType from aca_model.baseline.model import create_model from aca_model.baseline.regimes import SolverName -from aca_model.baseline.regimes._common import Grids, RegimeSpec from aca_model.benchmark import get_benchmark_params from aca_model.config import BENCHMARK_GRID_CONFIG _M1_REGIME = "nongroup_nomc_inelig_canwork" - -def _is_m1(spec: RegimeSpec) -> bool: - return spec["ss"] == "inelig" and spec["mc"] == "nomc" - - -@pytest.fixture -def m1_labor_live(monkeypatch: pytest.MonkeyPatch) -> None: - """Keep `labor_supply` live and fix only `buy_private` on M1, under every solver. - - The NBEGM wiring fixes both discrete actions on its own; this override keeps - labor supply as a genuine action (so the branch compiler solves it) and fixes - `buy_private`, applying the same choice to the brute build so both solvers share - one state-action space. - """ - original_build_functions = nongroup_mod._build_functions # noqa: SLF001 - original_build_actions = nongroup_mod.build_actions - - def build_functions_labor_live(spec: RegimeSpec, **kwargs: bool) -> dict: - if _is_m1(spec): - return original_build_functions( - spec, fix_buy_private=True, fix_labor_supply=False - ) - return original_build_functions(spec, **kwargs) - - def build_actions_labor_live( - spec: RegimeSpec, grids: Grids, **kwargs: bool - ) -> dict: - if _is_m1(spec): - return original_build_actions( - spec, grids, drop_buy_private=True, drop_labor_supply=False - ) - return original_build_actions(spec, grids, **kwargs) - - monkeypatch.setattr(nongroup_mod, "_build_functions", build_functions_labor_live) - monkeypatch.setattr(nongroup_mod, "build_actions", build_actions_labor_live) +_MULTIPLE_DISCRETE_ACTIONS = ( + "pylcm's ride-along discrete envelope is written over one action's grid and " + "refuses a regime declaring several; the M1 regime declares both " + "`labor_supply` and `buy_private`." +) def _solve_m1(solver: SolverName) -> tuple[dict[int, np.ndarray], int]: @@ -127,7 +95,7 @@ def _cliff_band_mask(reference: np.ndarray, assets_axis: int) -> np.ndarray: @pytest.mark.long_running -@pytest.mark.usefixtures("m1_labor_live") +@pytest.mark.xfail(strict=True, reason=_MULTIPLE_DISCRETE_ACTIONS) def test_nbegm_m1_labor_live_agrees_with_brute_split_by_cliff_band() -> None: """The M1 value functions agree cell-wise with labor live, gated per region. diff --git a/tests/test_nbegm_model_creation.py b/tests/test_nbegm_model_creation.py index d98bebb..3074ae7 100644 --- a/tests/test_nbegm_model_creation.py +++ b/tests/test_nbegm_model_creation.py @@ -1,11 +1,17 @@ """NBEGM solver wiring: `solver="nbegm"` is a per-regime option. Unlike DC-EGM (a global Euler solver on every living regime), NBEGM solves a -single 1-D consumption/savings regime with at most one discrete action, so it -attaches only to the M1 vertical-slice regime `nongroup_nomc_inelig_canwork`; -every other living regime keeps brute force. The savings-form spec is shared -with DC-EGM (NBEGM's budget node is `resources`, the post-decision function is -`savings`). +single 1-D consumption/savings regime, so it attaches only to the vertical-slice +regime `nongroup_nomc_inelig_canwork`; every other living regime keeps brute +force. The savings-form spec is shared with DC-EGM (NBEGM's budget node is +`resources`, the post-decision function is `savings`). + +The regime declares every choice its structure affords — whether to buy +non-group coverage, and how many hours to work. pylcm's ride-along discrete +envelope is written over a single action's grid and refuses a regime declaring +more than one, so model build under NBEGM raises until that arity is widened. +The regime is not narrowed to fit: a solver that cannot carry a choice refuses +the regime rather than being handed a model that omits it. """ import dataclasses @@ -40,6 +46,20 @@ _M1_REGIME = "nongroup_nomc_inelig_canwork" _BRUTE_REGIME = "retiree_nomc_inelig_canwork" +_MULTIPLE_DISCRETE_ACTIONS = ( + "pylcm's ride-along discrete envelope is written over one action's grid and " + "refuses a regime declaring several; the M1 regime declares both " + "`labor_supply` and `buy_private`." +) + +# `labor_supply` enters `countable_income`, which carries the SSI income test, so each +# labor level puts that breakpoint at a different liquid level. The one-sided read +# publishes its cliff limits on a single query grid shared across branches, which the +# two cannot both satisfy, so a regime carrying `labor_supply` needs the bridged read. +_BRIDGED_GRID_CONFIG = dataclasses.replace( + BENCHMARK_GRID_CONFIG, nbegm_jump_read="bridged" +) + def _build_regimes(solver: SolverName) -> dict[str, Regime]: return build_all_regimes( @@ -64,7 +84,7 @@ def _build_model_with(solver: SolverName, grid_config: GridConfig) -> Model: def _build_model(solver: SolverName) -> Model: - return _build_model_with(solver, BENCHMARK_GRID_CONFIG) + return _build_model_with(solver, _BRIDGED_GRID_CONFIG) def _grids() -> Grids: @@ -120,30 +140,42 @@ def test_build_nbegm_solver_forwards_the_jump_read_mode() -> None: assert solver.jump_read == "bridged" -def test_nbegm_m1_regime_fixes_buy_private() -> None: - """The NBEGM M1 slice drops `buy_private` as an action (fixed to purchase), - so the only choice is continuous consumption; the brute M1 regime keeps it.""" +def test_nbegm_m1_regime_declares_the_same_actions_as_brute_force() -> None: + """Which solver runs a regime does not change the choices the household has. + + The M1 regime affords a coverage choice and an hours choice, so it declares + both under either solver. + """ nbegm_m1 = _build_regimes("nbegm")[_M1_REGIME] brute_m1 = _build_regimes("brute_force")[_M1_REGIME] - assert "buy_private" not in nbegm_m1.actions - assert "buy_private" in brute_m1.actions + assert set(nbegm_m1.actions) == set(brute_m1.actions) -def test_nbegm_m1_regime_fixes_labor_supply() -> None: - """The NBEGM M1 slice drops `labor_supply` as an action (fixed to full-time - work), so no discrete action remains and the only choice is continuous - consumption; the brute M1 regime keeps `labor_supply`.""" +def test_nbegm_m1_regime_declares_both_discrete_actions() -> None: + """The M1 regime's discrete choices are whether to buy non-group coverage and + how many hours to work.""" nbegm_m1 = _build_regimes("nbegm")[_M1_REGIME] - brute_m1 = _build_regimes("brute_force")[_M1_REGIME] - assert "labor_supply" not in nbegm_m1.actions - assert "labor_supply" in brute_m1.actions + discrete = { + name + for name, grid in nbegm_m1.actions.items() + if isinstance(grid, DiscreteGrid) + } + assert discrete == {"buy_private", "labor_supply"} -def test_nbegm_m1_regime_has_no_discrete_action() -> None: - """With both discrete actions fixed, the NBEGM M1 slice leaves only the - continuous consumption choice — no `DiscreteGrid` action remains.""" - nbegm_m1 = _build_regimes("nbegm")[_M1_REGIME] - assert not any(isinstance(grid, DiscreteGrid) for grid in nbegm_m1.actions.values()) +def test_nbegm_model_build_refuses_a_regime_with_several_discrete_actions() -> None: + """Model build under NBEGM refuses the M1 regime, naming its discrete actions. + + The ride-along discrete envelope is written over a single action's grid. The + regime is not narrowed to fit it — the refusal is the honest outcome until + the envelope carries several actions. + """ + with pytest.raises( + RegimeInitializationError, match="exactly one discrete action" + ) as excinfo: + _build_model("nbegm") + assert "labor_supply" in str(excinfo.value) + assert "buy_private" in str(excinfo.value) def test_nbegm_m1_regime_takes_the_savings_form_assets_laws() -> None: @@ -162,6 +194,7 @@ def test_nbegm_m1_regime_takes_the_savings_form_assets_laws() -> None: assert law is expected, target_name +@pytest.mark.xfail(strict=True, reason=_MULTIPLE_DISCRETE_ACTIONS) def test_nbegm_savings_form_functions_are_scoped_to_the_m1_regime() -> None: """Under NBEGM only the M1 regime carries the savings-form budget functions (`resources`, `savings`); brute regimes keep the cash-on-hand form and carry @@ -173,6 +206,7 @@ def test_nbegm_savings_form_functions_are_scoped_to_the_m1_regime() -> None: assert "resources" not in model.user_regimes[_BRUTE_REGIME].functions +@pytest.mark.xfail(strict=True, reason=_MULTIPLE_DISCRETE_ACTIONS) def test_nbegm_m1_regime_does_not_carry_inverse_marginal_utility() -> None: """NBEGM inverts the Euler equation internally, so the M1 regime never carries the DC-EGM `inverse_marginal_utility` function (whose @@ -182,6 +216,7 @@ def test_nbegm_m1_regime_does_not_carry_inverse_marginal_utility() -> None: assert "inverse_marginal_utility" not in model.user_regimes[_M1_REGIME].functions +@pytest.mark.xfail(strict=True, reason=_MULTIPLE_DISCRETE_ACTIONS) def test_nbegm_m1_regime_keeps_the_borrowing_constraint() -> None: """The M1 regime declares the borrowing constraint like every brute regime. @@ -234,66 +269,28 @@ def test_ssi_benefit_declares_the_income_test_kink() -> None: assert income_test.indexed_by == "spousal_income" -def test_nbegm_keeps_labor_supply_live_when_configured() -> None: - """With `nbegm_live_labor_supply=True`, the M1 regime carries `labor_supply` - as a genuine discrete action under NBEGM while `buy_private` stays fixed, so - the branch compiler solves the labor choice against the cliffed budget.""" - grid_config = dataclasses.replace( - BENCHMARK_GRID_CONFIG, nbegm_live_labor_supply=True - ) - regimes = build_all_regimes( - grid_config=grid_config, - fixed_params=_FIXED_PARAMS, - wage_params=_WAGE_PARAMS, - pref_type_grid=DiscreteGrid(BenchmarkPrefType), - solver="nbegm", - ) - actions = regimes[_M1_REGIME].actions - assert "labor_supply" in actions - assert "buy_private" not in actions - - -def test_nbegm_live_labor_supply_requires_the_bridged_cliff_read() -> None: - """A live `labor_supply` action builds only under `nbegm_jump_read="bridged"`. +@pytest.mark.xfail(strict=True, reason=_MULTIPLE_DISCRETE_ACTIONS) +def test_nbegm_labor_supply_requires_the_bridged_cliff_read() -> None: + """The M1 regime builds under NBEGM only with `nbegm_jump_read="bridged"`. `labor_supply` enters `countable_income`, which carries the SSI income test, so each labor level puts that breakpoint at a different liquid level. The one-sided read publishes its cliff limits on one query grid shared across branches, so the two cannot both hold and the build is refused. """ - live_labor = dataclasses.replace( - BENCHMARK_GRID_CONFIG, nbegm_live_labor_supply=True - ) with pytest.raises(RegimeInitializationError, match="must not enter any schedule"): - _build_model_with( - "nbegm", dataclasses.replace(live_labor, nbegm_jump_read="one_sided") - ) + _build_model_with("nbegm", BENCHMARK_GRID_CONFIG) - bridged = dataclasses.replace(live_labor, nbegm_jump_read="bridged") - assert isinstance(_build_model_with("nbegm", bridged), Model) - - -def test_nbegm_fixes_labor_supply_by_default() -> None: - """By default NBEGM fixes both discrete actions on the M1 regime, so the - only remaining choice is continuous consumption against the cliffed budget.""" - regimes = _build_regimes("nbegm") - actions = regimes[_M1_REGIME].actions - assert "labor_supply" not in actions - assert "buy_private" not in actions + assert isinstance(_build_model_with("nbegm", _BRIDGED_GRID_CONFIG), Model) +@pytest.mark.xfail(strict=True, reason=_MULTIPLE_DISCRETE_ACTIONS) @pytest.mark.parametrize("policy", list(PolicyVariant)) def test_nbegm_builds_every_aca_policy_variant(policy: PolicyVariant) -> None: """Every ACA policy variant builds a model under NBEGM with the M1 regime on - the solver and labor live — the overlay's function swaps compose with the - branch compiler's per-regime wiring.""" - # Live labor requires the bridged cliff read — see - # `test_nbegm_live_labor_supply_requires_the_bridged_cliff_read`. - grid_config = dataclasses.replace( - BENCHMARK_GRID_CONFIG, - nbegm_live_labor_supply=True, - nbegm_jump_read="bridged", - ) + the solver and both discrete choices live — the overlay's function swaps + compose with the branch compiler's per-regime wiring.""" + grid_config = _BRIDGED_GRID_CONFIG model = create_aca_model( n_subjects=1, policy=policy, @@ -317,21 +314,18 @@ def test_nbegm_builds_every_aca_policy_variant(policy: PolicyVariant) -> None: assert "labor_supply" in regimes[_M1_REGIME].actions +@pytest.mark.xfail(strict=True, reason=_MULTIPLE_DISCRETE_ACTIONS) @pytest.mark.parametrize("policy", list(PolicyVariant)) def test_nbegm_aca_variants_leave_no_free_buy_private_params( policy: PolicyVariant, ) -> None: - """With `buy_private` fixed under the NBEGM M1 slice, no ACA-swapped - function may leave `buy_private` as a free parameter — the params template - holds no `buy_private` leaves, so solve/simulate never demand a - `*__buy_private` entry the pipeline cannot supply.""" - # Live labor requires the bridged cliff read — see - # `test_nbegm_live_labor_supply_requires_the_bridged_cliff_read`. - grid_config = dataclasses.replace( - BENCHMARK_GRID_CONFIG, - nbegm_live_labor_supply=True, - nbegm_jump_read="bridged", - ) + """`buy_private` is a choice, never a parameter. + + No ACA-swapped function may leave it as a free parameter — the params + template holds no `buy_private` leaves, so solve/simulate never demand a + `*__buy_private` entry the pipeline cannot supply. + """ + grid_config = _BRIDGED_GRID_CONFIG model = create_aca_model( n_subjects=1, policy=policy, diff --git a/tests/test_nbegm_solve_validation.py b/tests/test_nbegm_solve_validation.py index 6d4449d..14f2e45 100644 --- a/tests/test_nbegm_solve_validation.py +++ b/tests/test_nbegm_solve_validation.py @@ -1,9 +1,9 @@ """NBEGM vs brute-force agreement on the M1 slice value function. -Both solvers run the full 18-regime model at the benchmark grid with the M1 -regime's discrete actions fixed (labor supply to full-time, `buy_private` to -purchase) so the value functions are defined on identical state spaces and -differ only through the continuous-consumption solver. +Both solvers run the full 18-regime model at the benchmark grid on the same +state-action space — the M1 regime declares its `labor_supply` and `buy_private` +choices under either solver — so the value functions differ only through the +continuous-consumption solver. NBEGM runs its `"bridged"` cliff-read mode so both solvers share the same read convention (a finite brute reads child values by linear interpolation @@ -22,47 +22,19 @@ from helpers.model import _DERIVED_CATEGORICALS # ty: ignore[unresolved-import] from lcm import DiscreteGrid -import aca_model.baseline.regimes._nongroup as nongroup_mod from aca_model.agent.preferences import BenchmarkPrefType from aca_model.baseline.model import create_model from aca_model.baseline.regimes import SolverName -from aca_model.baseline.regimes._common import Grids, RegimeSpec from aca_model.benchmark import get_benchmark_params from aca_model.config import BENCHMARK_GRID_CONFIG _M1_REGIME = "nongroup_nomc_inelig_canwork" - -@pytest.fixture -def m1_actions_fixed_for_brute(monkeypatch: pytest.MonkeyPatch) -> None: - """Fix the M1 regime's discrete actions under every solver. - - The NBEGM wiring fixes `labor_supply` and `buy_private` on its own; this - fixture applies the same fixing to the brute build so both solvers produce - an M1 value function on the same state-action space. - """ - original_build_functions = nongroup_mod._build_functions # noqa: SLF001 - original_build_actions = nongroup_mod.build_actions - - def build_functions_fixed(spec: RegimeSpec, **kwargs: bool) -> dict: - if spec["ss"] == "inelig" and spec["mc"] == "nomc": - return original_build_functions( - spec, fix_buy_private=True, fix_labor_supply=True - ) - return original_build_functions(spec, **kwargs) - - def build_actions_fixed(spec: RegimeSpec, grids: Grids, **kwargs: bool) -> dict: - if spec["ss"] == "inelig" and spec["mc"] == "nomc": - return original_build_actions( - spec, - grids, - drop_buy_private=True, - drop_labor_supply=True, - ) - return original_build_actions(spec, grids, **kwargs) - - monkeypatch.setattr(nongroup_mod, "_build_functions", build_functions_fixed) - monkeypatch.setattr(nongroup_mod, "build_actions", build_actions_fixed) +_MULTIPLE_DISCRETE_ACTIONS = ( + "pylcm's ride-along discrete envelope is written over one action's grid and " + "refuses a regime declaring several; the M1 regime declares both " + "`labor_supply` and `buy_private`." +) def _solve_m1(solver: SolverName) -> dict[int, np.ndarray]: @@ -87,7 +59,7 @@ def _solve_m1(solver: SolverName) -> dict[int, np.ndarray]: @pytest.mark.long_running -@pytest.mark.usefixtures("m1_actions_fixed_for_brute") +@pytest.mark.xfail(strict=True, reason=_MULTIPLE_DISCRETE_ACTIONS) def test_nbegm_m1_value_function_agrees_with_brute_in_the_bulk() -> None: """The M1 value functions agree cell-wise away from the cliff tail. From 3b2a0b1f1d10e20de4dd1a7f13eb51c69932294a Mon Sep 17 00:00:00 2001 From: Hans-Martin von Gaudecker Date: Sat, 15 Aug 2026 06:58:42 +0200 Subject: [PATCH 02/16] Consume pylcm's multi-action discrete envelope The M1 regime declares both buy_private and labor_supply, and pylcm's envelope now branches over the Cartesian product of a regime's discrete action grids, so the eight strict xfails pinned to the old single-action refusal all XPASS. Removes them and converts the refusal test into the capability it replaced: the model builds with both actions live. 25 passed where 15 had failed. Every one of the 15 was a stale expectation -- including the five policy-variant builds, which are reported as FAILED rather than XPASS under strict=True and so read like real build failures until the XPASS(strict) marker is inspected. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_018yHhuFzqsDw1MhdB1i2Ljm --- tests/test_nbegm_labor_live_validation.py | 7 ---- tests/test_nbegm_model_creation.py | 44 +++++++++-------------- tests/test_nbegm_solve_validation.py | 7 ---- 3 files changed, 16 insertions(+), 42 deletions(-) diff --git a/tests/test_nbegm_labor_live_validation.py b/tests/test_nbegm_labor_live_validation.py index c9fbc74..a0ef0f0 100644 --- a/tests/test_nbegm_labor_live_validation.py +++ b/tests/test_nbegm_labor_live_validation.py @@ -27,12 +27,6 @@ _M1_REGIME = "nongroup_nomc_inelig_canwork" -_MULTIPLE_DISCRETE_ACTIONS = ( - "pylcm's ride-along discrete envelope is written over one action's grid and " - "refuses a regime declaring several; the M1 regime declares both " - "`labor_supply` and `buy_private`." -) - def _solve_m1(solver: SolverName) -> tuple[dict[int, np.ndarray], int]: # The CPU XLA backend does not fuse the ride-cell fan-out and materialises the @@ -95,7 +89,6 @@ def _cliff_band_mask(reference: np.ndarray, assets_axis: int) -> np.ndarray: @pytest.mark.long_running -@pytest.mark.xfail(strict=True, reason=_MULTIPLE_DISCRETE_ACTIONS) def test_nbegm_m1_labor_live_agrees_with_brute_split_by_cliff_band() -> None: """The M1 value functions agree cell-wise with labor live, gated per region. diff --git a/tests/test_nbegm_model_creation.py b/tests/test_nbegm_model_creation.py index 3074ae7..a8ed368 100644 --- a/tests/test_nbegm_model_creation.py +++ b/tests/test_nbegm_model_creation.py @@ -7,11 +7,10 @@ `resources`, the post-decision function is `savings`). The regime declares every choice its structure affords — whether to buy -non-group coverage, and how many hours to work. pylcm's ride-along discrete -envelope is written over a single action's grid and refuses a regime declaring -more than one, so model build under NBEGM raises until that arity is widened. -The regime is not narrowed to fit: a solver that cannot carry a choice refuses -the regime rather than being handed a model that omits it. +non-group coverage, and how many hours to work — and the discrete envelope +branches over the Cartesian product of those grids. The regime is never narrowed +to fit the solver: a solver that cannot carry a choice refuses the regime rather +than being handed a model that omits it. """ import dataclasses @@ -46,12 +45,6 @@ _M1_REGIME = "nongroup_nomc_inelig_canwork" _BRUTE_REGIME = "retiree_nomc_inelig_canwork" -_MULTIPLE_DISCRETE_ACTIONS = ( - "pylcm's ride-along discrete envelope is written over one action's grid and " - "refuses a regime declaring several; the M1 regime declares both " - "`labor_supply` and `buy_private`." -) - # `labor_supply` enters `countable_income`, which carries the SSI income test, so each # labor level puts that breakpoint at a different liquid level. The one-sided read # publishes its cliff limits on a single query grid shared across branches, which the @@ -163,19 +156,20 @@ def test_nbegm_m1_regime_declares_both_discrete_actions() -> None: assert discrete == {"buy_private", "labor_supply"} -def test_nbegm_model_build_refuses_a_regime_with_several_discrete_actions() -> None: - """Model build under NBEGM refuses the M1 regime, naming its discrete actions. +def test_nbegm_model_builds_a_regime_declaring_several_discrete_actions() -> None: + """The M1 regime builds under NBEGM with both of its discrete actions live. - The ride-along discrete envelope is written over a single action's grid. The - regime is not narrowed to fit it — the refusal is the honest outcome until - the envelope carries several actions. + The discrete envelope branches over the Cartesian product of the regime's + discrete action grids, so a regime is never narrowed to fit the solver. """ - with pytest.raises( - RegimeInitializationError, match="exactly one discrete action" - ) as excinfo: - _build_model("nbegm") - assert "labor_supply" in str(excinfo.value) - assert "buy_private" in str(excinfo.value) + model = _build_model("nbegm") + nbegm_m1 = model.user_regimes[_M1_REGIME] + discrete = { + name + for name, grid in nbegm_m1.actions.items() + if isinstance(grid, DiscreteGrid) + } + assert discrete == {"buy_private", "labor_supply"} def test_nbegm_m1_regime_takes_the_savings_form_assets_laws() -> None: @@ -194,7 +188,6 @@ def test_nbegm_m1_regime_takes_the_savings_form_assets_laws() -> None: assert law is expected, target_name -@pytest.mark.xfail(strict=True, reason=_MULTIPLE_DISCRETE_ACTIONS) def test_nbegm_savings_form_functions_are_scoped_to_the_m1_regime() -> None: """Under NBEGM only the M1 regime carries the savings-form budget functions (`resources`, `savings`); brute regimes keep the cash-on-hand form and carry @@ -206,7 +199,6 @@ def test_nbegm_savings_form_functions_are_scoped_to_the_m1_regime() -> None: assert "resources" not in model.user_regimes[_BRUTE_REGIME].functions -@pytest.mark.xfail(strict=True, reason=_MULTIPLE_DISCRETE_ACTIONS) def test_nbegm_m1_regime_does_not_carry_inverse_marginal_utility() -> None: """NBEGM inverts the Euler equation internally, so the M1 regime never carries the DC-EGM `inverse_marginal_utility` function (whose @@ -216,7 +208,6 @@ def test_nbegm_m1_regime_does_not_carry_inverse_marginal_utility() -> None: assert "inverse_marginal_utility" not in model.user_regimes[_M1_REGIME].functions -@pytest.mark.xfail(strict=True, reason=_MULTIPLE_DISCRETE_ACTIONS) def test_nbegm_m1_regime_keeps_the_borrowing_constraint() -> None: """The M1 regime declares the borrowing constraint like every brute regime. @@ -269,7 +260,6 @@ def test_ssi_benefit_declares_the_income_test_kink() -> None: assert income_test.indexed_by == "spousal_income" -@pytest.mark.xfail(strict=True, reason=_MULTIPLE_DISCRETE_ACTIONS) def test_nbegm_labor_supply_requires_the_bridged_cliff_read() -> None: """The M1 regime builds under NBEGM only with `nbegm_jump_read="bridged"`. @@ -284,7 +274,6 @@ def test_nbegm_labor_supply_requires_the_bridged_cliff_read() -> None: assert isinstance(_build_model_with("nbegm", _BRIDGED_GRID_CONFIG), Model) -@pytest.mark.xfail(strict=True, reason=_MULTIPLE_DISCRETE_ACTIONS) @pytest.mark.parametrize("policy", list(PolicyVariant)) def test_nbegm_builds_every_aca_policy_variant(policy: PolicyVariant) -> None: """Every ACA policy variant builds a model under NBEGM with the M1 regime on @@ -314,7 +303,6 @@ def test_nbegm_builds_every_aca_policy_variant(policy: PolicyVariant) -> None: assert "labor_supply" in regimes[_M1_REGIME].actions -@pytest.mark.xfail(strict=True, reason=_MULTIPLE_DISCRETE_ACTIONS) @pytest.mark.parametrize("policy", list(PolicyVariant)) def test_nbegm_aca_variants_leave_no_free_buy_private_params( policy: PolicyVariant, diff --git a/tests/test_nbegm_solve_validation.py b/tests/test_nbegm_solve_validation.py index 14f2e45..028fce7 100644 --- a/tests/test_nbegm_solve_validation.py +++ b/tests/test_nbegm_solve_validation.py @@ -30,12 +30,6 @@ _M1_REGIME = "nongroup_nomc_inelig_canwork" -_MULTIPLE_DISCRETE_ACTIONS = ( - "pylcm's ride-along discrete envelope is written over one action's grid and " - "refuses a regime declaring several; the M1 regime declares both " - "`labor_supply` and `buy_private`." -) - def _solve_m1(solver: SolverName) -> dict[int, np.ndarray]: grid_config = dataclasses.replace(BENCHMARK_GRID_CONFIG, nbegm_jump_read="bridged") @@ -59,7 +53,6 @@ def _solve_m1(solver: SolverName) -> dict[int, np.ndarray]: @pytest.mark.long_running -@pytest.mark.xfail(strict=True, reason=_MULTIPLE_DISCRETE_ACTIONS) def test_nbegm_m1_value_function_agrees_with_brute_in_the_bulk() -> None: """The M1 value functions agree cell-wise away from the cliff tail. From 9a5fe8a1610e03ac7676095c89a0726ba66ba527 Mon Sep 17 00:00:00 2001 From: Hans-Martin von Gaudecker Date: Sat, 15 Aug 2026 08:39:02 +0200 Subject: [PATCH 03/16] Solve every living regime with NB-EGM `solver="nbegm"` attached the solver to one regime and left the other 17 on brute force, so the model's stated solver was not the one that produced most of its result. Every living regime now gets the NBEGM config, and the retiree and tied builders carry the savings-form budget (`resources`, `savings`) that the solver's contract reads -- without it those regimes cannot be built at all. The build-time affinity and interval-constancy probes cannot run on the added regimes: they differentiate the budget on scalar inputs and `assets_and_income.capital_income` declares `rate_of_return: ScalarFloat`, which the probe's one-element array violates. The probes fall back to `assume_declared` and warn, so NB-EGM's exactness precondition is asserted rather than checked there, and the solve needs validating against an independent reference. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_018yHhuFzqsDw1MhdB1i2Ljm --- src/aca_model/baseline/regimes/__init__.py | 7 +-- src/aca_model/baseline/regimes/_retiree.py | 9 +++- src/aca_model/baseline/regimes/_tied.py | 9 +++- tests/test_nbegm_model_creation.py | 53 ++++++++++++---------- 4 files changed, 47 insertions(+), 31 deletions(-) diff --git a/src/aca_model/baseline/regimes/__init__.py b/src/aca_model/baseline/regimes/__init__.py index a0c5543..b64543e 100644 --- a/src/aca_model/baseline/regimes/__init__.py +++ b/src/aca_model/baseline/regimes/__init__.py @@ -36,11 +36,6 @@ from aca_model.baseline.regimes._nbegm import build_nbegm_solver from aca_model.config import GridConfig -# NBEGM is a per-regime (not global) solver: it solves one 1-D -# consumption/savings regime with at most one discrete action, so it attaches -# only to the M1 vertical-slice regime. -_NBEGM_REGIME = "nongroup_nomc_inelig_canwork" - __all__ = [ "REGIME_SPECS", "RegimeId", @@ -115,7 +110,7 @@ def build_all_regimes( name, grids, dcegm_solver=dcegm_solver, - nbegm_solver=nbegm_solver if name == _NBEGM_REGIME else None, + nbegm_solver=nbegm_solver, ) regimes["dead"] = build_dead_regime(solver=solver) return regimes diff --git a/src/aca_model/baseline/regimes/_retiree.py b/src/aca_model/baseline/regimes/_retiree.py index 3ac87d2..76e19ef 100644 --- a/src/aca_model/baseline/regimes/_retiree.py +++ b/src/aca_model/baseline/regimes/_retiree.py @@ -20,6 +20,7 @@ build_actions, build_common_functions, build_granular_regime_transition, + build_nbegm_functions, build_pension_functions, build_regime_probs, build_state_transitions, @@ -133,6 +134,12 @@ def build_regime( if egm_solver is None else ("nbegm" if nbegm_solver is not None else "dcegm") ) + functions = _build_functions(spec) + if nbegm_solver is not None: + # NBEGM's solver contract is stated per regime: it reads the budget in + # savings form off `resources` and the post-decision node off + # `savings`, neither of which the brute-force build needs. + functions = {**functions, **build_nbegm_functions()} return Regime( transition=build_granular_regime_transition( transition_func=transition_func, target_ids=(*own.values(), *ng.values()) @@ -141,6 +148,6 @@ def build_regime( states=states, state_transitions=build_state_transitions(spec, solver=state_solver), actions=build_actions(spec, grids), - functions=_build_functions(spec), + functions=functions, **solver_kwargs, ) diff --git a/src/aca_model/baseline/regimes/_tied.py b/src/aca_model/baseline/regimes/_tied.py index 08a2196..348c34a 100644 --- a/src/aca_model/baseline/regimes/_tied.py +++ b/src/aca_model/baseline/regimes/_tied.py @@ -21,6 +21,7 @@ build_actions, build_common_functions, build_granular_regime_transition, + build_nbegm_functions, build_pension_functions, build_regime_probs, build_state_transitions, @@ -103,6 +104,12 @@ def build_regime( if egm_solver is None else ("nbegm" if nbegm_solver is not None else "dcegm") ) + functions = _build_functions(spec) + if nbegm_solver is not None: + # NBEGM's solver contract is stated per regime: it reads the budget in + # savings form off `resources` and the post-decision node off + # `savings`, neither of which the brute-force build needs. + functions = {**functions, **build_nbegm_functions()} return Regime( transition=build_granular_regime_transition( transition_func=transition_func, target_ids=(*own.values(), *ng.values()) @@ -111,6 +118,6 @@ def build_regime( states=states, state_transitions=build_state_transitions(spec, solver=state_solver), actions=build_actions(spec, grids), - functions=_build_functions(spec), + functions=functions, **solver_kwargs, ) diff --git a/tests/test_nbegm_model_creation.py b/tests/test_nbegm_model_creation.py index a8ed368..8efde5d 100644 --- a/tests/test_nbegm_model_creation.py +++ b/tests/test_nbegm_model_creation.py @@ -1,10 +1,9 @@ -"""NBEGM solver wiring: `solver="nbegm"` is a per-regime option. +"""NBEGM solver wiring. -Unlike DC-EGM (a global Euler solver on every living regime), NBEGM solves a -single 1-D consumption/savings regime, so it attaches only to the vertical-slice -regime `nongroup_nomc_inelig_canwork`; every other living regime keeps brute -force. The savings-form spec is shared with DC-EGM (NBEGM's budget node is -`resources`, the post-decision function is `savings`). +`solver="nbegm"` solves every living regime with NB-EGM, so each of them carries +the savings-form budget the solver reads: the budget node is `resources` and the +post-decision function is `savings`, the same spec DC-EGM uses. The `dead` regime +is terminal and keeps its own solver. The regime declares every choice its structure affords — whether to buy non-group coverage, and how many hours to work — and the discrete envelope @@ -21,7 +20,7 @@ from helpers.model import _DERIVED_CATEGORICALS # ty: ignore[unresolved-import] from lcm import DiscreteGrid, Model, Regime from lcm.exceptions import RegimeInitializationError -from lcm.solvers import NBEGM, GridSearch +from lcm.solvers import NBEGM from aca_model.aca import PolicyVariant from aca_model.aca.model import create_model as create_aca_model @@ -89,15 +88,18 @@ def _grids() -> Grids: ) -def test_nbegm_attaches_only_to_the_m1_regime() -> None: - """`solver="nbegm"` gives the M1 slice regime a `NBEGM` config and leaves - every other living regime on brute force.""" +def test_nbegm_attaches_to_every_living_regime() -> None: + """`solver="nbegm"` solves every living regime with NB-EGM. + + A solver that reached only some regimes would leave the rest on brute + force while still reporting itself as the model's solver, so the choice of + solver would not be visible in the result it produced. + """ regimes = _build_regimes("nbegm") - assert isinstance(regimes[_M1_REGIME].solver, NBEGM) - for name in REGIME_SPECS: - if name == _M1_REGIME: - continue - assert isinstance(regimes[name].solver, GridSearch), name + on_brute_force = [ + name for name in REGIME_SPECS if not isinstance(regimes[name].solver, NBEGM) + ] + assert on_brute_force == [] def test_build_nbegm_solver_uses_the_savings_form_resources_budget() -> None: @@ -188,15 +190,20 @@ def test_nbegm_m1_regime_takes_the_savings_form_assets_laws() -> None: assert law is expected, target_name -def test_nbegm_savings_form_functions_are_scoped_to_the_m1_regime() -> None: - """Under NBEGM only the M1 regime carries the savings-form budget functions - (`resources`, `savings`); brute regimes keep the cash-on-hand form and carry - neither.""" +def test_nbegm_gives_every_living_regime_the_savings_form_budget() -> None: + """Under NBEGM every living regime carries `resources` and `savings`. + + They are the solver's budget contract, so a regime NB-EGM solves without + them cannot be built at all; a regime it does not solve has no use for + them, which is why the brute-force build omits them. + """ model = _build_model("nbegm") - m1_functions = model.user_regimes[_M1_REGIME].functions - assert "resources" in m1_functions - assert "savings" in m1_functions - assert "resources" not in model.user_regimes[_BRUTE_REGIME].functions + missing = [ + name + for name in REGIME_SPECS + if not {"resources", "savings"} <= set(model.user_regimes[name].functions) + ] + assert missing == [] def test_nbegm_m1_regime_does_not_carry_inverse_marginal_utility() -> None: From ace111bcdf1e64162bc01f642e54490302d915cd Mon Sep 17 00:00:00 2001 From: Hans-Martin von Gaudecker Date: Sun, 16 Aug 2026 12:52:28 +0200 Subject: [PATCH 04/16] Let pylcm's NBEGM probes check the ACA budget The probes run on the first solve against the model's complete parameter vector, so the tax tables and threshold schedules the budget reads are the model's own rather than synthesized stand-ins. Affinity and constancy are checked rather than asserted. --- src/aca_model/baseline/regimes/_nbegm.py | 5 ----- 1 file changed, 5 deletions(-) diff --git a/src/aca_model/baseline/regimes/_nbegm.py b/src/aca_model/baseline/regimes/_nbegm.py index d1cda4f..8f8685b 100644 --- a/src/aca_model/baseline/regimes/_nbegm.py +++ b/src/aca_model/baseline/regimes/_nbegm.py @@ -40,11 +40,6 @@ def build_nbegm_solver(grids: Grids) -> NBEGM: continuous_state="assets", budget_target="resources", post_decision_function="savings", - # The budget DAG mixes scalar- and array-annotated parameters (tax - # tables, threshold schedules) that the build-time derivative probes - # cannot synthesize, so affinity/constancy is asserted here and - # validated by the full-model brute-agreement gates. - probe_failure="assume_declared", # Splay the child stochastic-node expectation per the grid config: `0` (the # default) reads the whole node mesh in one pass on a memory-rich device; a # positive value loops it in blocks to fit a tighter budget (a CPU run). From 4f258186c71304dc0bf2f102b299f5d062e263cb Mon Sep 17 00:00:00 2001 From: Hans-Martin von Gaudecker Date: Wed, 19 Aug 2026 16:14:19 +0200 Subject: [PATCH 05/16] Declare the liquid Euler margin on the EGM regimes pylcm takes the liquid roles from the regime: a `ConsumptionSavingsRegime` names the liquid state, the action paid from resources, the resources node and the post-decision state, and the solver carries numerical configuration only. `ACA_LIQUID_MARGIN` states that once for the three living-regime builders, which now route through `build_alive_regime`; a brute-force regime has no `resources` or `savings` node to name and stays a plain `Regime`. The role assertions move with the declaration: they read `regime.liquid`, which is public API, rather than the bound solver's attributes. --- src/aca_model/baseline/regimes/_common.py | 31 ++++++++++++++++ src/aca_model/baseline/regimes/_dcegm.py | 4 -- src/aca_model/baseline/regimes/_nbegm.py | 9 ++--- src/aca_model/baseline/regimes/_nongroup.py | 6 +-- src/aca_model/baseline/regimes/_retiree.py | 6 +-- src/aca_model/baseline/regimes/_tied.py | 6 +-- tests/test_dcegm_model_creation.py | 21 +++++++---- tests/test_nbegm_model_creation.py | 41 +++++++++++++++------ 8 files changed, 86 insertions(+), 38 deletions(-) diff --git a/src/aca_model/baseline/regimes/_common.py b/src/aca_model/baseline/regimes/_common.py index eddfa52..5f24287 100644 --- a/src/aca_model/baseline/regimes/_common.py +++ b/src/aca_model/baseline/regimes/_common.py @@ -13,9 +13,11 @@ import jax.numpy as jnp import numpy as np from lcm import ( + ConsumptionSavingsRegime, DiscreteGrid, IrregSpacedGrid, LinSpacedGrid, + LiquidMargin, MarkovTransition, NormalIIDProcess, Phased, @@ -26,6 +28,7 @@ categorical, fixed_transition, ) +from lcm.solvers import OneMarginSolver from lcm.typing import BoolND, FloatND, IntND, RegimeName, ScalarInt, UserParams from aca_model.agent import ( @@ -44,6 +47,34 @@ SolverName = Literal["brute_force", "dcegm", "nbegm"] +# The liquid Euler margin every EGM-solved ACA regime shares. `resources` is +# post-transfer cash-on-hand (`max(cash_on_hand, floor)`) and `savings` is the +# post-decision assets node; both are supplied by the savings-form rewiring in +# `build_dcegm_functions` / `build_nbegm_functions`, so the margin is +# declarable only on a regime that carries them. +ACA_LIQUID_MARGIN = LiquidMargin( + state="assets", + action="consumption_dollars", + resources="resources", + post_decision_state="savings", +) + + +def build_alive_regime( + *, egm_solver: OneMarginSolver | None, **regime_kwargs: Any +) -> Regime: + """Build one alive regime, declaring the liquid margin when EGM solves it. + + A brute-force regime has no `resources` or `savings` node to name, so it + stays a plain `Regime`. An EGM-solved regime owns the four margin names + and hands them to the solver, which takes numerical configuration only. + """ + if egm_solver is None: + return Regime(**regime_kwargs) + return ConsumptionSavingsRegime( + solver=egm_solver, liquid=ACA_LIQUID_MARGIN, **regime_kwargs + ) + @categorical(ordered=False) class RegimeId: diff --git a/src/aca_model/baseline/regimes/_dcegm.py b/src/aca_model/baseline/regimes/_dcegm.py index b0ee129..76cf075 100644 --- a/src/aca_model/baseline/regimes/_dcegm.py +++ b/src/aca_model/baseline/regimes/_dcegm.py @@ -35,10 +35,6 @@ def build_dcegm_solver(grids: Grids) -> DCEGM: batch_size=grids.grid_config.n_savings_batch_size, ) return DCEGM( - continuous_state="assets", - continuous_action="consumption_dollars", - resources="resources", - post_decision_function="savings", savings_grid=savings_grid, stochastic_node_batch_size=grids.grid_config.n_stochastic_node_batch_size, ) diff --git a/src/aca_model/baseline/regimes/_nbegm.py b/src/aca_model/baseline/regimes/_nbegm.py index 8f8685b..61c656a 100644 --- a/src/aca_model/baseline/regimes/_nbegm.py +++ b/src/aca_model/baseline/regimes/_nbegm.py @@ -23,9 +23,9 @@ def build_nbegm_solver(grids: Grids) -> NBEGM: The savings grid mirrors DC-EGM's: lower bound 0 (the borrowing constraint in post-decision form), upper bound the assets span, cubically clustered - toward the constraint. The budget node is `resources` (post-floor - cash-on-hand) and the post-decision function is `savings`, matching the - shared savings-form spec. + toward the constraint. Which DAG nodes play the liquid roles is the + regime's declaration (`ACA_LIQUID_MARGIN`), not the solver's; this config + carries numerical settings only. """ n_points = grids.grid_config.n_savings_gridpoints _fail_if_too_few_savings_gridpoints(n_points) @@ -37,9 +37,6 @@ def build_nbegm_solver(grids: Grids) -> NBEGM: ) return NBEGM( savings_grid=savings_grid, - continuous_state="assets", - budget_target="resources", - post_decision_function="savings", # Splay the child stochastic-node expectation per the grid config: `0` (the # default) reads the whole node mesh in one pass on a memory-rich device; a # positive value loops it in blocks to fit a tighter budget (a CPU run). diff --git a/src/aca_model/baseline/regimes/_nongroup.py b/src/aca_model/baseline/regimes/_nongroup.py index 5189cf9..f13e3c9 100644 --- a/src/aca_model/baseline/regimes/_nongroup.py +++ b/src/aca_model/baseline/regimes/_nongroup.py @@ -17,6 +17,7 @@ Grids, RegimeSpec, build_actions, + build_alive_regime, build_common_functions, build_granular_regime_transition, build_nbegm_functions, @@ -119,7 +120,6 @@ def build_regime( states = build_states(spec, grids) egm_solver = dcegm_solver if dcegm_solver is not None else nbegm_solver - solver_kwargs: dict = {} if egm_solver is None else {"solver": egm_solver} state_solver = ( "brute_force" if egm_solver is None @@ -135,7 +135,8 @@ def build_regime( # consumption by an argmax over the consumption grid and needs the # explicit feasibility mask. functions = {**functions, **build_nbegm_functions()} - return Regime( + return build_alive_regime( + egm_solver=egm_solver, transition=build_granular_regime_transition( transition_func=transition_func, target_ids=own.values() ), @@ -145,5 +146,4 @@ def build_regime( actions=build_actions(spec, grids), functions=functions, constraints=constraints, - **solver_kwargs, ) diff --git a/src/aca_model/baseline/regimes/_retiree.py b/src/aca_model/baseline/regimes/_retiree.py index 76e19ef..309a71c 100644 --- a/src/aca_model/baseline/regimes/_retiree.py +++ b/src/aca_model/baseline/regimes/_retiree.py @@ -18,6 +18,7 @@ Grids, RegimeSpec, build_actions, + build_alive_regime, build_common_functions, build_granular_regime_transition, build_nbegm_functions, @@ -128,7 +129,6 @@ def build_regime( states = build_states(spec, grids) egm_solver = dcegm_solver if dcegm_solver is not None else nbegm_solver - solver_kwargs: dict = {} if egm_solver is None else {"solver": egm_solver} state_solver = ( "brute_force" if egm_solver is None @@ -140,7 +140,8 @@ def build_regime( # savings form off `resources` and the post-decision node off # `savings`, neither of which the brute-force build needs. functions = {**functions, **build_nbegm_functions()} - return Regime( + return build_alive_regime( + egm_solver=egm_solver, transition=build_granular_regime_transition( transition_func=transition_func, target_ids=(*own.values(), *ng.values()) ), @@ -149,5 +150,4 @@ def build_regime( state_transitions=build_state_transitions(spec, solver=state_solver), actions=build_actions(spec, grids), functions=functions, - **solver_kwargs, ) diff --git a/src/aca_model/baseline/regimes/_tied.py b/src/aca_model/baseline/regimes/_tied.py index 348c34a..ff13664 100644 --- a/src/aca_model/baseline/regimes/_tied.py +++ b/src/aca_model/baseline/regimes/_tied.py @@ -19,6 +19,7 @@ Grids, RegimeSpec, build_actions, + build_alive_regime, build_common_functions, build_granular_regime_transition, build_nbegm_functions, @@ -98,7 +99,6 @@ def build_regime( states = build_states(spec, grids) egm_solver = dcegm_solver if dcegm_solver is not None else nbegm_solver - solver_kwargs: dict = {} if egm_solver is None else {"solver": egm_solver} state_solver = ( "brute_force" if egm_solver is None @@ -110,7 +110,8 @@ def build_regime( # savings form off `resources` and the post-decision node off # `savings`, neither of which the brute-force build needs. functions = {**functions, **build_nbegm_functions()} - return Regime( + return build_alive_regime( + egm_solver=egm_solver, transition=build_granular_regime_transition( transition_func=transition_func, target_ids=(*own.values(), *ng.values()) ), @@ -119,5 +120,4 @@ def build_regime( state_transitions=build_state_transitions(spec, solver=state_solver), actions=build_actions(spec, grids), functions=functions, - **solver_kwargs, ) diff --git a/tests/test_dcegm_model_creation.py b/tests/test_dcegm_model_creation.py index 2a4b8de..d46a825 100644 --- a/tests/test_dcegm_model_creation.py +++ b/tests/test_dcegm_model_creation.py @@ -12,7 +12,13 @@ import numpy as np import pytest from helpers.model import _DERIVED_CATEGORICALS # ty: ignore[unresolved-import] -from lcm import DiscreteGrid, IrregSpacedGrid, Model, Regime +from lcm import ( + ConsumptionSavingsRegime, + DiscreteGrid, + IrregSpacedGrid, + Model, + Regime, +) from lcm.solvers import DCEGM from aca_model.agent import assets_and_income @@ -56,14 +62,15 @@ def _build_model(solver: SolverName) -> Model: def test_every_living_regime_gets_the_dcegm_solver() -> None: - """`solver="dcegm"` attaches a `DCEGM` config with assets as the Euler - state to every living regime; the terminal regime keeps the default.""" + """`solver="dcegm"` attaches a `DCEGM` config to every living regime, each + declaring `assets` as its liquid margin; the terminal regime keeps the + default.""" regimes = _build_regimes("dcegm") for name in REGIME_SPECS: - solver = regimes[name].solver - assert isinstance(solver, DCEGM), name - assert solver.continuous_state == "assets" - assert solver.continuous_action == "consumption_dollars" + regime = cast("ConsumptionSavingsRegime", regimes[name]) + assert isinstance(regime.solver, DCEGM), name + assert regime.liquid.state == "assets", name + assert regime.liquid.action == "consumption_dollars", name assert not isinstance(regimes["dead"].solver, DCEGM) diff --git a/tests/test_nbegm_model_creation.py b/tests/test_nbegm_model_creation.py index 8efde5d..0502370 100644 --- a/tests/test_nbegm_model_creation.py +++ b/tests/test_nbegm_model_creation.py @@ -18,7 +18,13 @@ import pytest from helpers.model import _DERIVED_CATEGORICALS # ty: ignore[unresolved-import] -from lcm import DiscreteGrid, Model, Regime +from lcm import ( + ConsumptionSavingsRegime, + DiscreteGrid, + LiquidMargin, + Model, + Regime, +) from lcm.exceptions import RegimeInitializationError from lcm.solvers import NBEGM @@ -88,6 +94,13 @@ def _grids() -> Grids: ) +def _liquid_margins(regimes: Mapping[str, Regime]) -> set[LiquidMargin]: + """Return the liquid margin every living regime declares.""" + return { + cast("ConsumptionSavingsRegime", regimes[name]).liquid for name in REGIME_SPECS + } + + def test_nbegm_attaches_to_every_living_regime() -> None: """`solver="nbegm"` solves every living regime with NB-EGM. @@ -102,20 +115,24 @@ def test_nbegm_attaches_to_every_living_regime() -> None: assert on_brute_force == [] -def test_build_nbegm_solver_uses_the_savings_form_resources_budget() -> None: - """The NBEGM config inverts against `resources` in post-decision savings +def test_nbegm_regimes_declare_the_savings_form_resources_budget() -> None: + """Every NB-EGM regime inverts against `resources` in post-decision savings form, matching the DC-EGM contract the regime is rewired into.""" - solver = build_nbegm_solver(_grids()) - assert isinstance(solver, NBEGM) - assert solver.budget_target == "resources" - assert solver.post_decision_function == "savings" + declared = { + (margin.resources, margin.post_decision_state) + for margin in _liquid_margins(_build_regimes("nbegm")) + } + assert declared == {("resources", "savings")} -def test_build_nbegm_solver_names_assets_as_the_euler_axis() -> None: - """`assets` is the liquid (Euler) axis; `aime` and the stochastic shock grids - ride along, so the solver names the Euler axis explicitly.""" - solver = build_nbegm_solver(_grids()) - assert solver.continuous_state == "assets" +def test_nbegm_regimes_declare_assets_as_the_liquid_euler_axis() -> None: + """`assets` paid down by `consumption_dollars` is the liquid margin of every + NB-EGM regime; `aime` and the stochastic shock grids ride along.""" + declared = { + (margin.state, margin.action) + for margin in _liquid_margins(_build_regimes("nbegm")) + } + assert declared == {("assets", "consumption_dollars")} def test_build_nbegm_solver_forwards_the_jump_read_mode() -> None: From 221bf4f0c5f69390bb90412b9c13e335e9356cbe Mon Sep 17 00:00:00 2001 From: Hans-Martin von Gaudecker Date: Wed, 19 Aug 2026 18:07:15 +0200 Subject: [PATCH 06/16] Describe the borrowing constraint's actual solver coverage Both callers of `build_model_constraints` still said DC-EGM gets no broadcast constraint. It is declared under every solver: the EGM solve enforces the limit through the savings grid's lower bound, but forward simulation re-decides consumption by an argmax over the consumption grid and needs the explicit feasibility mask. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_018yHhuFzqsDw1MhdB1i2Ljm --- src/aca_model/baseline/regimes/__init__.py | 2 +- src/aca_model/baseline/regimes/_common.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/src/aca_model/baseline/regimes/__init__.py b/src/aca_model/baseline/regimes/__init__.py index b64543e..6e2d68c 100644 --- a/src/aca_model/baseline/regimes/__init__.py +++ b/src/aca_model/baseline/regimes/__init__.py @@ -131,7 +131,7 @@ def build_model_slots( Both the baseline and the ACA `create_model` consume this — the ACA overlay swaps only regime-level functions, so the broadcast slots are policy-invariant. Under DC-EGM the solver-contract functions join the - broadcast set and no borrowing constraint is declared. + broadcast set; the borrowing constraint is declared under every solver. """ grids = build_grids( grid_config=grid_config, diff --git a/src/aca_model/baseline/regimes/_common.py b/src/aca_model/baseline/regimes/_common.py index 5f24287..6de310d 100644 --- a/src/aca_model/baseline/regimes/_common.py +++ b/src/aca_model/baseline/regimes/_common.py @@ -567,8 +567,8 @@ def build_dead_regime(*, solver: SolverName = "brute_force") -> Regime: inputs (e.g. `pension_benefit`) don't surface as params in the dead template. - constraints: the borrowing constraint is masked — `dead` has no - consumption action. (Under DC-EGM no constraint is broadcast, so - there is nothing to mask.) + consumption action. It is broadcast under every solver, so there is + always exactly one mask to apply. - `pension_wealth` is masked explicitly: a carried state is rejected in terminal regimes before pruning could drop it. """ From 8466aacff7b654b9e2ba38ed51438c1b7c261e4f Mon Sep 17 00:00:00 2001 From: Hans-Martin von Gaudecker Date: Wed, 19 Aug 2026 18:28:46 +0200 Subject: [PATCH 07/16] Name the DC-EGM blocker that actually fires The xfail reason described the assets-law chain, but the build stops earlier: pylcm's DC-EGM refuses the broadcast borrowing constraint for reading continuous variables. Both blockers are now listed in the order they fire. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_018yHhuFzqsDw1MhdB1i2Ljm --- tests/test_dcegm_model_creation.py | 19 +++++++++++++------ 1 file changed, 13 insertions(+), 6 deletions(-) diff --git a/tests/test_dcegm_model_creation.py b/tests/test_dcegm_model_creation.py index d46a825..cbea60a 100644 --- a/tests/test_dcegm_model_creation.py +++ b/tests/test_dcegm_model_creation.py @@ -153,12 +153,19 @@ def test_benchmark_consumption_points_pin_both_floors() -> None: @pytest.mark.xfail( strict=False, reason=( - "pylcm's DC-EGM contract does not yet admit the ACA budget: the " - "assets law reaches `assets` outside the post-decision function — " - "through `oop_costs` (Medicaid eligibility → `countable_income` → " - "`capital_income`) and `pension_assets_adjustment` " - "(`marginal_tax_rate` → `gross_income` → `capital_income`). " - "Fixes land upstream in pylcm, not here." + "Two independent blockers, both upstream in pylcm. The one that " + "fires first: DC-EGM refuses the broadcast `borrowing_constraint` " + "because it reads the continuous `assets` and `consumption_dollars`, " + "while the solver enforces the borrowing limit through the savings " + "grid's lower bound instead. The model declares it under every " + "solver because forward simulation re-decides consumption by an " + "argmax over the consumption grid and needs the explicit mask, so " + "the two requirements conflict and neither side can drop its half " + "alone. Behind it: the assets law reaches `assets` outside the " + "post-decision function — through `oop_costs` (Medicaid eligibility " + "→ `countable_income` → `capital_income`) and " + "`pension_assets_adjustment` (`marginal_tax_rate` → `gross_income` " + "→ `capital_income`)." ), ) def test_dcegm_benchmark_model_builds() -> None: From a2ad9588e59744519a196e677852efef70658f27 Mon Sep 17 00:00:00 2001 From: Hans-Martin von Gaudecker Date: Thu, 20 Aug 2026 15:19:36 +0200 Subject: [PATCH 08/16] FIX Declare the ACA borrowing limit structurally --- src/aca_model/baseline/regimes/__init__.py | 2 +- src/aca_model/baseline/regimes/_common.py | 20 +++++++++++------ tests/test_dcegm_model_creation.py | 26 +--------------------- 3 files changed, 15 insertions(+), 33 deletions(-) diff --git a/src/aca_model/baseline/regimes/__init__.py b/src/aca_model/baseline/regimes/__init__.py index 6e2d68c..dbbb5aa 100644 --- a/src/aca_model/baseline/regimes/__init__.py +++ b/src/aca_model/baseline/regimes/__init__.py @@ -141,7 +141,7 @@ def build_model_slots( ) return { "functions": build_model_functions(solver=solver), - "constraints": build_model_constraints(), + "constraints": build_model_constraints(solver=solver), "states": build_model_states(grids), "state_transitions": build_model_state_transitions(), } diff --git a/src/aca_model/baseline/regimes/_common.py b/src/aca_model/baseline/regimes/_common.py index 6de310d..33805d7 100644 --- a/src/aca_model/baseline/regimes/_common.py +++ b/src/aca_model/baseline/regimes/_common.py @@ -27,6 +27,7 @@ RouwenhorstAR1Process, categorical, fixed_transition, + post_decision_lower_bound, ) from lcm.solvers import OneMarginSolver from lcm.typing import BoolND, FloatND, IntND, RegimeName, ScalarInt, UserParams @@ -577,7 +578,7 @@ def build_dead_regime(*, solver: SolverName = "brute_force") -> Regime: for name in build_model_functions(solver=solver) if name not in _DEAD_KEEPS } - constraint_masks = dict.fromkeys(build_model_constraints()) + constraint_masks = dict.fromkeys(build_model_constraints(solver=solver)) return Regime( transition=None, functions={"utility": preferences.bequest, **function_masks}, @@ -693,16 +694,21 @@ def build_nbegm_functions() -> dict: } -def build_model_constraints() -> dict: +def build_model_constraints(*, solver: SolverName) -> dict: """Build the model-level constraints broadcast into every regime. `dead` masks the borrowing constraint — it has no consumption action. - The constraint is broadcast under every solver: an EGM solve (DC-EGM or - NBEGM) enforces the borrowing limit through the savings grid's lower - bound, but forward simulation re-decides consumption by an argmax over - the consumption grid and needs the explicit feasibility mask. + Grid search evaluates the action-level predicate directly. An EGM-family + solve proves the equivalent post-decision lower bound from its savings + grid, while forward simulation receives the solver's intrinsic budget + mask over the consumption grid. """ - return {"borrowing_constraint": assets_and_income.borrowing_constraint} + borrowing_constraint = ( + assets_and_income.borrowing_constraint + if solver == "brute_force" + else post_decision_lower_bound(margin=ACA_LIQUID_MARGIN, lower=0.0) + ) + return {"borrowing_constraint": borrowing_constraint} def build_model_states(grids: Grids) -> dict: diff --git a/tests/test_dcegm_model_creation.py b/tests/test_dcegm_model_creation.py index cbea60a..5a73f23 100644 --- a/tests/test_dcegm_model_creation.py +++ b/tests/test_dcegm_model_creation.py @@ -150,32 +150,8 @@ def test_benchmark_consumption_points_pin_both_floors() -> None: np.testing.assert_allclose(points[:2], [floor, floor * 2.0**exponent], rtol=1e-12) -@pytest.mark.xfail( - strict=False, - reason=( - "Two independent blockers, both upstream in pylcm. The one that " - "fires first: DC-EGM refuses the broadcast `borrowing_constraint` " - "because it reads the continuous `assets` and `consumption_dollars`, " - "while the solver enforces the borrowing limit through the savings " - "grid's lower bound instead. The model declares it under every " - "solver because forward simulation re-decides consumption by an " - "argmax over the consumption grid and needs the explicit mask, so " - "the two requirements conflict and neither side can drop its half " - "alone. Behind it: the assets law reaches `assets` outside the " - "post-decision function — through `oop_costs` (Medicaid eligibility " - "→ `countable_income` → `capital_income`) and " - "`pension_assets_adjustment` (`marginal_tax_rate` → `gross_income` " - "→ `capital_income`)." - ), -) def test_dcegm_benchmark_model_builds() -> None: - """The benchmark model accepts `solver="dcegm"` end to end. - - The acceptance criterion for the upstream DC-EGM stack: once pylcm's - contract admits the ACA budget chains, this builds without error. The - construction-time consumption points are supplied so the build reaches - the upstream limitation rather than the missing-points guard. - """ + """The benchmark model accepts `solver="dcegm"` end to end.""" model = create_model( n_subjects=1, fixed_params=_FIXED_PARAMS, From 328e3710cfafb19c28ebcd46a27090238415fdd6 Mon Sep 17 00:00:00 2001 From: Hans-Martin von Gaudecker Date: Fri, 21 Aug 2026 13:41:24 +0200 Subject: [PATCH 09/16] Migrate ACA benchmark to breakpoint grids --- src/aca_model/baseline/regimes/_common.py | 32 ++++++++++------------- tests/test_aime_grid.py | 11 +++++--- tests/test_benchmark.py | 16 +++++++++++- tests/test_dcegm_model_creation.py | 9 ++----- tests/test_nbegm_model_creation.py | 9 ++----- 5 files changed, 41 insertions(+), 36 deletions(-) diff --git a/src/aca_model/baseline/regimes/_common.py b/src/aca_model/baseline/regimes/_common.py index 33805d7..3168213 100644 --- a/src/aca_model/baseline/regimes/_common.py +++ b/src/aca_model/baseline/regimes/_common.py @@ -13,20 +13,22 @@ import jax.numpy as jnp import numpy as np from lcm import ( - ConsumptionSavingsRegime, DiscreteGrid, + GridBreakpoint, IrregSpacedGrid, LinSpacedGrid, - LiquidMargin, MarkovTransition, NormalIIDProcess, Phased, - PiecewiseGridSegment, PiecewiseLinSpacedGrid, Regime, RouwenhorstAR1Process, categorical, fixed_transition, +) +from lcm.consumption_savings_regime import ( + ConsumptionSavingsRegime, + LiquidMargin, post_decision_lower_bound, ) from lcm.solvers import OneMarginSolver @@ -372,22 +374,16 @@ def _build_aime_grid( total is fixed by the PIA structure (`sum(_AIME_PIECE_N_POINTS)`). """ kinks = [float(k) for k in np.asarray(fixed_params["pia_aime_grid"])] - segments = ( - PiecewiseGridSegment( - interval=f"[{kinks[0]}, {kinks[1]})", n_points=_AIME_PIECE_N_POINTS[0] - ), - PiecewiseGridSegment( - interval=f"[{kinks[1]}, {kinks[2]})", n_points=_AIME_PIECE_N_POINTS[1] - ), - PiecewiseGridSegment( - interval=f"[{kinks[2]}, {kinks[3]})", n_points=_AIME_PIECE_N_POINTS[2] - ), - PiecewiseGridSegment( - interval=f"[{kinks[3]}, {kinks[4]}]", n_points=_AIME_PIECE_N_POINTS[3] - ), - ) return PiecewiseLinSpacedGrid( - segments=segments, batch_size=grid_config.n_aime_batch_size + start=kinks[0], + stop=kinks[4], + breakpoints=( + GridBreakpoint(value=kinks[1]), + GridBreakpoint(value=kinks[2]), + GridBreakpoint(value=kinks[3]), + ), + points_per_segment=_AIME_PIECE_N_POINTS, + batch_size=grid_config.n_aime_batch_size, ) diff --git a/tests/test_aime_grid.py b/tests/test_aime_grid.py index a249ed3..368fe37 100644 --- a/tests/test_aime_grid.py +++ b/tests/test_aime_grid.py @@ -18,12 +18,17 @@ _FIXED_PARAMS = MappingProxyType({"pia_aime_grid": _PIA_AIME_GRID}) -def test_build_aime_grid_has_four_segments() -> None: - """The AIME grid spans all four PIA intervals, including the extension.""" +def test_build_aime_grid_owns_each_pia_breakpoint_on_the_right() -> None: + """The AIME grid preserves four PIA pieces with right-owned interiors.""" grid = _build_aime_grid( grid_config=BENCHMARK_GRID_CONFIG, fixed_params=_FIXED_PARAMS ) - assert len(grid.segments) == 4 + np.testing.assert_allclose( + [point.value for point in grid.breakpoints], + _PIA_AIME_GRID[1:-1], + ) + assert tuple(point.owner for point in grid.breakpoints) == ("right",) * 3 + assert grid.points_per_segment == _AIME_PIECE_N_POINTS def test_build_aime_grid_top_point_is_extension_aime() -> None: diff --git a/tests/test_benchmark.py b/tests/test_benchmark.py index fb1a9d6..564e729 100644 --- a/tests/test_benchmark.py +++ b/tests/test_benchmark.py @@ -2,7 +2,7 @@ import numpy as np import pytest -from lcm import DiscreteGrid +from lcm import DiscreteGrid, GridBreakpoint, PiecewiseLinSpacedGrid from aca_model.agent.preferences import BenchmarkPrefType from aca_model.benchmark import ( @@ -12,6 +12,20 @@ ) +def test_benchmark_model_builds_with_the_current_pylcm_grid_api() -> None: + """The frozen benchmark model constructs with breakpoint-first AIME grids.""" + model = create_benchmark_model( + n_subjects=1, + pref_type_grid=DiscreteGrid(BenchmarkPrefType), + ) + + aime = model.user_regimes["retiree_nomc_inelig_canwork"].states["aime"] + assert isinstance(aime, PiecewiseLinSpacedGrid) + assert all(isinstance(point, GridBreakpoint) for point in aime.breakpoints) + assert all(point.owner == "right" for point in aime.breakpoints) + assert aime.n_points == sum(aime.points_per_segment) + + @pytest.mark.long_running def test_benchmark_model_simulates_end_to_end() -> None: n_subjects = 20 diff --git a/tests/test_dcegm_model_creation.py b/tests/test_dcegm_model_creation.py index 5a73f23..492fdd2 100644 --- a/tests/test_dcegm_model_creation.py +++ b/tests/test_dcegm_model_creation.py @@ -12,13 +12,8 @@ import numpy as np import pytest from helpers.model import _DERIVED_CATEGORICALS # ty: ignore[unresolved-import] -from lcm import ( - ConsumptionSavingsRegime, - DiscreteGrid, - IrregSpacedGrid, - Model, - Regime, -) +from lcm import DiscreteGrid, IrregSpacedGrid, Model, Regime +from lcm.consumption_savings_regime import ConsumptionSavingsRegime from lcm.solvers import DCEGM from aca_model.agent import assets_and_income diff --git a/tests/test_nbegm_model_creation.py b/tests/test_nbegm_model_creation.py index 0502370..ce70017 100644 --- a/tests/test_nbegm_model_creation.py +++ b/tests/test_nbegm_model_creation.py @@ -18,13 +18,8 @@ import pytest from helpers.model import _DERIVED_CATEGORICALS # ty: ignore[unresolved-import] -from lcm import ( - ConsumptionSavingsRegime, - DiscreteGrid, - LiquidMargin, - Model, - Regime, -) +from lcm import DiscreteGrid, Model, Regime +from lcm.consumption_savings_regime import ConsumptionSavingsRegime, LiquidMargin from lcm.exceptions import RegimeInitializationError from lcm.solvers import NBEGM From f61cff94fa02a075bb075194b8b983ffe66641ad Mon Sep 17 00:00:00 2001 From: Hans-Martin von Gaudecker Date: Sun, 30 Aug 2026 17:18:36 +0200 Subject: [PATCH 10/16] ENH Forward NB-EGM workspace controls --- src/aca_model/baseline/regimes/_nbegm.py | 8 ++++++++ src/aca_model/config.py | 8 ++++++++ tests/test_nbegm_model_creation.py | 26 ++++++++++++++++++++++++ 3 files changed, 42 insertions(+) diff --git a/src/aca_model/baseline/regimes/_nbegm.py b/src/aca_model/baseline/regimes/_nbegm.py index 61c656a..3bd740b 100644 --- a/src/aca_model/baseline/regimes/_nbegm.py +++ b/src/aca_model/baseline/regimes/_nbegm.py @@ -50,12 +50,20 @@ def build_nbegm_solver(grids: Grids) -> NBEGM: # "certified" is exact and can abstain, "ordinary" reads in the working # format at a fraction of the cost. envelope_arithmetic=grids.grid_config.nbegm_envelope_arithmetic, + # Stream continuation intervals, or let the active byte planner choose + # when the configured width is zero. + interval_batch_size=grids.grid_config.n_nbegm_interval_batch_size, # Stream both ride-along cores over ride-cell blocks per the grid config; # `0` vmaps the whole flattened mesh at once. cell_block_size=grids.grid_config.n_nbegm_cell_block_size, # Stream the discrete-action branch axis in blocks per the grid config; # `0` runs the whole axis in one vectorized pass. branch_batch_size=grids.grid_config.n_nbegm_branch_batch_size, + # Compile-only per-device memory preflight; None preserves the legacy + # manual block path. + max_device_workspace_bytes=( + grids.grid_config.n_nbegm_max_device_workspace_bytes + ), # Cliff-read mode: exact one-sided limits (default) or the fast bridged # read for inner estimation loops (see `GridConfig.nbegm_jump_read`). jump_read=grids.grid_config.nbegm_jump_read, diff --git a/src/aca_model/config.py b/src/aca_model/config.py index db5ed3e..38c6d64 100644 --- a/src/aca_model/config.py +++ b/src/aca_model/config.py @@ -107,6 +107,14 @@ class GridConfig: # either way — the knob trades peak device memory against a sequential scan. # Only consulted under `solver="nbegm"`. n_nbegm_envelope_segment_block_size: int = 0 + # Compiled batch width for the liquid-interval axis. With an active byte + # budget, 0 lets the planner choose up to the full axis; a positive value + # caps it. Only consulted under `solver="nbegm"`. + n_nbegm_interval_batch_size: int = 0 + # Authoritative busiest-device workspace budget. None preserves manual + # blocks; a positive value activates compile-only continuation/envelope + # planning before backward induction. Only consulted under `solver="nbegm"`. + n_nbegm_max_device_workspace_bytes: int | None = None # Which arithmetic decides ownership in the merged upper envelope: # - "certified" (default) — candidates are compared in double-double precision # and no winner is published where none is separated, so the reported owner is diff --git a/tests/test_nbegm_model_creation.py b/tests/test_nbegm_model_creation.py index ce70017..b47af6f 100644 --- a/tests/test_nbegm_model_creation.py +++ b/tests/test_nbegm_model_creation.py @@ -352,3 +352,29 @@ def test_nbegm_aca_variants_leave_no_free_buy_private_params( if isinstance(params, dict) and "buy_private" in params ] assert offenders == [], offenders + + +def test_nbegm_workspace_budget_and_all_streaming_axes_are_forwarded() -> None: + """ACA exposes every planner axis and the positive per-device byte budget.""" + grid_config = dataclasses.replace( + BENCHMARK_GRID_CONFIG, + n_nbegm_stochastic_node_batch_size=2, + n_nbegm_envelope_segment_block_size=3, + n_nbegm_interval_batch_size=4, + n_nbegm_cell_block_size=5, + n_nbegm_branch_batch_size=6, + n_nbegm_max_device_workspace_bytes=72 * 1024**3, + ) + grids = build_grids( + grid_config=grid_config, + fixed_params=_FIXED_PARAMS, + wage_params=_WAGE_PARAMS, + pref_type_grid=DiscreteGrid(BenchmarkPrefType), + ) + solver = build_nbegm_solver(grids) + assert solver.stochastic_node_batch_size == 2 + assert solver.envelope_segment_block_size == 3 + assert solver.interval_batch_size == 4 + assert solver.cell_block_size == 5 + assert solver.branch_batch_size == 6 + assert solver.max_device_workspace_bytes == 72 * 1024**3 From 5d1d4797474946dafc9509f9d83907f31dcb6dc6 Mon Sep 17 00:00:00 2001 From: Hans-Martin von Gaudecker Date: Wed, 2 Sep 2026 01:29:18 +0200 Subject: [PATCH 11/16] ENH Allow explicit fp32 ACA runs --- src/aca_model/__init__.py | 9 ++++++++- tests/test_precision.py | 28 ++++++++++++++++++++++++++++ 2 files changed, 36 insertions(+), 1 deletion(-) create mode 100644 tests/test_precision.py diff --git a/src/aca_model/__init__.py b/src/aca_model/__init__.py index 1fc43f7..d26e6e4 100644 --- a/src/aca_model/__init__.py +++ b/src/aca_model/__init__.py @@ -1,6 +1,13 @@ +import os + import jax -jax.config.update("jax_enable_x64", True) +_x64_requested = os.environ.get("ACA_JAX_ENABLE_X64", "1").lower() not in { + "0", + "false", + "no", +} +jax.config.update("jax_enable_x64", _x64_requested) # Import lcm before installing the claw so its `_jaxtyping_patch` (picklable # jaxtyping sentinel) and `MappingProxyType` pytree registration are in place. diff --git a/tests/test_precision.py b/tests/test_precision.py new file mode 100644 index 0000000..e807bd3 --- /dev/null +++ b/tests/test_precision.py @@ -0,0 +1,28 @@ +"""Process-level floating-point precision selection.""" + +import os +import subprocess +import sys + + +def test_aca_precision_environment_selects_fp32() -> None: + """`ACA_JAX_ENABLE_X64=0` makes ACA computations use 32-bit floats.""" + env = { + **os.environ, + "JAX_ENABLE_X64": "1", + "ACA_JAX_ENABLE_X64": "0", + } + + result = subprocess.run( + [ + sys.executable, + "-c", + "import aca_model, jax; print(jax.config.jax_enable_x64)", + ], + check=True, + capture_output=True, + env=env, + text=True, + ) + + assert result.stdout.strip() == "False" From 94977a39172a8b32bbf620a7e289a05cf5c20e16 Mon Sep 17 00:00:00 2001 From: Hans-Martin von Gaudecker Date: Wed, 2 Sep 2026 02:06:34 +0200 Subject: [PATCH 12/16] FIX Gate the experimental NB-EGM planner --- src/aca_model/baseline/regimes/_nbegm.py | 21 +++++++++++++++------ tests/test_nbegm_model_creation.py | 23 +++++++++++++++++++---- 2 files changed, 34 insertions(+), 10 deletions(-) diff --git a/src/aca_model/baseline/regimes/_nbegm.py b/src/aca_model/baseline/regimes/_nbegm.py index 3bd740b..89dcfa0 100644 --- a/src/aca_model/baseline/regimes/_nbegm.py +++ b/src/aca_model/baseline/regimes/_nbegm.py @@ -12,6 +12,8 @@ config. """ +import dataclasses + from lcm import IrregSpacedGrid from lcm.solvers import NBEGM @@ -35,7 +37,7 @@ def build_nbegm_solver(grids: Grids) -> NBEGM: points=tuple(savings_stop * (i / (n_points - 1)) ** 3 for i in range(n_points)), batch_size=grids.grid_config.n_savings_batch_size, ) - return NBEGM( + solver = NBEGM( savings_grid=savings_grid, # Splay the child stochastic-node expectation per the grid config: `0` (the # default) reads the whole node mesh in one pass on a memory-rich device; a @@ -59,15 +61,22 @@ def build_nbegm_solver(grids: Grids) -> NBEGM: # Stream the discrete-action branch axis in blocks per the grid config; # `0` runs the whole axis in one vectorized pass. branch_batch_size=grids.grid_config.n_nbegm_branch_batch_size, - # Compile-only per-device memory preflight; None preserves the legacy - # manual block path. - max_device_workspace_bytes=( - grids.grid_config.n_nbegm_max_device_workspace_bytes - ), # Cliff-read mode: exact one-sided limits (default) or the fast bridged # read for inner estimation loops (see `GridConfig.nbegm_jump_read`). jump_read=grids.grid_config.nbegm_jump_read, ) + budget = grids.grid_config.n_nbegm_max_device_workspace_bytes + if budget is None: + return solver + + fields = {field.name for field in dataclasses.fields(NBEGM)} + if "max_device_workspace_bytes" not in fields: + msg = ( + "GridConfig.n_nbegm_max_device_workspace_bytes requires a pylcm " + "build with the experimental NB-EGM workspace planner." + ) + raise RuntimeError(msg) + return dataclasses.replace(solver, max_device_workspace_bytes=budget) def _fail_if_too_few_savings_gridpoints(n_savings_gridpoints: int) -> None: diff --git a/tests/test_nbegm_model_creation.py b/tests/test_nbegm_model_creation.py index b47af6f..fdd3a4b 100644 --- a/tests/test_nbegm_model_creation.py +++ b/tests/test_nbegm_model_creation.py @@ -354,8 +354,8 @@ def test_nbegm_aca_variants_leave_no_free_buy_private_params( assert offenders == [], offenders -def test_nbegm_workspace_budget_and_all_streaming_axes_are_forwarded() -> None: - """ACA exposes every planner axis and the positive per-device byte budget.""" +def test_nbegm_all_streaming_axes_are_forwarded() -> None: + """ACA forwards every streaming axis supported by pylcm.""" grid_config = dataclasses.replace( BENCHMARK_GRID_CONFIG, n_nbegm_stochastic_node_batch_size=2, @@ -363,7 +363,6 @@ def test_nbegm_workspace_budget_and_all_streaming_axes_are_forwarded() -> None: n_nbegm_interval_batch_size=4, n_nbegm_cell_block_size=5, n_nbegm_branch_batch_size=6, - n_nbegm_max_device_workspace_bytes=72 * 1024**3, ) grids = build_grids( grid_config=grid_config, @@ -377,4 +376,20 @@ def test_nbegm_workspace_budget_and_all_streaming_axes_are_forwarded() -> None: assert solver.interval_batch_size == 4 assert solver.cell_block_size == 5 assert solver.branch_batch_size == 6 - assert solver.max_device_workspace_bytes == 72 * 1024**3 + + +def test_nbegm_workspace_budget_requires_the_experimental_planner() -> None: + """A planner budget is refused when pylcm has no workspace planner.""" + grid_config = dataclasses.replace( + BENCHMARK_GRID_CONFIG, + n_nbegm_max_device_workspace_bytes=72 * 1024**3, + ) + grids = build_grids( + grid_config=grid_config, + fixed_params=_FIXED_PARAMS, + wage_params=_WAGE_PARAMS, + pref_type_grid=DiscreteGrid(BenchmarkPrefType), + ) + + with pytest.raises(RuntimeError, match="experimental NB-EGM workspace planner"): + build_nbegm_solver(grids) From 3af83d075feb0619a75169a2da957282df36bb1f Mon Sep 17 00:00:00 2001 From: Hans-Martin von Gaudecker Date: Thu, 3 Sep 2026 09:04:28 +0200 Subject: [PATCH 13/16] MAINT Migrate the tests to pylcm's SolutionResult solve API pylcm's `Model.solve` returns a `SolutionResult` and `Model.simulate` takes `solution=` in place of `period_to_regime_to_V_arr=`. The value mapping these tests read is `SolutionResult.values`; passing `None` for the old keyword meant "solve inside simulate", which is the new default. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01Y7QCx9tkGD3TqwNw681QwD --- tests/test_benchmark.py | 3 --- tests/test_dcegm_parity.py | 5 ++--- tests/test_nbegm_labor_live_validation.py | 2 +- tests/test_nbegm_solve_validation.py | 2 +- 4 files changed, 4 insertions(+), 8 deletions(-) diff --git a/tests/test_benchmark.py b/tests/test_benchmark.py index 564e729..ac9b525 100644 --- a/tests/test_benchmark.py +++ b/tests/test_benchmark.py @@ -41,7 +41,6 @@ def test_benchmark_model_simulates_end_to_end() -> None: result = model.simulate( params=params, initial_conditions=initial_conditions, - period_to_regime_to_V_arr=None, log_level="off", ) @@ -76,7 +75,6 @@ def test_benchmark_panel_exposes_hic_premium_and_wage_targets() -> None: result = model.simulate( params=params, initial_conditions=initial_conditions, - period_to_regime_to_V_arr=None, log_level="off", ) @@ -139,7 +137,6 @@ def test_benchmark_simulate_obeys_borrowing_constraint() -> None: result = model.simulate( params=params, initial_conditions=initial_conditions, - period_to_regime_to_V_arr=None, log_level="off", ) diff --git a/tests/test_dcegm_parity.py b/tests/test_dcegm_parity.py index 49e1695..4c83e2c 100644 --- a/tests/test_dcegm_parity.py +++ b/tests/test_dcegm_parity.py @@ -134,7 +134,6 @@ def _seeded_panel(model: Model) -> pd.DataFrame: result = model.simulate( params=params, initial_conditions=initial_conditions, - period_to_regime_to_V_arr=None, log_level="off", ) df = result.to_dataframe() @@ -235,7 +234,7 @@ def test_dcegm_solves_the_parity_model() -> None: (regime, period) cell carries a finite value-function array.""" model = _make_model(solver="dcegm", grid_config=PARITY_GRID_CONFIG) _, _, params = get_benchmark_params(model=model) - period_to_regime_to_v = model.solve(params=params, log_level="off") - for regime_to_v in period_to_regime_to_v.values(): + solution = model.solve(params=params, log_level="off") + for regime_to_v in solution.values.values(): for v_arr in regime_to_v.values(): assert bool(np.isfinite(np.asarray(v_arr)).any()) diff --git a/tests/test_nbegm_labor_live_validation.py b/tests/test_nbegm_labor_live_validation.py index a0ef0f0..53a732c 100644 --- a/tests/test_nbegm_labor_live_validation.py +++ b/tests/test_nbegm_labor_live_validation.py @@ -58,7 +58,7 @@ def _solve_m1(solver: SolverName) -> tuple[dict[int, np.ndarray], int]: ).index("assets") return { period: np.asarray(regimes[_M1_REGIME]) - for period, regimes in solution.items() + for period, regimes in solution.values.items() if _M1_REGIME in regimes }, assets_axis diff --git a/tests/test_nbegm_solve_validation.py b/tests/test_nbegm_solve_validation.py index 028fce7..af16729 100644 --- a/tests/test_nbegm_solve_validation.py +++ b/tests/test_nbegm_solve_validation.py @@ -47,7 +47,7 @@ def _solve_m1(solver: SolverName) -> dict[int, np.ndarray]: solution = model.solve(params=params, log_level="off") return { period: np.asarray(regimes[_M1_REGIME]) - for period, regimes in solution.items() + for period, regimes in solution.values.items() if _M1_REGIME in regimes } From b941507c14932b85b482d77ed8d7f3034c5dbfee Mon Sep 17 00:00:00 2001 From: Hans-Martin von Gaudecker Date: Tue, 8 Sep 2026 20:37:37 +0200 Subject: [PATCH 14/16] ENH Move ACA execution controls to Model policy Forward ExecutionConfig through baseline, ACA and benchmark factories. Derive omitted accelerator budgets from selected JAX allocator limits and reject retired GridConfig execution fields. Preserve economic grids and numerical solver settings. Validation against pylcm f6d41d54: exact ASV CPU preflight, focused compatibility tests, Ruff, ty and repository hooks pass. DC-EGM construction still exposes an upstream continuation-template placement failure on three CPU devices. --- README.md | 40 ++++++ src/aca_model/aca/model.py | 12 +- src/aca_model/baseline/model.py | 12 +- src/aca_model/baseline/regimes/_common.py | 25 +--- src/aca_model/baseline/regimes/_dcegm.py | 6 +- src/aca_model/baseline/regimes/_nbegm.py | 63 +-------- src/aca_model/benchmark.py | 13 +- src/aca_model/config.py | 150 ++-------------------- src/aca_model/execution.py | 53 ++++++++ tests/test_aime_grid.py | 9 +- tests/test_beartype_claw.py | 4 +- tests/test_benchmark.py | 28 +++- tests/test_dcegm_model_creation.py | 38 +----- tests/test_execution_budget.py | 62 +++++++++ tests/test_execution_config.py | 117 +++++++++++++++++ tests/test_grid_config.py | 35 +++++ tests/test_model_creation.py | 40 +----- tests/test_nbegm_labor_live_validation.py | 16 +-- tests/test_nbegm_model_creation.py | 43 +------ tests/test_nbegm_solve_validation.py | 2 +- tests/test_social_security.py | 2 +- tests/test_ss_benefit_integration.py | 2 +- 22 files changed, 403 insertions(+), 369 deletions(-) create mode 100644 src/aca_model/execution.py create mode 100644 tests/test_execution_budget.py create mode 100644 tests/test_execution_config.py create mode 100644 tests/test_grid_config.py diff --git a/README.md b/README.md index 5f52c1a..9edf892 100644 --- a/README.md +++ b/README.md @@ -1,3 +1,43 @@ # aca-model Core lifecycle model for the ACA structural retirement project. + +Model factories accept `execution_config=ExecutionConfig(...)`. An omitted policy uses +the smallest JAX allocator `bytes_limit` among the selected accelerators as an explicit +per-device budget. This enables budget-aware width selection; CPU construction remains +unbudgeted. Missing accelerator limits require an explicit policy instead of silently +falling back to unbudgeted execution. + +To combine a measured budget with device selection and sharding: + +```python +from dataclasses import replace + +from aca_model.execution import execution_config_for_devices + +execution_config = replace( + execution_config_for_devices(devices=(0, 1, 2)), + sharded_states=("pref_type",), +) +``` + +Pass this policy to the baseline, ACA or benchmark factory. Record its exact +`device_memory_bytes`, the selected devices' allocator observations, executed precision +and selected program widths with performance evidence. The budget is an allocator +ceiling; it does not establish that all workload phases fit. + +Explicit policies pass through unchanged. In particular, `ExecutionConfig()` is an +unbudgeted control and retains pylcm's bootstrap width selection. Economic grids, +numerical solver settings and result retention are separate choices. + +`GridConfig` accepts economic grid sizes and numerical solver choices only. Execution +controls belong to `ExecutionConfig`: sharded state names go in `sharded_states`, +declared program widths in `axis_widths`, and the per-device byte ceiling in +`device_memory_bytes`. Obsolete grid execution keywords are rejected. A state-specific +batch size has no automatic conversion to a flattened program width; callers must choose +and record an explicit execution policy. + +For comparisons with an older pylcm version, retain its model/configuration pin and use +a separately reviewed API adapter. Preserve grids, transition declarations, numerical +solver choices, retention and seeds across the comparison. An unspecified width in the +new policy means planner selection, not a translation of an old zero-valued batch size. diff --git a/src/aca_model/aca/model.py b/src/aca_model/aca/model.py index 0db5ff6..94ad7bd 100644 --- a/src/aca_model/aca/model.py +++ b/src/aca_model/aca/model.py @@ -7,7 +7,7 @@ from collections.abc import Mapping from typing import Any -from lcm import AgeGrid, DiscreteGrid, Model +from lcm import AgeGrid, DiscreteGrid, ExecutionConfig, Model from lcm.typing import UserParams from aca_model.aca import PolicyVariant @@ -15,6 +15,7 @@ from aca_model.baseline.model import _fail_if_dcegm_without_consumption_points from aca_model.baseline.regimes import RegimeId, SolverName, build_model_slots from aca_model.config import MODEL_CONFIG, GridConfig +from aca_model.execution import execution_config_for_devices def create_model( @@ -28,6 +29,7 @@ def create_model( pref_type_grid: DiscreteGrid, solver: SolverName = "brute_force", consumption_dollars_points: tuple[float, ...] | None = None, + execution_config: ExecutionConfig | None = None, ) -> Model: """Create an ACA policy variant model. @@ -53,6 +55,9 @@ def create_model( consumption_dollars_points: Construction-time consumption action gridpoints; required under DC-EGM. See `aca_model.baseline.model.create_model`. + execution_config: Explicit hardware-local policy forwarded unchanged. + None uses the smallest selected accelerator allocator limit as the + device-memory budget; CPU construction remains unbudgeted. Returns: pylcm Model. @@ -92,6 +97,11 @@ def create_model( description=f"Structural retirement model ({policy.name})", fixed_params=fixed_params, derived_categoricals=derived_categoricals, + execution_config=( + execution_config_for_devices() + if execution_config is None + else execution_config + ), n_subjects=n_subjects, **model_slots, ) diff --git a/src/aca_model/baseline/model.py b/src/aca_model/baseline/model.py index 87acf7c..5574155 100644 --- a/src/aca_model/baseline/model.py +++ b/src/aca_model/baseline/model.py @@ -13,7 +13,7 @@ from collections.abc import Mapping from typing import Any -from lcm import AgeGrid, DiscreteGrid, Model +from lcm import AgeGrid, DiscreteGrid, ExecutionConfig, Model from lcm.typing import UserParams from aca_model.baseline.regimes import ( @@ -23,6 +23,7 @@ build_model_slots, ) from aca_model.config import MODEL_CONFIG, GridConfig +from aca_model.execution import execution_config_for_devices def create_model( @@ -35,6 +36,7 @@ def create_model( pref_type_grid: DiscreteGrid, solver: SolverName = "brute_force", consumption_dollars_points: tuple[float, ...] | None = None, + execution_config: ExecutionConfig | None = None, ) -> Model: """Create the baseline structural retirement model. @@ -69,6 +71,9 @@ def create_model( continuous-action grid at model construction); `None` keeps the runtime-points grid completed per iteration via `inject_consumption_dollars_points`. + execution_config: Explicit hardware-local policy forwarded unchanged. + None uses the smallest selected accelerator allocator limit as the + device-memory budget; CPU construction remains unbudgeted. Returns: A pylcm Model with 19 regimes (18 non-terminal + dead) spanning @@ -106,6 +111,11 @@ def create_model( description="Baseline structural retirement model (pre-ACA)", fixed_params=fixed_params, derived_categoricals=derived_categoricals, + execution_config=( + execution_config_for_devices() + if execution_config is None + else execution_config + ), n_subjects=n_subjects, **model_slots, ) diff --git a/src/aca_model/baseline/regimes/_common.py b/src/aca_model/baseline/regimes/_common.py index 3168213..6fb7e5b 100644 --- a/src/aca_model/baseline/regimes/_common.py +++ b/src/aca_model/baseline/regimes/_common.py @@ -237,10 +237,7 @@ class Grids: hcc_transitory: Any pref_type: DiscreteGrid grid_config: GridConfig - """The originating `GridConfig`. Exposed on `Grids` so `build_states` - can read per-axis `batch_size` settings for the discrete states it - constructs inline (health, spousal_income, lagged_labor_supply, - claimed_ss) without changing the `build_states`/`build_regime` API.""" + """Grid sizes and numerical settings used to construct each solver.""" # AIME piecewise grid: number of points per segment between the PIA @@ -304,7 +301,6 @@ def build_grids( rho=_WAGE_RHO, sigma=(1.0 - _WAGE_RHO**2) ** 0.5, mu=0.0, - batch_size=grid_config.n_wage_res_batch_size, ) hcc_persistent = get_hcc_persistent_shock(grid_config=grid_config) hcc_transitory = NormalIIDProcess( @@ -323,9 +319,8 @@ def build_grids( start=assets_start, stop=500_000.0, n_points=grid_config.n_assets_gridpoints, - batch_size=grid_config.n_assets_batch_size, ), - aime=_build_aime_grid(grid_config=grid_config, fixed_params=fixed_params), + aime=_build_aime_grid(fixed_params=fixed_params), pension_wealth=_PENSION_WEALTH_GRID, consumption_dollars=( IrregSpacedGrid(n_points=grid_config.n_consumption_dollars_gridpoints) @@ -361,9 +356,7 @@ def get_hcc_persistent_grid_points(*, grid_config: GridConfig) -> FloatND: return get_hcc_persistent_shock(grid_config=grid_config).to_jax() -def _build_aime_grid( - *, grid_config: GridConfig, fixed_params: UserParams -) -> PiecewiseLinSpacedGrid: +def _build_aime_grid(*, fixed_params: UserParams) -> PiecewiseLinSpacedGrid: """Return the AIME grid. The grid is piecewise-linspaced with breakpoints at the PIA bends @@ -383,7 +376,6 @@ def _build_aime_grid( GridBreakpoint(value=kinks[3]), ), points_per_segment=_AIME_PIECE_N_POINTS, - batch_size=grid_config.n_aime_batch_size, ) @@ -460,24 +452,20 @@ def build_states(spec: RegimeSpec, grids: Grids) -> dict: living regime are broadcast from the model level (`build_model_states`). """ can_work = spec["canwork"] == "canwork" - gc = grids.grid_config states: dict = {} states["health"] = DiscreteGrid( Health if spec["mc"] == "oamc" else HealthWithDisability, - batch_size=gc.n_health_batch_size, ) if can_work: states["log_ft_wage_res"] = grids.wage_res if can_work and spec["his"] != "tied": states["lagged_labor_supply"] = DiscreteGrid( LaggedLaborSupply, - batch_size=gc.n_lagged_labor_supply_batch_size, ) if spec["ss"] == "choose": states["claimed_ss"] = DiscreteGrid( ClaimedSS, - batch_size=gc.n_claimed_ss_batch_size, ) return states @@ -712,10 +700,9 @@ def build_model_states(grids: Grids) -> dict: These are the states every living regime carries with an identical grid. pylcm prunes them per regime by DAG reachability, so `dead` keeps only - `assets` and `pref_type` (the bequest DAG). `spousal_income` carries the - `distributed` flag — sharding is legal only on model-level states. + `assets` and `pref_type` (the bequest DAG). Placement is configured + separately through the model execution policy. """ - gc = grids.grid_config return { "assets": grids.assets, "aime": grids.aime, @@ -727,8 +714,6 @@ def build_model_states(grids: Grids) -> dict: "hcc_transitory": grids.hcc_transitory, "spousal_income": DiscreteGrid( SpousalIncome, - batch_size=gc.n_spousal_income_batch_size, - distributed=gc.spousal_income_distributed, ), "pref_type": grids.pref_type, } diff --git a/src/aca_model/baseline/regimes/_dcegm.py b/src/aca_model/baseline/regimes/_dcegm.py index 76cf075..54e3ab2 100644 --- a/src/aca_model/baseline/regimes/_dcegm.py +++ b/src/aca_model/baseline/regimes/_dcegm.py @@ -32,12 +32,8 @@ def build_dcegm_solver(grids: Grids) -> DCEGM: savings_stop = float(assets_points[-1]) - float(assets_points[0]) savings_grid = IrregSpacedGrid( points=tuple(savings_stop * (i / (n_points - 1)) ** 3 for i in range(n_points)), - batch_size=grids.grid_config.n_savings_batch_size, - ) - return DCEGM( - savings_grid=savings_grid, - stochastic_node_batch_size=grids.grid_config.n_stochastic_node_batch_size, ) + return DCEGM(savings_grid=savings_grid) def _fail_if_too_few_savings_gridpoints(n_savings_gridpoints: int) -> None: diff --git a/src/aca_model/baseline/regimes/_nbegm.py b/src/aca_model/baseline/regimes/_nbegm.py index 89dcfa0..23fb96a 100644 --- a/src/aca_model/baseline/regimes/_nbegm.py +++ b/src/aca_model/baseline/regimes/_nbegm.py @@ -1,19 +1,9 @@ -"""NBEGM solver configuration for the ACA M1 vertical-slice regime. +"""NB-EGM numerical configuration for the structural retirement model. -NBEGM is the case-piece endogenous-grid solver for a single 1-D -consumption/savings regime whose budget is split by institutional breakpoints -on a derived monotone income quantity. It shares DC-EGM's post-decision -(savings) spec — consumption is recovered from `resources = max(cash_on_hand, -floor)`, the assets laws are in savings form, the borrowing constraint is the -savings grid's lower bound — but solves only one regime with at most one -discrete action, so it attaches per regime rather than globally. The -function-level rewiring is shared with DC-EGM (`build_dcegm_functions`, the -savings-form assets laws in `_common`); this module holds only the solver -config. +Savings gridpoints, cliff reads, and envelope arithmetic belong to the solver. +Device placement and compiled program widths belong to the model execution policy. """ -import dataclasses - from lcm import IrregSpacedGrid from lcm.solvers import NBEGM @@ -21,62 +11,19 @@ def build_nbegm_solver(grids: Grids) -> NBEGM: - """Build the per-regime NBEGM configuration. - - The savings grid mirrors DC-EGM's: lower bound 0 (the borrowing constraint - in post-decision form), upper bound the assets span, cubically clustered - toward the constraint. Which DAG nodes play the liquid roles is the - regime's declaration (`ACA_LIQUID_MARGIN`), not the solver's; this config - carries numerical settings only. - """ + """Build NB-EGM with the model's savings grid and numerical read rules.""" n_points = grids.grid_config.n_savings_gridpoints _fail_if_too_few_savings_gridpoints(n_points) assets_points = grids.assets.to_jax() savings_stop = float(assets_points[-1]) - float(assets_points[0]) savings_grid = IrregSpacedGrid( points=tuple(savings_stop * (i / (n_points - 1)) ** 3 for i in range(n_points)), - batch_size=grids.grid_config.n_savings_batch_size, ) - solver = NBEGM( + return NBEGM( savings_grid=savings_grid, - # Splay the child stochastic-node expectation per the grid config: `0` (the - # default) reads the whole node mesh in one pass on a memory-rich device; a - # positive value loops it in blocks to fit a tighter budget (a CPU run). - stochastic_node_batch_size=grids.grid_config.n_nbegm_stochastic_node_batch_size, - # Stream the per-interval upper envelope over candidate-segment blocks per - # the grid config; `0` keeps the one-shot dense envelope. - envelope_segment_block_size=( - grids.grid_config.n_nbegm_envelope_segment_block_size - ), - # Which arithmetic decides envelope ownership per the grid config; - # "certified" is exact and can abstain, "ordinary" reads in the working - # format at a fraction of the cost. envelope_arithmetic=grids.grid_config.nbegm_envelope_arithmetic, - # Stream continuation intervals, or let the active byte planner choose - # when the configured width is zero. - interval_batch_size=grids.grid_config.n_nbegm_interval_batch_size, - # Stream both ride-along cores over ride-cell blocks per the grid config; - # `0` vmaps the whole flattened mesh at once. - cell_block_size=grids.grid_config.n_nbegm_cell_block_size, - # Stream the discrete-action branch axis in blocks per the grid config; - # `0` runs the whole axis in one vectorized pass. - branch_batch_size=grids.grid_config.n_nbegm_branch_batch_size, - # Cliff-read mode: exact one-sided limits (default) or the fast bridged - # read for inner estimation loops (see `GridConfig.nbegm_jump_read`). jump_read=grids.grid_config.nbegm_jump_read, ) - budget = grids.grid_config.n_nbegm_max_device_workspace_bytes - if budget is None: - return solver - - fields = {field.name for field in dataclasses.fields(NBEGM)} - if "max_device_workspace_bytes" not in fields: - msg = ( - "GridConfig.n_nbegm_max_device_workspace_bytes requires a pylcm " - "build with the experimental NB-EGM workspace planner." - ) - raise RuntimeError(msg) - return dataclasses.replace(solver, max_device_workspace_bytes=budget) def _fail_if_too_few_savings_gridpoints(n_savings_gridpoints: int) -> None: diff --git a/src/aca_model/benchmark.py b/src/aca_model/benchmark.py index 52b8691..cdeb522 100644 --- a/src/aca_model/benchmark.py +++ b/src/aca_model/benchmark.py @@ -29,7 +29,7 @@ import jax.numpy as jnp import numpy as np from jax import Array -from lcm import DiscreteGrid, Model +from lcm import DiscreteGrid, ExecutionConfig, Model from aca_model.agent.health import GoodHealth from aca_model.agent.labor_market import IsMarried @@ -73,19 +73,21 @@ def create_benchmark_model( *, n_subjects: int, pref_type_grid: DiscreteGrid, + execution_config: ExecutionConfig | None = None, ) -> Model: """Create the aca baseline with `BENCHMARK_GRID_CONFIG` and frozen fixed_params. - The benchmark uses a 2-type `BenchmarkPrefType`. No `batch_size != 0` - on any grid (continuous grids inherit - `BENCHMARK_GRID_CONFIG.n_assets_batch_size = 0` and - `n_aime_batch_size = 0`). + The benchmark uses a 2-type `BenchmarkPrefType`. Grids describe economic + outcomes; the execution policy selects devices and program widths. Args: n_subjects: Forwarded to `lcm.Model(n_subjects=...)`. When set, the first matching `simulate(...)` call AOT-compiles all simulate functions for that batch shape. pref_type_grid: Pref-type grid; pass `DiscreteGrid(BenchmarkPrefType)`. + execution_config: Explicit hardware-local policy forwarded unchanged. + None uses the smallest selected accelerator allocator limit as the + device-memory budget; CPU construction remains unbudgeted. """ fixed_params, wage_params, _ = get_benchmark_params(model=None) return create_model( @@ -94,6 +96,7 @@ def create_benchmark_model( wage_params=wage_params, derived_categoricals=_DERIVED_CATEGORICALS, pref_type_grid=pref_type_grid, + execution_config=execution_config, n_subjects=n_subjects, ) diff --git a/src/aca_model/config.py b/src/aca_model/config.py index 38c6d64..4d65551 100644 --- a/src/aca_model/config.py +++ b/src/aca_model/config.py @@ -1,4 +1,4 @@ -"""Configuration for the aca_model package.""" +"""Economic grids and numerical solver choices for the aca_model package.""" from dataclasses import dataclass from pathlib import Path @@ -21,149 +21,27 @@ class ModelConfig: @dataclass(frozen=True) class GridConfig: + """Economic resolution and numerical choices, independent of device execution.""" + n_assets_gridpoints: int = 24 n_aime_gridpoints: int = 12 n_consumption_dollars_gridpoints: int = 70 n_wage_res_gridpoints: int = 5 n_hcc_persistent_gridpoints: int = 3 n_hcc_transitory_gridpoints: int = 5 - # `batch_size` on the assets / AIME grids: chunked vmap stride for the - # outer state loop. `1` shrinks the per-period Q intermediate by that - # axis's cardinality on hosts where the unsplayed kernel doesn't fit; - # `0` lets a single kernel span the axis. - n_assets_batch_size: int = 0 - n_aime_batch_size: int = 0 - # Sharding flags for discrete state grids. pylcm distributes the - # grid across available devices when the flag is `True`. Sharding - # is only supported on discrete state grids; continuous axes - # (`assets`, `aime`, `wage_res`, `hcc_*`) compile to an all-gather - # of the full V-array per device and are rejected at grid - # construction. Mutually exclusive with `batch_size>0` on the same - # axis (pylcm rejects the combination). `spousal_income_distributed` - # routes through `baseline/regimes/_common.py:build_states` to its - # inline-built `DiscreteGrid(...)` call. - pref_type_distributed: bool = False - spousal_income_distributed: bool = False - # `batch_size` on the inline-constructed discrete state grids — - # health, spousal_income, lagged_labor_supply, claimed_ss. These - # are read in `build_states` via `grids.grid_config`. Setting any - # of them to `1` puts that axis in a Python-level outer loop within - # the discrete-states block of the productmap - # (`_ordered_state_action_names`), shrinking the per-call Q - # intermediate by that axis's cardinality at the cost of one extra - # lax.scan layer. Defaults to `0`; production overrides set to `1` - # to compound the splay across the unsharded discretes. - n_health_batch_size: int = 0 - n_spousal_income_batch_size: int = 0 - n_lagged_labor_supply_batch_size: int = 0 - n_claimed_ss_batch_size: int = 0 - # `batch_size` on the `pref_type` discrete grid: chunked vmap stride - # for the pref-type axis during solve. `1` (one pref-type per Python - # dispatch) shrinks the per-period Q intermediate by `n_pref_types` - # at the cost of an outer Python loop; `0` lets a single kernel span - # all pref-types. Defaults to `0` — the production overrides set it - # to `1` on hardware where the unsplayed kernel doesn't fit. - n_pref_type_batch_size: int = 0 - # `batch_size` on the `wage_res` stochastic shock process: chunked - # productmap stride along the wage-residual stoch axis inside Q_and_F. - # `1` shrinks the per-target Q intermediate by `n_wage_res_gridpoints` - # at the cost of an inner Python loop; `0` lets the productmap span - # the full axis. Defaults to `0` — production overrides set it to `1` - # on hardware where the ACA-overlay per-cell DAG blows the kernel's - # compile-time working set past device HBM. - n_wage_res_batch_size: int = 0 - # Number of nodes on the DC-EGM savings grid (the post-decision endogenous - # grid), cubically clustered toward the borrowing constraint. Drives the - # padded-grid dimension the egm_step kernel carries, so it scales both the - # rolling carry and the gather mesh roughly linearly — the dominant lever on - # DC-EGM device memory. Only consulted under `solver="dcegm"`. + # The post-decision savings grid is cubically clustered toward the borrowing + # constraint. Its node count controls resolution under DC-EGM and NB-EGM. n_savings_gridpoints: int = 200 - # `batch_size` on the DC-EGM savings grid (the post-decision endogenous - # grid). pylcm splays the per-savings-node continuation into `lax.map` - # blocks of this size, shrinking the binding `egm_step` working buffer by - # roughly the block factor while the upper envelope still runs on the full - # gathered grid (value function unchanged). `0` keeps the whole grid in one - # kernel. Only consulted under `solver="dcegm"`. - n_savings_batch_size: int = 0 - # `batch_size` on the DC-EGM child stochastic-node expectation (the - # process-state mesh the egm_step kernel sums over). pylcm splays that - # expectation into `lax.map` blocks of this size, collapsing the gather - # buffer's node axis by roughly the block factor (value function - # unchanged) — the grid-independent lever for the residual egm_step - # transient that savings-batching leaves behind. `0` evaluates the whole - # mesh in one kernel. Only consulted under `solver="dcegm"`. - n_stochastic_node_batch_size: int = 0 - # Block size for splaying the NBEGM continuation's child stochastic-node - # expectation (health, health-cost shocks, the wage residual). `0` reads the - # whole joint node mesh in one pass — fast, but its peak intermediate scales - # with the full ride-along × node × child-grid product. A positive value loops - # the mesh in blocks of that size, trading runtime for a much smaller peak; `1` - # (one node at a time) is the memory-minimal setting for a CPU validation grid. - # Only consulted under `solver="nbegm"`. - n_nbegm_stochastic_node_batch_size: int = 0 - # Streams the per-interval upper envelope over candidate-segment blocks of this - # size instead of materialising the full (query x candidate) bracket matrix per - # ride cell. `0` keeps the one-shot dense envelope; the result is identical - # either way — the knob trades peak device memory against a sequential scan. - # Only consulted under `solver="nbegm"`. - n_nbegm_envelope_segment_block_size: int = 0 - # Compiled batch width for the liquid-interval axis. With an active byte - # budget, 0 lets the planner choose up to the full axis; a positive value - # caps it. Only consulted under `solver="nbegm"`. - n_nbegm_interval_batch_size: int = 0 - # Authoritative busiest-device workspace budget. None preserves manual - # blocks; a positive value activates compile-only continuation/envelope - # planning before backward induction. Only consulted under `solver="nbegm"`. - n_nbegm_max_device_workspace_bytes: int | None = None - # Which arithmetic decides ownership in the merged upper envelope: - # - "certified" (default) — candidates are compared in double-double precision - # and no winner is published where none is separated, so the reported owner is - # one the arithmetic could prove. Ordering survives the cancellation a nearly - # tied crossing produces, at roughly an order of magnitude more arithmetic per - # read — and the envelope read is the dominant per-cell cost of a case-piece - # solve, so this is the setting that decides the solver's runtime. - # - "ordinary" — each candidate is read in the working format and the largest - # owns the query. Adequate wherever candidate values are separated by much - # more than the format's resolution at their own magnitude. - # Incompatible with a positive `n_nbegm_envelope_segment_block_size`, which - # selects a blocked scan that carries the certified arithmetic only. - # Only consulted under `solver="nbegm"`. + # Arithmetic used to compare upper-envelope candidates: + # - "certified": double-double comparisons publish only separated winners. + # - "ordinary": comparisons use the working floating-point format. nbegm_envelope_arithmetic: Literal["certified", "ordinary"] = "certified" - # Streams both NBEGM ride-along cores (continuation fan-out and envelope - # solve) over ride-cell blocks of this size instead of vmapping the whole - # flattened ride mesh at once — the dominant peak-memory term at production - # mesh sizes. `0` keeps the whole-mesh vmap; the result is identical either - # way. Only consulted under `solver="nbegm"`. - # - # Backend-dependent tuning at production ride-mesh sizes (under both cliff-read - # modes). On GPU the whole-mesh vmap stays within a few GiB, so `0` is fine. The - # CPU XLA backend does not fuse the fan-out and materialises the whole flattened - # ride mesh at once — a production-grid solve then needs hundreds of GiB even at - # a small asset grid (the blow-up rides the aime/shock/health mesh, not assets). - # Set this to a positive block (e.g. 64) for a CPU solve; it bounds the peak to - # the GPU's few-GiB footprint at the cost of serialising the mesh into a - # `lax.map` scan. - n_nbegm_cell_block_size: int = 0 - # How NBEGM parents read the child value's institutional cliffs: - # - "one_sided" (default) — carry rows hold each cliff preimage as a duplicated - # abscissa with exact one-sided limits; reads never average across a cliff, - # but publishing the topology gates the stochastic-dim fold off (slower). - # - "bridged" — plain carry rows; interpolation may bridge a cliff like any - # finite-grid solver, and the fold stays available. The fast setting for - # inner estimation loops, polished afterwards under "one_sided". - # Only consulted under `solver="nbegm"`. + # How NB-EGM parents read institutional cliffs in the child value: + # - "one_sided": duplicated abscissae preserve one-sided limits. + # - "bridged": interpolation may bridge finite-grid discontinuities. + # Live labor choices require "bridged" because branch-dependent income + # thresholds cannot share one query grid of one-sided cliff limits. nbegm_jump_read: Literal["one_sided", "bridged"] = "one_sided" - # Block size for streaming the discrete-action branch axis in both NBEGM - # ride-along cores (one continuation row / one continuous subproblem per live - # labor level). `0` runs the whole branch axis in one vectorized pass; a positive - # value scans it in blocks of that many branches (identical result, per-branch - # intermediates bounded by one block). Only consulted under `solver="nbegm"`. - n_nbegm_branch_batch_size: int = 0 - # `labor_supply` enters `countable_income`, which carries the SSI income test, so - # each labor level puts that breakpoint at a different liquid level. The one-sided - # read publishes its cliff limits on a single query grid shared across branches, - # which the two cannot both satisfy, so a regime carrying `labor_supply` builds - # only under `nbegm_jump_read="bridged"`. MODEL_CONFIG = ModelConfig() @@ -176,6 +54,4 @@ class GridConfig: n_wage_res_gridpoints=3, n_hcc_persistent_gridpoints=3, n_hcc_transitory_gridpoints=3, - n_assets_batch_size=0, - n_aime_batch_size=0, ) diff --git a/src/aca_model/execution.py b/src/aca_model/execution.py new file mode 100644 index 0000000..b940674 --- /dev/null +++ b/src/aca_model/execution.py @@ -0,0 +1,53 @@ +"""Hardware-local execution policies derived from allocator observations.""" + +import jax +from lcm import ExecutionConfig + + +def execution_config_for_devices( + *, devices: tuple[int, ...] | None = None +) -> ExecutionConfig: + """Create an execution policy bounded by the selected accelerator pools. + + The smallest allocator limit is a total per-device ceiling. Pylcm accounts + for its resident arrays within that ceiling, so current allocator usage is + not subtracted here. CPU devices do not supply an accelerator-pool budget. + + Args: + devices: Visible device IDs to select, in order; None selects all devices + on JAX's default backend. + + Returns: + ExecutionConfig with explicit device IDs and the measured accelerator + budget. Axis widths and sharded states remain unspecified. + + Raises: + ValueError: A selected device is not visible or an accelerator has no + positive integer allocator limit. Missing limits require an explicit + policy. + """ + policy = ExecutionConfig(devices=devices) + visible = {device.id: device for device in jax.devices()} + selected = tuple(visible) if policy.devices is None else policy.devices + missing = tuple(device_id for device_id in selected if device_id not in visible) + if missing: + msg = f"Selected devices are not visible: {missing}." + raise ValueError(msg) + limits = [] + for device_id in selected: + device = visible[device_id] + if device.platform == "cpu": + continue + stats = device.memory_stats() or {} + limit = stats.get("bytes_limit") + if type(limit) is not int or limit <= 0: + msg = ( + "No positive integer allocator byte limit reported for " + f"device {device_id}." + ) + raise ValueError(msg) + limits.append(limit) + return ExecutionConfig( + devices=selected, + device_memory_bytes=min(limits) if limits else None, + ) diff --git a/tests/test_aime_grid.py b/tests/test_aime_grid.py index 368fe37..d63da86 100644 --- a/tests/test_aime_grid.py +++ b/tests/test_aime_grid.py @@ -9,7 +9,6 @@ _AIME_PIECE_N_POINTS, _build_aime_grid, ) -from aca_model.config import BENCHMARK_GRID_CONFIG # Production SSA bend points plus the delayed-retirement-credit extension: # 0, kink_0, kink_1, taxable-max, and the extension point that carries the @@ -20,9 +19,7 @@ def test_build_aime_grid_owns_each_pia_breakpoint_on_the_right() -> None: """The AIME grid preserves four PIA pieces with right-owned interiors.""" - grid = _build_aime_grid( - grid_config=BENCHMARK_GRID_CONFIG, fixed_params=_FIXED_PARAMS - ) + grid = _build_aime_grid(fixed_params=_FIXED_PARAMS) np.testing.assert_allclose( [point.value for point in grid.breakpoints], _PIA_AIME_GRID[1:-1], @@ -33,9 +30,7 @@ def test_build_aime_grid_owns_each_pia_breakpoint_on_the_right() -> None: def test_build_aime_grid_top_point_is_extension_aime() -> None: """The grid reaches the delayed-credit extension AIME at its top.""" - grid = _build_aime_grid( - grid_config=BENCHMARK_GRID_CONFIG, fixed_params=_FIXED_PARAMS - ) + grid = _build_aime_grid(fixed_params=_FIXED_PARAMS) np.testing.assert_allclose(float(grid.to_jax().max()), 187954.752, rtol=1e-5) diff --git a/tests/test_beartype_claw.py b/tests/test_beartype_claw.py index e4130be..2fb7873 100644 --- a/tests/test_beartype_claw.py +++ b/tests/test_beartype_claw.py @@ -12,7 +12,7 @@ import pytest from beartype.roar import BeartypeCallHintViolation -from helpers.model import make_baseline_model # ty: ignore[unresolved-import] +from helpers.model import make_baseline_model def test_claw_checks_aca_model() -> None: @@ -22,4 +22,4 @@ def test_claw_checks_aca_model() -> None: by the claw before the value reaches pylcm's own `Model` perimeter. """ with pytest.raises(BeartypeCallHintViolation): - make_baseline_model(n_subjects="not an int") + make_baseline_model(n_subjects="not an int") # ty: ignore[invalid-argument-type] diff --git a/tests/test_benchmark.py b/tests/test_benchmark.py index ac9b525..56922ef 100644 --- a/tests/test_benchmark.py +++ b/tests/test_benchmark.py @@ -2,7 +2,14 @@ import numpy as np import pytest -from lcm import DiscreteGrid, GridBreakpoint, PiecewiseLinSpacedGrid +from lcm import ( + DiscreteGrid, + GridBreakpoint, + IrregSpacedGrid, + LinSpacedGrid, + PiecewiseLinSpacedGrid, + RouwenhorstAR1Process, +) from aca_model.agent.preferences import BenchmarkPrefType from aca_model.benchmark import ( @@ -13,7 +20,7 @@ def test_benchmark_model_builds_with_the_current_pylcm_grid_api() -> None: - """The frozen benchmark model constructs with breakpoint-first AIME grids.""" + """The frozen benchmark preserves its state and action grid extents.""" model = create_benchmark_model( n_subjects=1, pref_type_grid=DiscreteGrid(BenchmarkPrefType), @@ -23,7 +30,22 @@ def test_benchmark_model_builds_with_the_current_pylcm_grid_api() -> None: assert isinstance(aime, PiecewiseLinSpacedGrid) assert all(isinstance(point, GridBreakpoint) for point in aime.breakpoints) assert all(point.owner == "right" for point in aime.breakpoints) - assert aime.n_points == sum(aime.points_per_segment) + assert aime.n_points == 38 + regime = model.user_regimes["retiree_nomc_inelig_canwork"] + assert len(model.user_regimes) == 19 + assert model.n_periods == 45 + assets = regime.states["assets"] + wage_res = regime.states["log_ft_wage_res"] + pref_type = regime.states["pref_type"] + consumption = regime.actions["consumption_dollars"] + assert isinstance(assets, LinSpacedGrid) + assert isinstance(wage_res, RouwenhorstAR1Process) + assert isinstance(pref_type, DiscreteGrid) + assert isinstance(consumption, IrregSpacedGrid) + assert assets.n_points == 3 + assert wage_res.n_points == 3 + assert len(pref_type.categories) == 2 + assert consumption.n_points == 5 @pytest.mark.long_running diff --git a/tests/test_dcegm_model_creation.py b/tests/test_dcegm_model_creation.py index 492fdd2..45274ff 100644 --- a/tests/test_dcegm_model_creation.py +++ b/tests/test_dcegm_model_creation.py @@ -11,7 +11,7 @@ import numpy as np import pytest -from helpers.model import _DERIVED_CATEGORICALS # ty: ignore[unresolved-import] +from helpers.model import _DERIVED_CATEGORICALS from lcm import DiscreteGrid, IrregSpacedGrid, Model, Regime from lcm.consumption_savings_regime import ConsumptionSavingsRegime from lcm.solvers import DCEGM @@ -162,42 +162,6 @@ def test_dcegm_benchmark_model_builds() -> None: assert isinstance(model.user_regimes["retiree_nomc_inelig_canwork"].solver, DCEGM) -def test_savings_grid_batch_size_follows_grid_config() -> None: - """`GridConfig.n_savings_batch_size` sets the `batch_size` on every - living regime's DC-EGM savings grid, so the post-decision continuation - splays into `lax.map` blocks of that width.""" - grid_config = dataclasses.replace(BENCHMARK_GRID_CONFIG, n_savings_batch_size=50) - regimes = build_all_regimes( - grid_config=grid_config, - fixed_params=_FIXED_PARAMS, - wage_params=_WAGE_PARAMS, - pref_type_grid=DiscreteGrid(BenchmarkPrefType), - solver="dcegm", - ) - for name in REGIME_SPECS: - solver = cast("DCEGM", regimes[name].solver) - assert solver.savings_grid.batch_size == 50, name - - -def test_stochastic_node_batch_size_follows_grid_config() -> None: - """`GridConfig.n_stochastic_node_batch_size` sets `stochastic_node_batch_size` - on every living regime's DC-EGM solver, so the child stochastic-node - expectation splays into `lax.map` blocks of that width.""" - grid_config = dataclasses.replace( - BENCHMARK_GRID_CONFIG, n_stochastic_node_batch_size=7 - ) - regimes = build_all_regimes( - grid_config=grid_config, - fixed_params=_FIXED_PARAMS, - wage_params=_WAGE_PARAMS, - pref_type_grid=DiscreteGrid(BenchmarkPrefType), - solver="dcegm", - ) - for name in REGIME_SPECS: - solver = cast("DCEGM", regimes[name].solver) - assert solver.stochastic_node_batch_size == 7, name - - def test_savings_grid_length_follows_grid_config() -> None: """`GridConfig.n_savings_gridpoints` sets the number of nodes on every living regime's DC-EGM savings grid.""" diff --git a/tests/test_execution_budget.py b/tests/test_execution_budget.py new file mode 100644 index 0000000..0dafef2 --- /dev/null +++ b/tests/test_execution_budget.py @@ -0,0 +1,62 @@ +"""Allocator observations determine the execution ceiling on selected devices.""" + +from dataclasses import dataclass + +import jax +import pytest + +from aca_model import execution + + +@dataclass +class _Device: + id: int + platform: str + stats: dict[str, int] | None + + def memory_stats(self): + return self.stats + + +def test_budget_uses_smallest_selected_allocator_limit(monkeypatch): + """Only selected GPU pools bound execution; current usage is not subtracted.""" + devices = [ + _Device(0, "gpu", {"bytes_limit": 1000, "bytes_in_use": 200}), + _Device(1, "gpu", {"bytes_limit": 800, "bytes_in_use": 100}), + _Device(2, "gpu", None), + ] + monkeypatch.setattr(jax, "devices", lambda: devices) + + config = execution.execution_config_for_devices(devices=(0, 1)) + + assert config.devices == (0, 1) + assert config.device_memory_bytes == 800 + assert config.axis_widths == {} + assert config.sharded_states == () + + +@pytest.mark.parametrize("stats", [None, {}, {"bytes_limit": 0}]) +def test_missing_gpu_limit_is_refused(monkeypatch, stats): + """An unknown selected GPU ceiling must not silently select bootstrap widths.""" + monkeypatch.setattr(jax, "devices", lambda: [_Device(0, "gpu", stats)]) + + with pytest.raises(ValueError, match=r"allocator.*limit.*device 0"): + execution.execution_config_for_devices() + + +def test_cpu_construction_does_not_invent_an_allocator_budget(monkeypatch): + """CPU devices retain an unspecified byte budget when no pool limit exists.""" + monkeypatch.setattr(jax, "devices", lambda: [_Device(0, "cpu", None)]) + + config = execution.execution_config_for_devices() + + assert config.devices == (0,) + assert config.device_memory_bytes is None + + +def test_selected_device_must_be_visible(monkeypatch): + """An unavailable selected device is rejected before reading allocator stats.""" + monkeypatch.setattr(jax, "devices", lambda: [_Device(0, "gpu", None)]) + + with pytest.raises(ValueError, match="Selected devices are not visible"): + execution.execution_config_for_devices(devices=(1,)) diff --git a/tests/test_execution_config.py b/tests/test_execution_config.py new file mode 100644 index 0000000..04f19a7 --- /dev/null +++ b/tests/test_execution_config.py @@ -0,0 +1,117 @@ +"""Model factories expose pylcm's execution policy without changing economic grids.""" + +from functools import partial +from types import SimpleNamespace + +import jax +import pytest +from lcm import DiscreteGrid, ExecutionConfig +from lcm.exceptions import ExecutionPlanningError + +from aca_model.aca import model as aca_model_module +from aca_model.aca.health_insurance import PolicyVariant +from aca_model.aca.model import create_model as create_aca_model +from aca_model.agent.health import GoodHealth +from aca_model.agent.labor_market import IsMarried +from aca_model.agent.preferences import BenchmarkPrefType +from aca_model.baseline import model as baseline_model_module +from aca_model.baseline.health_insurance import HealthInsuranceState +from aca_model.baseline.model import create_model +from aca_model.benchmark import create_benchmark_model, get_benchmark_params +from aca_model.config import BENCHMARK_GRID_CONFIG + + +def _factory(kind): + if kind == "benchmark": + return partial( + create_benchmark_model, + n_subjects=1, + pref_type_grid=DiscreteGrid(BenchmarkPrefType), + ) + fixed_params, wage_params, _ = get_benchmark_params(model=None) + factory = ( + partial(create_aca_model, policy=PolicyVariant.ACA) + if kind == "aca" + else create_model + ) + return partial( + factory, + n_subjects=1, + fixed_params=fixed_params, + wage_params=wage_params, + derived_categoricals={ + "good_health": DiscreteGrid(GoodHealth), + "is_married": DiscreteGrid(IsMarried), + "his": DiscreteGrid(HealthInsuranceState), + "target_his": DiscreteGrid(HealthInsuranceState), + "pref_type": DiscreteGrid(BenchmarkPrefType), + }, + grid_config=BENCHMARK_GRID_CONFIG, + pref_type_grid=DiscreteGrid(BenchmarkPrefType), + ) + + +@pytest.mark.parametrize("kind", ["baseline", "aca", "benchmark"]) +def test_factory_preserves_explicit_device_selection(kind): + """A factory uses exactly the selected device and preserves the model grids.""" + devices = (jax.devices()[-1].id,) + model = _factory(kind)(execution_config=ExecutionConfig(devices=devices)) + + assert model.execution_devices == devices + assert len(model.user_regimes) == 19 + regime = model.user_regimes["retiree_nomc_inelig_canwork"] + assert regime.states["assets"].n_points == 3 + assert regime.states["aime"].n_points == 38 + assert regime.actions["consumption_dollars"].n_points == 5 + + +@pytest.mark.parametrize("kind", ["baseline", "aca", "benchmark"]) +def test_factory_preserves_axis_width_validation(kind): + """A width for an undeclared program axis is refused by the model.""" + config = ExecutionConfig(axis_widths={"aca_undeclared_axis": 1}) + with pytest.raises(ExecutionPlanningError, match="aca_undeclared_axis"): + _factory(kind)(execution_config=config) + + +@pytest.mark.parametrize("kind", ["baseline", "aca", "benchmark"]) +@pytest.mark.parametrize( + ("requested_policy", "expected_budget"), + [ + (None, 800), + ( + ExecutionConfig( + devices=(0,), + device_memory_bytes=400, + sharded_states=("pref_type",), + axis_widths={"cell": 32}, + ), + 400, + ), + (ExecutionConfig(devices=(0,)), None), + ], +) +def test_factory_supplies_the_requested_or_measured_budget( + monkeypatch, kind, requested_policy, expected_budget +): + """Defaults use measured limits; explicit budgets and unbudgeted controls survive.""" + factory = _factory(kind) + device = SimpleNamespace( + id=0, platform="gpu", memory_stats=lambda: {"bytes_limit": 800} + ) + monkeypatch.setattr(jax, "devices", lambda: [device]) + + class ConstructionObservedError(Exception): + pass + + def observe_pylcm_constructor(**kwargs): + policy = kwargs["execution_config"] + assert policy.device_memory_bytes == expected_budget + assert policy.devices == (0,) + if requested_policy is not None: + assert policy == requested_policy + raise ConstructionObservedError + + monkeypatch.setattr(baseline_model_module, "Model", observe_pylcm_constructor) + monkeypatch.setattr(aca_model_module, "Model", observe_pylcm_constructor) + with pytest.raises(ConstructionObservedError): + factory(execution_config=requested_policy) diff --git a/tests/test_grid_config.py b/tests/test_grid_config.py new file mode 100644 index 0000000..7eb6564 --- /dev/null +++ b/tests/test_grid_config.py @@ -0,0 +1,35 @@ +"""Economic-grid configuration rejects hardware execution controls.""" + +import pytest + +from aca_model.config import GridConfig + + +@pytest.mark.parametrize( + "field", + [ + "n_assets_batch_size", + "n_aime_batch_size", + "pref_type_distributed", + "spousal_income_distributed", + "n_health_batch_size", + "n_spousal_income_batch_size", + "n_lagged_labor_supply_batch_size", + "n_claimed_ss_batch_size", + "n_pref_type_batch_size", + "n_wage_res_batch_size", + "n_savings_batch_size", + "n_stochastic_node_batch_size", + "n_nbegm_stochastic_node_batch_size", + "n_nbegm_envelope_segment_block_size", + "n_nbegm_interval_batch_size", + "n_nbegm_max_device_workspace_bytes", + "n_nbegm_cell_block_size", + "n_nbegm_branch_batch_size", + ], +) +def test_grid_config_refuses_execution_controls(field): + """A misplaced execution control fails visibly instead of being ignored.""" + value = True if field.endswith("distributed") else 1 + with pytest.raises(TypeError, match=field): + GridConfig(**{field: value}) # ty: ignore[invalid-argument-type] diff --git a/tests/test_model_creation.py b/tests/test_model_creation.py index 5b2415c..68b7d77 100644 --- a/tests/test_model_creation.py +++ b/tests/test_model_creation.py @@ -2,10 +2,9 @@ import inspect from collections.abc import Mapping -from dataclasses import replace import pytest -from helpers.model import ( # ty: ignore[unresolved-import] +from helpers.model import ( make_aca_model, make_baseline_model, ) @@ -367,43 +366,6 @@ def test_baseline_model_creates() -> None: assert len(model.user_regimes) == 19 -@pytest.mark.parametrize( - ("config_field", "state_name"), - [ - ("spousal_income_distributed", "spousal_income"), - ], -) -def test_discrete_state_distributed_flag_propagates_to_model_states( - config_field: str, state_name: str -) -> None: - """`GridConfig._distributed=True` sets `distributed=True` on the - model-level `DiscreteGrid` for that axis (sharding is legal only on - model-level states).""" - gc = replace(BENCHMARK_GRID_CONFIG, **{config_field: True}) - grids = build_grids( - grid_config=gc, - fixed_params=_FIXED_PARAMS, - wage_params=_WAGE_PARAMS, - pref_type_grid=DiscreteGrid(BenchmarkPrefType), - ) - model_states = build_model_states(grids) - assert model_states[state_name].distributed is True - - -@pytest.mark.parametrize( - "state_name", - ["lagged_labor_supply", "claimed_ss", "spousal_income"], -) -def test_discrete_state_distributed_flag_defaults_to_false(state_name: str) -> None: - """`distributed` on inline-built discrete states defaults to `False` so - configurations that do not opt in see no behaviour change.""" - if state_name == "spousal_income": - grid = build_model_states(_GRIDS)[state_name] - else: - grid = build_regime("retiree_dimc_choose_canwork").states[state_name] - assert grid.distributed is False - - def test_dead_regime_prunes_unused_broadcast_states() -> None: """States every living regime shares are declared once at the model level; `dead` keeps only what the bequest DAG reads (`assets`, `pref_type`). diff --git a/tests/test_nbegm_labor_live_validation.py b/tests/test_nbegm_labor_live_validation.py index 53a732c..279c6fd 100644 --- a/tests/test_nbegm_labor_live_validation.py +++ b/tests/test_nbegm_labor_live_validation.py @@ -16,8 +16,8 @@ import numpy as np import pytest -from helpers.model import _DERIVED_CATEGORICALS # ty: ignore[unresolved-import] -from lcm import DiscreteGrid +from helpers.model import _DERIVED_CATEGORICALS +from lcm import DiscreteGrid, ExecutionConfig from aca_model.agent.preferences import BenchmarkPrefType from aca_model.baseline.model import create_model @@ -29,16 +29,11 @@ def _solve_m1(solver: SolverName) -> tuple[dict[int, np.ndarray], int]: - # The CPU XLA backend does not fuse the ride-cell fan-out and materialises the - # whole flattened ride mesh at once, so a full-model solve needs hundreds of GiB on - # host even at a tiny asset grid. `n_nbegm_cell_block_size` streams the mesh in - # blocks (identical result) to bound the peak to the GPU's few-GiB footprint; the - # live-labor branch axis makes this essential on CPU. A coarser savings grid keeps - # the check quick. On GPU the whole-mesh vmap stays small, so production sets 0. + # A bounded ride-cell width limits continuation fan-out on CPU. + # Savings resolution is shared by both numerical comparison arms. grid_config = dataclasses.replace( BENCHMARK_GRID_CONFIG, nbegm_jump_read="bridged", - n_nbegm_cell_block_size=32, n_savings_gridpoints=50, ) fixed_params, wage_params, _ = get_benchmark_params(model=None) @@ -50,6 +45,9 @@ def _solve_m1(solver: SolverName) -> tuple[dict[int, np.ndarray], int]: grid_config=grid_config, pref_type_grid=DiscreteGrid(BenchmarkPrefType), solver=solver, + execution_config=ExecutionConfig( + axis_widths={"cell": 32} if solver == "nbegm" else {} + ), ) _, _, params = get_benchmark_params(model=model) solution = model.solve(params=params, log_level="off") diff --git a/tests/test_nbegm_model_creation.py b/tests/test_nbegm_model_creation.py index fdd3a4b..79d0265 100644 --- a/tests/test_nbegm_model_creation.py +++ b/tests/test_nbegm_model_creation.py @@ -17,7 +17,7 @@ from typing import cast import pytest -from helpers.model import _DERIVED_CATEGORICALS # ty: ignore[unresolved-import] +from helpers.model import _DERIVED_CATEGORICALS from lcm import DiscreteGrid, Model, Regime from lcm.consumption_savings_regime import ConsumptionSavingsRegime, LiquidMargin from lcm.exceptions import RegimeInitializationError @@ -352,44 +352,3 @@ def test_nbegm_aca_variants_leave_no_free_buy_private_params( if isinstance(params, dict) and "buy_private" in params ] assert offenders == [], offenders - - -def test_nbegm_all_streaming_axes_are_forwarded() -> None: - """ACA forwards every streaming axis supported by pylcm.""" - grid_config = dataclasses.replace( - BENCHMARK_GRID_CONFIG, - n_nbegm_stochastic_node_batch_size=2, - n_nbegm_envelope_segment_block_size=3, - n_nbegm_interval_batch_size=4, - n_nbegm_cell_block_size=5, - n_nbegm_branch_batch_size=6, - ) - grids = build_grids( - grid_config=grid_config, - fixed_params=_FIXED_PARAMS, - wage_params=_WAGE_PARAMS, - pref_type_grid=DiscreteGrid(BenchmarkPrefType), - ) - solver = build_nbegm_solver(grids) - assert solver.stochastic_node_batch_size == 2 - assert solver.envelope_segment_block_size == 3 - assert solver.interval_batch_size == 4 - assert solver.cell_block_size == 5 - assert solver.branch_batch_size == 6 - - -def test_nbegm_workspace_budget_requires_the_experimental_planner() -> None: - """A planner budget is refused when pylcm has no workspace planner.""" - grid_config = dataclasses.replace( - BENCHMARK_GRID_CONFIG, - n_nbegm_max_device_workspace_bytes=72 * 1024**3, - ) - grids = build_grids( - grid_config=grid_config, - fixed_params=_FIXED_PARAMS, - wage_params=_WAGE_PARAMS, - pref_type_grid=DiscreteGrid(BenchmarkPrefType), - ) - - with pytest.raises(RuntimeError, match="experimental NB-EGM workspace planner"): - build_nbegm_solver(grids) diff --git a/tests/test_nbegm_solve_validation.py b/tests/test_nbegm_solve_validation.py index af16729..4c840da 100644 --- a/tests/test_nbegm_solve_validation.py +++ b/tests/test_nbegm_solve_validation.py @@ -19,7 +19,7 @@ import numpy as np import pytest -from helpers.model import _DERIVED_CATEGORICALS # ty: ignore[unresolved-import] +from helpers.model import _DERIVED_CATEGORICALS from lcm import DiscreteGrid from aca_model.agent.preferences import BenchmarkPrefType diff --git a/tests/test_social_security.py b/tests/test_social_security.py index 2ee11a0..30fc799 100644 --- a/tests/test_social_security.py +++ b/tests/test_social_security.py @@ -6,7 +6,7 @@ import jax.numpy as jnp import numpy as np import pandas as pd -from helpers.social_security import ( # ty: ignore[unresolved-import] +from helpers.social_security import ( compute_di_dropout_scale, compute_pia_table, ) diff --git a/tests/test_ss_benefit_integration.py b/tests/test_ss_benefit_integration.py index 0b70990..e111b96 100644 --- a/tests/test_ss_benefit_integration.py +++ b/tests/test_ss_benefit_integration.py @@ -5,7 +5,7 @@ """ import jax.numpy as jnp -from helpers.social_security import compute_pia_table # ty: ignore[unresolved-import] +from helpers.social_security import compute_pia_table from aca_model.agent.labor_market import LaborSupply from aca_model.environment import social_security From 735a25090ec647ef8d42cbb919369acc46f1a493 Mon Sep 17 00:00:00 2001 From: Hans-Martin von Gaudecker Date: Sun, 13 Sep 2026 00:11:05 +0200 Subject: [PATCH 15/16] API Remove constructor population hints from model factories --- src/aca_model/aca/model.py | 3 --- src/aca_model/baseline/model.py | 5 +---- src/aca_model/benchmark.py | 5 ----- tests/helpers/model.py | 6 ++---- tests/test_beartype_claw.py | 9 +++++---- tests/test_benchmark.py | 20 +++++++++++++++---- tests/test_dcegm_model_creation.py | 3 --- tests/test_dcegm_parity.py | 1 - tests/test_execution_config.py | 9 +++++++-- .../test_initial_conditions_extreme_assets.py | 3 +-- tests/test_model_creation.py | 20 +++++++++---------- tests/test_nbegm_labor_live_validation.py | 1 - tests/test_nbegm_model_creation.py | 3 --- tests/test_nbegm_solve_validation.py | 1 - 14 files changed, 42 insertions(+), 47 deletions(-) diff --git a/src/aca_model/aca/model.py b/src/aca_model/aca/model.py index 94ad7bd..5cc6139 100644 --- a/src/aca_model/aca/model.py +++ b/src/aca_model/aca/model.py @@ -20,7 +20,6 @@ def create_model( *, - n_subjects: int, policy: PolicyVariant, fixed_params: UserParams, wage_params: Mapping[str, Any], @@ -34,7 +33,6 @@ def create_model( """Create an ACA policy variant model. Args: - n_subjects: Forwarded to `lcm.Model(n_subjects=...)`. policy: Which ACA policy combination to apply (e.g. `PolicyVariant.ACA`). fixed_params: Parameters to fix at model creation time. Pass @@ -102,6 +100,5 @@ def create_model( if execution_config is None else execution_config ), - n_subjects=n_subjects, **model_slots, ) diff --git a/src/aca_model/baseline/model.py b/src/aca_model/baseline/model.py index 5574155..6040c45 100644 --- a/src/aca_model/baseline/model.py +++ b/src/aca_model/baseline/model.py @@ -5,7 +5,7 @@ Usage: from aca_model.baseline.model import create_model - model = create_model(n_subjects=..., fixed_params=..., wage_params=..., ...) + model = create_model(fixed_params=..., wage_params=..., ...) params = get_default_params() V = model.solve(params) """ @@ -28,7 +28,6 @@ def create_model( *, - n_subjects: int, fixed_params: UserParams, wage_params: Mapping[str, Any], derived_categoricals: Mapping[str, DiscreteGrid], @@ -41,7 +40,6 @@ def create_model( """Create the baseline structural retirement model. Args: - n_subjects: Forwarded to `lcm.Model(n_subjects=...)`. fixed_params: Parameters to fix at model creation time. Fixed params are partialled into compiled functions and removed from the params template. Pass data-derived constants here; @@ -116,7 +114,6 @@ def create_model( if execution_config is None else execution_config ), - n_subjects=n_subjects, **model_slots, ) diff --git a/src/aca_model/benchmark.py b/src/aca_model/benchmark.py index cdeb522..54fe41c 100644 --- a/src/aca_model/benchmark.py +++ b/src/aca_model/benchmark.py @@ -71,7 +71,6 @@ def create_benchmark_model( *, - n_subjects: int, pref_type_grid: DiscreteGrid, execution_config: ExecutionConfig | None = None, ) -> Model: @@ -81,9 +80,6 @@ def create_benchmark_model( outcomes; the execution policy selects devices and program widths. Args: - n_subjects: Forwarded to `lcm.Model(n_subjects=...)`. When set, the - first matching `simulate(...)` call AOT-compiles all simulate - functions for that batch shape. pref_type_grid: Pref-type grid; pass `DiscreteGrid(BenchmarkPrefType)`. execution_config: Explicit hardware-local policy forwarded unchanged. None uses the smallest selected accelerator allocator limit as the @@ -97,7 +93,6 @@ def create_benchmark_model( derived_categoricals=_DERIVED_CATEGORICALS, pref_type_grid=pref_type_grid, execution_config=execution_config, - n_subjects=n_subjects, ) diff --git a/tests/helpers/model.py b/tests/helpers/model.py index 55571e9..4818827 100644 --- a/tests/helpers/model.py +++ b/tests/helpers/model.py @@ -26,11 +26,10 @@ } -def make_baseline_model(*, n_subjects: int) -> Model: +def make_baseline_model() -> Model: """Baseline model on `BENCHMARK_GRID_CONFIG` with the benchmark snapshot params.""" fixed_params, wage_params, _ = get_benchmark_params(model=None) return _create_baseline_model( - n_subjects=n_subjects, fixed_params=fixed_params, wage_params=wage_params, derived_categoricals=_DERIVED_CATEGORICALS, @@ -39,11 +38,10 @@ def make_baseline_model(*, n_subjects: int) -> Model: ) -def make_aca_model(*, n_subjects: int, policy: PolicyVariant) -> Model: +def make_aca_model(*, policy: PolicyVariant) -> Model: """ACA model on `BENCHMARK_GRID_CONFIG` with the benchmark snapshot params.""" fixed_params, wage_params, _ = get_benchmark_params(model=None) return _create_aca_model( - n_subjects=n_subjects, policy=policy, fixed_params=fixed_params, wage_params=wage_params, diff --git a/tests/test_beartype_claw.py b/tests/test_beartype_claw.py index 2fb7873..b5824d1 100644 --- a/tests/test_beartype_claw.py +++ b/tests/test_beartype_claw.py @@ -12,14 +12,15 @@ import pytest from beartype.roar import BeartypeCallHintViolation -from helpers.model import make_baseline_model + +from aca_model.benchmark import create_benchmark_model def test_claw_checks_aca_model() -> None: """An ill-typed argument to an `aca_model` function is rejected by beartype. - `create_model` annotates `n_subjects` as `int`; passing a string is caught - by the claw before the value reaches pylcm's own `Model` perimeter. + The benchmark factory requires a DiscreteGrid for preference types. + Invalid types are rejected before model construction. """ with pytest.raises(BeartypeCallHintViolation): - make_baseline_model(n_subjects="not an int") # ty: ignore[invalid-argument-type] + create_benchmark_model(pref_type_grid="not a grid") # ty: ignore[invalid-argument-type] diff --git a/tests/test_benchmark.py b/tests/test_benchmark.py index 56922ef..c3d54f6 100644 --- a/tests/test_benchmark.py +++ b/tests/test_benchmark.py @@ -22,7 +22,6 @@ def test_benchmark_model_builds_with_the_current_pylcm_grid_api() -> None: """The frozen benchmark preserves its state and action grid extents.""" model = create_benchmark_model( - n_subjects=1, pref_type_grid=DiscreteGrid(BenchmarkPrefType), ) @@ -52,7 +51,6 @@ def test_benchmark_model_builds_with_the_current_pylcm_grid_api() -> None: def test_benchmark_model_simulates_end_to_end() -> None: n_subjects = 20 model = create_benchmark_model( - n_subjects=n_subjects, pref_type_grid=DiscreteGrid(BenchmarkPrefType), ) _, _, params = get_benchmark_params(model=model) @@ -86,7 +84,6 @@ def test_benchmark_panel_exposes_hic_premium_and_wage_targets() -> None: """ n_subjects = 20 model = create_benchmark_model( - n_subjects=n_subjects, pref_type_grid=DiscreteGrid(BenchmarkPrefType), ) _, _, params = get_benchmark_params(model=model) @@ -148,7 +145,6 @@ def test_benchmark_simulate_obeys_borrowing_constraint() -> None: """ n_subjects = 4 model = create_benchmark_model( - n_subjects=n_subjects, pref_type_grid=DiscreteGrid(BenchmarkPrefType), ) _, _, params = get_benchmark_params(model=model) @@ -172,3 +168,19 @@ def test_benchmark_simulate_obeys_borrowing_constraint() -> None: f"borrowing_constraint violated on {int((slack < 0).sum())} row(s); " f"min slack = {slack.min():.6g}" ) + + +def test_initial_conditions_choose_population_size_after_model_construction() -> None: + """One model supports reproducible initial conditions with different row counts.""" + model = create_benchmark_model(pref_type_grid=DiscreteGrid(BenchmarkPrefType)) + small = get_benchmark_initial_conditions(model=model, n_subjects=2, seed=17) + repeated = get_benchmark_initial_conditions(model=model, n_subjects=2, seed=17) + large = get_benchmark_initial_conditions(model=model, n_subjects=5, seed=17) + + assert small.keys() == repeated.keys() == large.keys() + for name, values in small.items(): + assert values.shape == (2,) + assert large[name].shape == (5,) + np.testing.assert_array_equal(values, repeated[name]) + np.testing.assert_array_equal(small["age"], [51.0, 51.0]) + np.testing.assert_array_equal(small["claimed_ss"], [0, 0]) diff --git a/tests/test_dcegm_model_creation.py b/tests/test_dcegm_model_creation.py index 45274ff..6670eba 100644 --- a/tests/test_dcegm_model_creation.py +++ b/tests/test_dcegm_model_creation.py @@ -46,7 +46,6 @@ def _build_regimes(solver: SolverName) -> dict[str, Regime]: def _build_model(solver: SolverName) -> Model: return create_model( - n_subjects=1, fixed_params=_FIXED_PARAMS, wage_params=_WAGE_PARAMS, derived_categoricals=_DERIVED_CATEGORICALS, @@ -124,7 +123,6 @@ def test_dcegm_requires_construction_time_consumption_points() -> None: at model construction, so the runtime-injection path cannot be used.""" with pytest.raises(ValueError, match="consumption_dollars_points"): create_model( - n_subjects=1, fixed_params=_FIXED_PARAMS, wage_params=_WAGE_PARAMS, derived_categoricals=_DERIVED_CATEGORICALS, @@ -148,7 +146,6 @@ def test_benchmark_consumption_points_pin_both_floors() -> None: def test_dcegm_benchmark_model_builds() -> None: """The benchmark model accepts `solver="dcegm"` end to end.""" model = create_model( - n_subjects=1, fixed_params=_FIXED_PARAMS, wage_params=_WAGE_PARAMS, derived_categoricals=_DERIVED_CATEGORICALS, diff --git a/tests/test_dcegm_parity.py b/tests/test_dcegm_parity.py index 4c83e2c..70bfe05 100644 --- a/tests/test_dcegm_parity.py +++ b/tests/test_dcegm_parity.py @@ -114,7 +114,6 @@ def _make_model(*, solver: SolverName, grid_config: GridConfig) -> Model: else None ) return create_model( - n_subjects=N_SUBJECTS, fixed_params=fixed_params, wage_params=wage_params, derived_categoricals=_DERIVED_CATEGORICALS, diff --git a/tests/test_execution_config.py b/tests/test_execution_config.py index 04f19a7..77bd52d 100644 --- a/tests/test_execution_config.py +++ b/tests/test_execution_config.py @@ -25,7 +25,6 @@ def _factory(kind): if kind == "benchmark": return partial( create_benchmark_model, - n_subjects=1, pref_type_grid=DiscreteGrid(BenchmarkPrefType), ) fixed_params, wage_params, _ = get_benchmark_params(model=None) @@ -36,7 +35,6 @@ def _factory(kind): ) return partial( factory, - n_subjects=1, fixed_params=fixed_params, wage_params=wage_params, derived_categoricals={ @@ -115,3 +113,10 @@ def observe_pylcm_constructor(**kwargs): monkeypatch.setattr(aca_model_module, "Model", observe_pylcm_constructor) with pytest.raises(ConstructionObservedError): factory(execution_config=requested_policy) + + +@pytest.mark.parametrize("kind", ["baseline", "aca", "benchmark"]) +def test_factory_rejects_construction_subject_count(kind): + """Population size is supplied by initial conditions, not model construction.""" + with pytest.raises(TypeError, match="n_subjects"): + _factory(kind)(n_subjects=1) diff --git a/tests/test_initial_conditions_extreme_assets.py b/tests/test_initial_conditions_extreme_assets.py index 407b1cd..6455bee 100644 --- a/tests/test_initial_conditions_extreme_assets.py +++ b/tests/test_initial_conditions_extreme_assets.py @@ -8,8 +8,8 @@ """ import jax.numpy as jnp +from _lcm.simulation.initial_conditions import validate_initial_conditions from lcm import DiscreteGrid -from lcm.model import validate_initial_conditions from aca_model.agent.assets_and_income import borrowing_constraint from aca_model.agent.preferences import BenchmarkPrefType @@ -101,7 +101,6 @@ def test_extreme_negative_assets_subject_passes_validation() -> None: """ n_subjects = 1 model = create_benchmark_model( - n_subjects=n_subjects, pref_type_grid=DiscreteGrid(BenchmarkPrefType), ) _, _, params = get_benchmark_params(model=model) diff --git a/tests/test_model_creation.py b/tests/test_model_creation.py index 68b7d77..8e4398f 100644 --- a/tests/test_model_creation.py +++ b/tests/test_model_creation.py @@ -51,24 +51,24 @@ def build_regime(name: str): def test_model_creates_successfully() -> None: - model = make_baseline_model(n_subjects=1) + model = make_baseline_model() assert len(model.user_regimes) == 19 assert model.n_periods == 45 def test_model_age_range() -> None: - model = make_baseline_model(n_subjects=1) + model = make_baseline_model() assert model.ages.values[0] == 51.0 assert model.ages.values[-1] == 95.0 def test_dead_regime_is_terminal() -> None: - model = make_baseline_model(n_subjects=1) + model = make_baseline_model() assert model.user_regimes["dead"].terminal def test_non_terminal_regimes_not_terminal() -> None: - model = make_baseline_model(n_subjects=1) + model = make_baseline_model() for name in REGIME_SPECS: assert not model.user_regimes[name].terminal @@ -233,7 +233,7 @@ def test_all_non_terminal_regimes_carry_pension_wealth_as_carried_state() -> Non build_model_state_transitions()["pension_wealth"] is pensions.wealth_next_before_adjustment ) - model = make_baseline_model(n_subjects=1) + model = make_baseline_model() for name in REGIME_SPECS: assert isinstance(model.user_regimes[name].states["pension_wealth"], Phased), ( name @@ -281,7 +281,7 @@ def test_hcc_persistent_and_transitory_are_shock_grids() -> None: def test_aca_model_creates_successfully() -> None: - model = make_aca_model(n_subjects=1, policy=PolicyVariant.ACA) + model = make_aca_model(policy=PolicyVariant.ACA) assert len(model.user_regimes) == 19 assert model.n_periods == 45 @@ -322,7 +322,7 @@ def test_aca_other_regimes_have_no_aca_policy_keys() -> None: @pytest.mark.parametrize("policy", list(PolicyVariant)) def test_all_policy_variants_create(policy: PolicyVariant) -> None: """All policy variants create valid models.""" - model = make_aca_model(n_subjects=1, policy=policy) + model = make_aca_model(policy=policy) assert len(model.user_regimes) == 19 @@ -362,7 +362,7 @@ def test_aca_only_medicaid_expansion() -> None: def test_baseline_model_creates() -> None: """Baseline model creates successfully without PolicyVariant.""" - model = make_baseline_model(n_subjects=1) + model = make_baseline_model() assert len(model.user_regimes) == 19 @@ -371,7 +371,7 @@ def test_dead_regime_prunes_unused_broadcast_states() -> None: `dead` keeps only what the bequest DAG reads (`assets`, `pref_type`). `pension_wealth` is masked (carried states are illegal in terminal regimes); the other unused broadcast states are pruned by reachability.""" - model = make_baseline_model(n_subjects=1) + model = make_baseline_model() assert model.pruned_variables["dead"] == frozenset( {"aime", "spousal_income", "hcc_persistent", "hcc_transitory"} ) @@ -381,6 +381,6 @@ def test_dead_regime_prunes_unused_broadcast_states() -> None: def test_living_regimes_keep_every_broadcast_state() -> None: """Every model-level state is read by each living regime's DAG, so pruning removes nothing outside `dead`.""" - model = make_baseline_model(n_subjects=1) + model = make_baseline_model() for name in REGIME_SPECS: assert model.pruned_variables[name] == frozenset() diff --git a/tests/test_nbegm_labor_live_validation.py b/tests/test_nbegm_labor_live_validation.py index 279c6fd..62b4f9e 100644 --- a/tests/test_nbegm_labor_live_validation.py +++ b/tests/test_nbegm_labor_live_validation.py @@ -38,7 +38,6 @@ def _solve_m1(solver: SolverName) -> tuple[dict[int, np.ndarray], int]: ) fixed_params, wage_params, _ = get_benchmark_params(model=None) model = create_model( - n_subjects=1, fixed_params=fixed_params, wage_params=wage_params, derived_categoricals=_DERIVED_CATEGORICALS, diff --git a/tests/test_nbegm_model_creation.py b/tests/test_nbegm_model_creation.py index 79d0265..f853fcc 100644 --- a/tests/test_nbegm_model_creation.py +++ b/tests/test_nbegm_model_creation.py @@ -66,7 +66,6 @@ def _build_regimes(solver: SolverName) -> dict[str, Regime]: def _build_model_with(solver: SolverName, grid_config: GridConfig) -> Model: return create_model( - n_subjects=1, fixed_params=_FIXED_PARAMS, wage_params=_WAGE_PARAMS, derived_categoricals=_DERIVED_CATEGORICALS, @@ -300,7 +299,6 @@ def test_nbegm_builds_every_aca_policy_variant(policy: PolicyVariant) -> None: compose with the branch compiler's per-regime wiring.""" grid_config = _BRIDGED_GRID_CONFIG model = create_aca_model( - n_subjects=1, policy=policy, fixed_params=_FIXED_PARAMS, wage_params=_WAGE_PARAMS, @@ -334,7 +332,6 @@ def test_nbegm_aca_variants_leave_no_free_buy_private_params( """ grid_config = _BRIDGED_GRID_CONFIG model = create_aca_model( - n_subjects=1, policy=policy, fixed_params=_FIXED_PARAMS, wage_params=_WAGE_PARAMS, diff --git a/tests/test_nbegm_solve_validation.py b/tests/test_nbegm_solve_validation.py index 4c840da..e8e6484 100644 --- a/tests/test_nbegm_solve_validation.py +++ b/tests/test_nbegm_solve_validation.py @@ -35,7 +35,6 @@ def _solve_m1(solver: SolverName) -> dict[int, np.ndarray]: grid_config = dataclasses.replace(BENCHMARK_GRID_CONFIG, nbegm_jump_read="bridged") fixed_params, wage_params, _ = get_benchmark_params(model=None) model = create_model( - n_subjects=1, fixed_params=fixed_params, wage_params=wage_params, derived_categoricals=_DERIVED_CATEGORICALS, From 3bc89f74572fe48c9ff5401bf022bf585b47d02b Mon Sep 17 00:00:00 2001 From: Hans-Martin von Gaudecker Date: Sat, 26 Sep 2026 12:27:34 +0200 Subject: [PATCH 16/16] Floor leisure at a small positive value so float32 utility stays finite (#18) Co-authored-by: Claude Opus 5.5 --- .github/workflows/main.yml | 8 +- src/aca_model/agent/preferences.py | 33 +++-- tests/test_leisure_floor.py | 190 +++++++++++++++++++++++++++++ 3 files changed, 218 insertions(+), 13 deletions(-) create mode 100644 tests/test_leisure_floor.py diff --git a/.github/workflows/main.yml b/.github/workflows/main.yml index 632668d..897c37f 100644 --- a/.github/workflows/main.yml +++ b/.github/workflows/main.yml @@ -30,12 +30,12 @@ jobs: - uses: actions/setup-python@v6 with: python-version: ${{ matrix.python-version }} - - name: Install pylcm from feat/nb-egm - # This slice depends on the NB-EGM solver, which lives on pylcm's - # feat/nb-egm branch; revert to @main once that branch merges. + - name: Install pylcm at the revision aca-dev pins + # Pinned by commit so a deleted or rewritten branch cannot break the + # install; move it together with the aca-dev pylcm submodule pointer. run: >- pip install "pylcm @ - git+https://github.com/OpenSourceEconomics/pylcm.git@feat/nb-egm" + git+https://github.com/OpenSourceEconomics/pylcm.git@7e8e4c223a63a07725015ebdd5af84cd9092cee2" - name: Install aca-model with test deps run: pip install -e . pytest pdbp - name: Run pytest diff --git a/src/aca_model/agent/preferences.py b/src/aca_model/agent/preferences.py index 438b2c4..3bb4b80 100644 --- a/src/aca_model/agent/preferences.py +++ b/src/aca_model/agent/preferences.py @@ -23,6 +23,13 @@ # endowment; it only bends the map near and beyond the endowment. _LEISURE_SMOOTHING_FRACTION = 0.01 +# Positive floor that leisure approaches as work costs exceed the endowment, as a +# fraction of the time endowment. It keeps felicity and its consumption derivative +# inside float32's range when leisure_available is far below zero, down to the +# NB-EGM inverse bracket's lowest consumption; away from the floor its effect decays +# exponentially. +_LEISURE_FLOOR_FRACTION = 1e-5 + @categorical(ordered=False) class PrefType: @@ -71,17 +78,25 @@ def fixed_cost_of_work( def _smooth_leisure_floor( leisure_available: FloatND, time_endowment: ScalarFloat ) -> FloatND: - """Bend leisure to a strictly positive floor as work costs approach the endowment. - - `softplus(x) = log(1 + e^x)` via `jnp.logaddexp(0, x)`, scaled by a small fraction - of the endowment. Where `leisure_available` is large relative to the smoothing width - the map reduces to `leisure_available` (bulk unchanged); as it falls to zero leisure - bends to `0⁺` — never negative, never a kinked clamp — so the CRRA aggregator never - receives a non-positive base. The smoothing width scales with the endowment, so the - map is scale-invariant. + """Bend leisure smoothly onto a positive floor as work costs approach the endowment. + + Leisure is a smooth maximum of `leisure_available` and the floor + `F = _LEISURE_FLOOR_FRACTION * time_endowment`: + `s * logaddexp(F / s, leisure_available / s)`, with smoothing width + `s = _LEISURE_SMOOTHING_FRACTION * time_endowment`. + + - `leisure_available` large relative to `s`: leisure equals `leisure_available` + up to a term that decays like `e^(-leisure_available / s)`. + - `leisure_available` at or below zero: leisure bends to `F⁺` — never below the + floor, never a kinked clamp — so every work choice stays feasible and the CRRA + aggregator's base stays large enough for float32 felicity and marginal + felicity. + + Both widths scale with the endowment, so the map is scale-invariant. """ smoothing = _LEISURE_SMOOTHING_FRACTION * time_endowment - return smoothing * jnp.logaddexp(0.0, leisure_available / smoothing) + floor = _LEISURE_FLOOR_FRACTION * time_endowment + return smoothing * jnp.logaddexp(floor / smoothing, leisure_available / smoothing) def leisure_canwork_retiree_or_nongroup( diff --git a/tests/test_leisure_floor.py b/tests/test_leisure_floor.py new file mode 100644 index 0000000..2c1e470 --- /dev/null +++ b/tests/test_leisure_floor.py @@ -0,0 +1,190 @@ +"""Float32 felicity at the leisure floor, against a 60-digit reference. + +The class covers the canwork cells at production preference values: every pref +type, health, lagged labor supply, hours choice and canwork age, at consumption +points spanning the NB-EGM numeric-inverse bracket (`1e-8` to the top of the +action range). Marginal utility is `jax.grad` of the consumption-dollar +felicity, the route pylcm's NB-EGM uses. +""" + +import itertools +from decimal import Decimal, getcontext + +import jax +import jax.numpy as jnp +import numpy as np +import pytest + +from aca_model.agent import preferences + +getcontext().prec = 60 + +CONSUMPTION_WEIGHTS = (0.6776541629520845, 0.8805772328591686, 0.0718086283445225) +COEFFICIENTS_RRA = (3.841252231680976, 0.9990771146810682, 3.8328505891095137) +TIME_ENDOWMENT = 3926.9478390365557 +BAD_HEALTH_COST = 408.9190313043897 +REENTRY_COST = 119.61095526191421 +FIXED_COST_INTERCEPT = 337.5223349543413 +FIXED_COST_AGE_TREND = 84.09242960636178 +REFERENCE_AGE = 50 +REFERENCE_HOURS = 1000.0 +AVERAGE_CONSUMPTION = 20000.0 +HOURS = (0.0, 1000.0, 1500.0, 2000.0, 2500.0) +CANWORK_AGES = tuple(range(51, 72)) +CONSUMPTION = (1e-8, 1597.0921419521899, 20000.0, 300_000.0, 721271853.846792) +FLOAT32_REL_TOL = 1e-4 +# Leisure-map widths as fractions of the endowment: smoothing width and floor. +SMOOTHING_FRACTION = 0.01 +FLOOR_FRACTION = 1e-5 + + +def _f32(x: float) -> jnp.ndarray: + return jnp.asarray(x, dtype=jnp.float32) + + +def _felicity_of_dollars(consumption, leisure, weight, rra, scale): + return preferences.u_alive( + consumption_equiv=preferences.consumption_equiv( + consumption_dollars=consumption, + equivalence_scale=jnp.ones_like(consumption), + ), + leisure=leisure, + consumption_weight=weight, + coefficient_rra=rra, + utility_scale_factor=scale, + ) + + +def _leisure_available(good: int, lagged: int, hours: float, age: int) -> float: + fixed_cost = FIXED_COST_INTERCEPT + FIXED_COST_AGE_TREND * (age - REFERENCE_AGE) + reentry = REENTRY_COST if lagged == 0 else 0.0 + work = hours + fixed_cost + reentry if hours > 0 else 0.0 + return TIME_ENDOWMENT - (0.0 if good else BAD_HEALTH_COST) - work + + +def _cells() -> list[tuple[float, float]]: + """All `(leisure_available, consumption)` pairs of the class.""" + return [ + (_leisure_available(good, lagged, hours, age), c) + for good, lagged, hours, age, c in itertools.product( + (0, 1), (0, 1), HOURS, CANWORK_AGES, CONSUMPTION + ) + ] + + +def _float32_u_and_marginal( + pref_type: int, cells: list[tuple[float, float]] +) -> tuple[np.ndarray, np.ndarray]: + available = jnp.asarray([a for a, _ in cells], dtype=jnp.float32) + consumption = jnp.asarray([c for _, c in cells], dtype=jnp.float32) + weight = _f32(CONSUMPTION_WEIGHTS[pref_type]) + rra = _f32(COEFFICIENTS_RRA[pref_type]) + leisure = _floored_leisure(available, _f32(TIME_ENDOWMENT)) + scale = preferences.utility_scale_factor( + average_consumption_equiv=_f32(AVERAGE_CONSUMPTION), + consumption_weight=weight, + coefficient_rra=rra, + time_endowment=_f32(TIME_ENDOWMENT), + fixed_cost_of_work_intercept=_f32(FIXED_COST_INTERCEPT), + reference_hours=_f32(REFERENCE_HOURS), + ) + args = (weight, rra, scale) + u = jax.vmap(_felicity_of_dollars, in_axes=(0, 0, None, None, None))( + consumption, leisure, *args + ) + marginal = jax.vmap( + jax.grad(_felicity_of_dollars), in_axes=(0, 0, None, None, None) + )(consumption, leisure, *args) + assert u.dtype == marginal.dtype == jnp.float32 + return np.asarray(u, dtype=np.float64), np.asarray(marginal, dtype=np.float64) + + +def _floored_leisure(available, endowment): + """Production leisure at a given `leisure_available`, via the tied-regime map. + + Good health and no fixed cost make the work loss the only deduction, so hours of + `endowment - available` leave exactly `available`. + """ + return preferences.leisure_canwork_tied( + working_hours_value=endowment - available, + good_health=jnp.ones(available.shape, dtype=jnp.int32), + time_endowment=endowment, + leisure_cost_of_bad_health=jnp.zeros_like(endowment), + fixed_cost_of_work=jnp.zeros_like(endowment), + ) + + +def _d(x: float) -> Decimal: + return Decimal(repr(float(x))) + + +def _reference_u_and_marginal( + pref_type: int, cells: list[tuple[float, float]] +) -> tuple[np.ndarray, np.ndarray]: + """Closed forms in 60-digit log space, independent of the JAX code path.""" + weight, rra = _d(CONSUMPTION_WEIGHTS[pref_type]), _d(COEFFICIENTS_RRA[pref_type]) + endowment = _d(TIME_ENDOWMENT) + smoothing = _d(SMOOTHING_FRACTION) * endowment + floor = _d(FLOOR_FRACTION) * endowment + average_leisure = endowment - _d(REFERENCE_HOURS) - _d(FIXED_COST_INTERCEPT) + log_average = weight * _d(AVERAGE_CONSUMPTION).ln() + (1 - weight) * ( + average_leisure.ln() + ) + scale = abs((1 - rra) / ((1 - rra) * log_average).exp()) + u, marginal = [], [] + for available, c in cells: + leisure = ( + smoothing + * ((floor / smoothing).exp() + (_d(available) / smoothing).exp()).ln() + ) + log_composite = weight * _d(c).ln() + (1 - weight) * leisure.ln() + power = ((1 - rra) * log_composite).exp() + u.append(float(scale * power / (1 - rra))) + marginal.append(float(scale * weight / _d(c) * power)) + return np.array(u), np.array(marginal) + + +OUTPUTS = {"u": 0, "marginal": 1} +WITNESS = [(_leisure_available(0, 0, 2500.0, 71), c) for c in CONSUMPTION] + + +@pytest.mark.parametrize("output", ["u", "marginal"]) +def test_float32_felicity_at_witness_cell_matches_reference(output: str) -> None: + """Pref type 2, bad health, did not work, 2500 h, age 71: finite and accurate.""" + got = _float32_u_and_marginal(2, WITNESS)[OUTPUTS[output]] + ref = _reference_u_and_marginal(2, WITNESS)[OUTPUTS[output]] + np.testing.assert_allclose(got, ref, rtol=FLOAT32_REL_TOL) + + +@pytest.mark.parametrize("output", ["u", "marginal"]) +@pytest.mark.parametrize("pref_type", [0, 1, 2]) +def test_float32_felicity_over_canwork_class_matches_reference( + pref_type: int, output: str +) -> None: + """Every canwork cell and consumption point: float32 equals the reference.""" + cells = _cells() + got = _float32_u_and_marginal(pref_type, cells)[OUTPUTS[output]] + ref = _reference_u_and_marginal(pref_type, cells)[OUTPUTS[output]] + np.testing.assert_allclose(got, ref, rtol=FLOAT32_REL_TOL) + + +@pytest.mark.parametrize("pref_type", [0, 1, 2]) +def test_leisure_floor_leaves_utility_unchanged_away_from_the_floor( + pref_type: int, +) -> None: + """With leisure_available above 10% of the endowment, u moves by at most 1e-7. + + Compared against the unfloored softplus `s * log(1 + e^(x/s))`. + """ + available = np.array([a for a, _ in _cells() if a > 0.1 * TIME_ENDOWMENT]) + assert available.size > 0 + smoothing = SMOOTHING_FRACTION * TIME_ENDOWMENT + unfloored = smoothing * np.logaddexp(0.0, available / smoothing) + floored = np.asarray( + _floored_leisure(jnp.asarray(available), jnp.asarray(TIME_ENDOWMENT)) + ) + exponent = (1.0 - CONSUMPTION_WEIGHTS[pref_type]) * ( + 1.0 - COEFFICIENTS_RRA[pref_type] + ) + relative_change = np.abs((floored / unfloored) ** exponent - 1.0) + np.testing.assert_array_less(relative_change, 1e-7)