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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 4 additions & 4 deletions .github/workflows/main.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
40 changes: 40 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
@@ -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.
9 changes: 8 additions & 1 deletion src/aca_model/__init__.py
Original file line number Diff line number Diff line change
@@ -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.
Expand Down
15 changes: 11 additions & 4 deletions src/aca_model/aca/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,19 +7,19 @@
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
from aca_model.aca.regimes import build_all_regimes
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(
*,
n_subjects: int,
policy: PolicyVariant,
fixed_params: UserParams,
wage_params: Mapping[str, Any],
Expand All @@ -28,11 +28,11 @@ 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.

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
Expand All @@ -53,6 +53,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.
Expand Down Expand Up @@ -92,6 +95,10 @@ def create_model(
description=f"Structural retirement model ({policy.name})",
fixed_params=fixed_params,
derived_categoricals=derived_categoricals,
n_subjects=n_subjects,
execution_config=(
execution_config_for_devices()
if execution_config is None
else execution_config
),
**model_slots,
)
33 changes: 24 additions & 9 deletions src/aca_model/agent/preferences.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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(
Expand Down
17 changes: 12 additions & 5 deletions src/aca_model/baseline/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,15 +5,15 @@

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)
"""

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 (
Expand All @@ -23,23 +23,23 @@
build_model_slots,
)
from aca_model.config import MODEL_CONFIG, GridConfig
from aca_model.execution import execution_config_for_devices


def create_model(
*,
n_subjects: int,
fixed_params: UserParams,
wage_params: Mapping[str, Any],
derived_categoricals: Mapping[str, DiscreteGrid],
grid_config: GridConfig,
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.

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;
Expand Down Expand Up @@ -69,6 +69,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
Expand Down Expand Up @@ -106,7 +109,11 @@ def create_model(
description="Baseline structural retirement model (pre-ACA)",
fixed_params=fixed_params,
derived_categoricals=derived_categoricals,
n_subjects=n_subjects,
execution_config=(
execution_config_for_devices()
if execution_config is None
else execution_config
),
**model_slots,
)

Expand Down
11 changes: 3 additions & 8 deletions src/aca_model/baseline/regimes/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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
Expand All @@ -136,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,
Expand All @@ -146,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(),
}
Loading
Loading