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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
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
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
190 changes: 190 additions & 0 deletions tests/test_leisure_floor.py
Original file line number Diff line number Diff line change
@@ -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)
Loading