diff --git a/docs/source/development/internals/batching.rst b/docs/source/development/internals/batching.rst new file mode 100644 index 00000000..cc48a30e --- /dev/null +++ b/docs/source/development/internals/batching.rst @@ -0,0 +1,355 @@ +.. _batching_internals: + +Batching: Internals, Alternatives, and Open Questions +======================================================= + +.. warning:: + + This page is AI-generated (drafted with Claude, based on reading the source and a + design discussion, without access to GPU hardware to verify the performance and + memory claims empirically). It is intended as internal documentation for + contributors and as context for other AI coding agents working on this codebase, + not as a peer-reviewed reference. Claims marked as unverified or plausible in the + text below have not been checked against real hardware; verify before relying on + them for a decision. + +.. note:: + + This page is for contributors working on the solver itself. If you only want to + *configure* batching for your model, see the :ref:`batching_guide` in the + background section instead. This page explains the mechanism behind that + configuration, discusses algorithmic alternatives, and relates the design to the + wider computational-science literature on batching irregular workloads. + +Why batching exists +-------------------- + +Backward induction in ``dcegm`` solves one period at a time, from the terminal +period backward, and each period's continuation values depend on the (already +solved) next period. The number of feasible state-choice combinations changes over +the life cycle -- e.g. a choice becomes unavailable, or a deterministic state stops +being reachable -- so the *natural* unit of work (one period) does not have a fixed +size. + +That would not matter for a plain Python loop, but the backward induction loop in +``dcegm`` is implemented as a single :func:`jax.lax.scan` per segment +(:func:`dcegm.backward_induction.backward_induction`). ``lax.scan`` stacks its +per-step inputs into one array and traces the step function *once*, so every step +must consume an array slice of the *same shape*. Batching exists to reconcile these +two facts: it repackages a horizon of unevenly-sized periods into a sequence of +equal-sized chunks that ``lax.scan`` can iterate over, while preserving the +dependency order backward induction requires (a state-choice's children must be +solved in a strictly earlier batch). + +How it is currently done +------------------------- + +All batching logic lives under ``src/dcegm/pre_processing/batches/``. The user-facing +entry point is :func:`dcegm.pre_processing.batches.batch_creation.create_batches_and_information`, +which splits the horizon into one or more *segments* (see below), and delegates the +construction of batches within each segment to +:func:`dcegm.pre_processing.batches.single_segment.create_single_segment_of_batches`. + +Within a segment, ``dcegm`` supports two modes: + +``largest_block`` + Implemented in + :func:`dcegm.pre_processing.batches.algo_batch_size.determine_optimal_batch_size`. + All eligible state-choices in the segment are sorted once, ascending, by the + minimum raw state-choice index of their child states. This is the only ordering + step -- see the open question below. The sorted sequence is then reversed and + split into contiguous chunks of size ``current_batch_size``, starting from + ``current_batch_size = size_last_period`` (the number of state-choices in the + segment's last period). + + Each candidate size is checked for validity: for every chunk, the maximum raw + state-choice index in the chunk must be *smaller* than the minimum raw + state-choice index among that chunk's required children. If this fails for any + chunk, ``current_batch_size`` is shrunk by 2% (``current_batch_size = int + (current_batch_size * 0.98)``) and *every* chunk is re-validated from scratch. + This repeats until a valid uniform size is found. In pseudocode: + + .. code-block:: text + + sort state-choices ascending by min(child state-choice index) + current_batch_size = size_last_period + loop: + chunks = split(reverse(sorted_state_choices), current_batch_size) + if all(chunk.max_index < min(child_indices(chunk)) for chunk in chunks): + return chunks + current_batch_size = int(current_batch_size * 0.98) + + This is a correctness-constrained search for a large uniform batch size, not a + padding-vs-waste tradeoff: because ``lax.scan`` reuses one compiled step + regardless of how many iterations it runs, a larger valid batch size is always + preferable within a segment (fewer scan steps, same compiled kernel, no extra + compile cost). See :ref:`batching_alternatives` for a cheaper way to find it. + +``period_max`` + Implemented in ``determine_period_max_batch_size`` in the same module. Each + period within the segment becomes exactly one batch. Batches are padded to the + segment's largest per-period state-choice count with a deterministic dummy + state-choice index (the first valid one in the same batch), so padding never + changes the solution -- it only wastes compute on the padded slots. + +Segmenting the horizon +~~~~~~~~~~~~~~~~~~~~~~~ + +``min_period_batch_segments`` (handled in ``batch_creation.py``) splits the horizon +into multiple segments *before* either mode above runs. In +:func:`dcegm.backward_induction.backward_induction`, each segment gets its own +``jax.lax.scan`` call, via a plain Python ``for id_segment in range(n_segments):`` +loop -- ``n_segments`` is a static Python int fixed at model-setup time, not a traced +value. This is the mechanism that lets different parts of the life cycle use +different batch modes or (implicitly, via ``largest_block``) different uniform batch +sizes: a single global batch size is bottlenecked by the most dependency-dense +region of the horizon, and segmenting lets the rest of the horizon avoid that +bottleneck. + +What that costs depends on *how* ``backward_induction`` is called, and this matters +because the two call sites behave differently: + +- :meth:`~dcegm.interfaces.model_class.setup_model.solve` calls it eagerly, with no + enclosing ``jax.jit``. Each segment's ``lax.scan`` is then traced and compiled as + its own separate XLA computation, dispatched one after another. +- :meth:`~dcegm.interfaces.model_class.setup_model.get_solve_func` (the recommended + path for repeated solves, e.g. inside :mod:`dcegm.likelihood`, which follows the + same pattern) wraps the *whole* ``backward_induction`` call -- Python loop included + -- inside one outer ``jax.jit``. Because ``n_segments`` is static, tracing unrolls + that loop at trace time into :math:`n_{\text{segments}}` sequential ``lax.scan`` + primitives inside a single jaxpr, which XLA then compiles as **one** executable + containing :math:`n_{\text{segments}}` back-to-back scan sub-computations -- not as + separate executables. + +Wrapping everything in one outer ``jax.jit`` is not free (see the warning below), so +it is worth being explicit about what it buys, since compiling each segment +separately would also get you the standard JAX benefit of "compile once on the first +call, reuse the executable for every later ``params``" -- that part is not what +distinguishes the two designs. What the *single* enclosing jit adds on top: + +- **No host round-trip between segments.** Under eager ``.solve()``, each segment's + ``lax.scan`` is a separate dispatch: Python-side argument/pytree handling and a + host-device synchronization boundary between segments. On GPU, dispatch and kernel + launch latency is frequently the actual bottleneck for a sequence of small-to-medium + ops, more so than raw compute throughput -- collapsing :math:`n_{\text{segments}}` + dispatches into one removes that overhead entirely for everything after the first + compile. +- **Whole-program scheduling.** XLA's scheduler and buffer-assignment see the entire + unrolled sequence at once and can reorder or overlap work across segment boundaries + where data dependencies allow. :math:`n_{\text{segments}}` separately-compiled + programs are each optimized in isolation, with zero visibility past their own jit + boundary. + +.. warning:: + + Even under the single-``jax.jit`` path, more segments plausibly still means more + memory, but the mechanism is different from "many separate compiled programs + competing for device memory": + + - **Compile-time/host cost (the solid part).** More segments means a larger + unrolled jaxpr/HLO program for XLA to schedule, buffer-assign, and (if enabled) + autotune. Compile time and the host RAM used *during* compilation both grow with + program size -- a general, well-documented property of XLA for large unrolled + graphs, not specific to ``dcegm``. + - **Device/runtime cost (plausible, not verified here).** XLA's buffer-assignment + does whole-program liveness analysis, so in principle it can still reuse buffers + *across* segment boundaries within that one compiled program -- the earlier + framing of "no reuse across separate executables" does not apply here, because + there is only one executable. What can still block reuse is a *shape change* at + a segment boundary: the ``value``/``policy``/``endog_grid`` arrays threaded as + the scan carry from one segment into the next only alias cleanly in place when + shapes match, and adjacent segments deliberately using different batch sizes or + modes is the entire point of segmenting. So the more precise (but unverified) + claim is that the number of *segment boundaries* -- not the number of segments + or executables as such -- is what may force extra allocations instead of + in-place reuse. + - **Batch-index metadata.** ``batch_info`` holds every segment's index arrays + (``batches_state_choice_idx`` and friends, built by + ``prepare_and_align_batch_arrays``) simultaneously, since all segments are + constructed up front before ``backward_induction`` runs. More segments means + more such arrays resident at once, though this is likely small next to the + solution containers. + + None of this is currently measured in ``dcegm``. Confirming it, and by how much, + would need ``jax.devices()[0].memory_stats()["peak_bytes_in_use"]`` (with + ``XLA_PYTHON_CLIENT_PREALLOCATE=false``, since JAX preallocates most GPU memory by + default and would otherwise hide the effect) compared across a fixed model solved + through :meth:`~dcegm.interfaces.model_class.setup_model.get_solve_func` with a + varying number of segments. See the cost model below, which currently only + accounts for compile *time*, not memory. + +Today, segment boundaries and per-segment modes are chosen by hand (see the +:ref:`batching_guide` for the recommended workflow using +``get_n_state_choices_per_period``). This manual step is one of the main things +:ref:`batching_alternatives` below tries to replace. + +Relation to the wider literature +--------------------------------- + +``dcegm``'s batching is an instance of a problem that shows up, with different names, +across computational science whenever irregular workloads need to run on hardware +that wants fixed shapes. Three literatures map onto the two mechanisms above: + +- **Dependency-safe batch construction** (``largest_block``'s validity check) is what + parallel computing calls *wavefront* or *level-scheduled* parallelism: group a DAG's + nodes into levels such that everything in a level depends only on already-computed + levels. It is classically used for sequence alignment and PDE stencils on GPUs. + + - Kartik Hegde et al. `Memory-Optimized Wavefront Parallelism on GPUs `_. *International Journal of Parallel Programming* (2020). + - `Taskflow: Wavefront Parallelism `_ (pedagogical overview of the pattern). + +- **Sizing batches under a hard uniformity constraint** is the classical *multiprocessor + scheduling* / *bin-packing* problem: assign items to a fixed number of equal-capacity + bins to minimize the number of bins (equivalently, maximize bin size). The Longest + Processing Time (LPT) heuristic gives a 4/3-approximation guarantee. + + - E. G. Coffman, M. R. Garey, D. S. Johnson (1978). `An Application of Bin-Packing to Multiprocessor Scheduling `_. *SIAM Journal on Computing*. + +- **Padding variable-size groups to a fixed shape** (``period_max``) is the same + mechanism as sequence bucketing in batched ML inference, including the same + padding-vs-compute-waste tradeoff and the same underlying cause: neither XLA nor + most deep-learning runtimes have first-class ragged-array support. + + - `Continuous batching from first principles `_. Hugging Face (2025). + - `Support for ragged arrays, like torch.nested `_. jax-ml/jax issue tracker. + +- For the domain itself, GPU-batched Bellman backward induction shows up directly in + recent economics/OR work, useful as motivation and cross-validation of the general + approach rather than as an algorithmic source: + + - `GPU-Accelerated Dynamic Programming for Multistage Stochastic Energy Storage Arbitrage `_. + - `Structural Reinforcement Learning for Heterogeneous Agent Macroeconomics `_. Moll et al. + +.. _batching_alternatives: + +Alternatives worth considering +-------------------------------- + +Splitting the current design into its two decisions clarifies which improvements are +"free" and which involve a genuine, hardware-dependent tradeoff. + +Within a segment: replace the multiplicative-decrease search +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +The validity check used by ``largest_block`` is monotonic in batch size: a larger +uniform batch size spans a wider range of raw indices, which can only make the +"children already solved" check harder to satisfy, never easier. That means the set +of valid batch sizes is downward-closed, and the true maximum can be found by +**binary search** instead of shrinking by 2% and re-validating every chunk from +scratch at every step: + +.. code-block:: text + + function max_valid_batch_size(state_choices_sorted_by_child_idx): + lo, hi = 1, len(state_choices_sorted_by_child_idx) + best = lo + while lo <= hi: + mid = (lo + hi) // 2 + if feasible(state_choices_sorted_by_child_idx, mid): + best = mid + lo = mid + 1 + else: + hi = mid - 1 + return best + +``feasible`` is the same per-chunk check already implemented today; only the search +strategy changes, from O(chunks-per-attempt :math:`\times` attempts-to-converge) with +an arbitrary 2%-undershoot, to O(chunks-per-attempt :math:`\times \log n`) landing on +the true maximum. This is a strict improvement, not a hardware-dependent tuning +choice -- it should be paired with an equivalence test (brute-force linear scan vs. +binary search on a couple of existing toy models) before being trusted, both to +confirm monotonicity holds in practice and following the project's convention of +pairing solver-internals refactors with a golden-value/equivalence test. + +Across segments: a calibrated cost model instead of a manual knob +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +Choosing how many segments to use, and where to split them, *is* genuinely +hardware- and workload-specific, because each extra segment buys a larger batch size +for the rest of the horizon at the cost of one extra XLA compile. Formalizing this +needs two device- and workload-specific constants: + +- :math:`C_{\text{compile}}`: fixed cost of one additional segment (a separate + ``lax.scan`` compilation). Matters most for a single ``solve()`` call; matters much + less inside an estimation loop that reuses the compiled function across many + parameter draws. +- :math:`C_{\text{step}}`: marginal cost per scan iteration (kernel dispatch plus + actual FLOPs). Shrinks, relative to a fixed batch size, the more GPU-bound the + workload is. + +Given those, choosing segment boundaries becomes a shortest-path / dynamic-programming +problem over candidate splits of the period axis: + +.. code-block:: text + + function best_segmentation(periods, C_compile, C_step): + n = len(periods) + cost = [0] + [infinity] * n + split_at = [None] * (n + 1) + for j in 1..n: + for i in 0..j-1: + B = max_valid_batch_size(periods[i:j]) + n_scan_steps = ceil(size(periods[i:j]) / B) + segment_cost = C_compile + n_scan_steps * C_step + if cost[i] + segment_cost < cost[j]: + cost[j] = cost[i] + segment_cost + split_at[j] = i + return backtrack(split_at) + +:math:`C_{\text{compile}}` and :math:`C_{\text{step}}` should be measured on the +target device (time a representative ``lax.scan`` compile, and a few iterations at +different batch sizes to fit the marginal per-step cost) rather than guessed. With +:math:`C_{\text{compile}} \to 0` the optimum drifts toward many small segments (one +period per block, minimal batch-size compromise); with a large +:math:`C_{\text{compile}}` it collapses to a single segment (today's default when +``min_period_batch_segments`` is not set). This also explains *why* segmenting ever +helps: a single global batch size is bottlenecked by the horizon's most +dependency-dense region, and paying for an extra compile lets the rest of the horizon +escape that bottleneck. + +As modeled above, :math:`C_{\text{compile}}` only represents compile *time* -- but, +per the warning above, each extra segment also has a device-memory cost from holding +more separately-compiled executables at once. Time is a cost worth trading off +against faster steps; memory is closer to a hard constraint, since exceeding it means +the solve fails outright rather than merely running slower. A more faithful version +of this DP would therefore track a *memory budget* alongside the additive time cost +(minimize total time subject to peak memory across live segments staying under the +device limit) rather than folding both into one scalar cost -- but that requires +measuring the per-segment memory cost first, which nothing in ``dcegm`` does today. + +A further-out alternative for ``period_max`` +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +``period_max``'s padding could in principle be avoided altogether by flattening all +periods into one array and using masked/segmented reductions (JAX's +``jax.ops.segment_sum``-style primitives) instead of materializing padded batches -- +the same idea as "continuous" or "packed" batching in ML serving. The real tradeoff +is that if batch shape varies per scan step, ``lax.scan`` can no longer reuse one +compiled step, trading padding waste against recompilation cost. Whether that trade +is worth it needs profiling on representative model sizes; it is not a clear win +either way and is listed here as a direction, not a recommendation. + +Open question: does the state-choice ordering matter? +-------------------------------------------------------- + +Before splitting into chunks, ``largest_block`` sorts the eligible state-choices by a +single, fixed key -- ascending minimum raw index of their child states -- and never +searches over alternative orderings. This ordering determines how tightly packed a +state-choice's dependencies are relative to its own position, which in turn +determines how large a uniform batch can get before the validity check fails. + +It is currently unknown whether this particular ordering is a good one, or even a +reasonable default, relative to the alternatives. The closest known analogue is +**bandwidth minimization / fill-reducing reordering** in sparse linear algebra +(reverse Cuthill-McKee, nested dissection), where the entire point of choosing an +ordering is to keep each row's dependencies as close as possible to the row itself, +which directly shrinks the "distance" a dependency-safe block needs to span. If the +same intuition transfers here, a bandwidth-minimizing ordering might permit +substantially larger valid batch sizes than the current ascending-child-index sort +achieves -- but this has not been tested, and it is also possible that the current +sort is already close to optimal for typical ``dcegm`` life-cycle structures (where +dependencies are inherently local, since a state-choice's children mostly live in the +immediately next period). Resolving this would need a small experiment: compare the +maximum valid batch size (via :ref:`batching_alternatives`'s binary search) under the +current sort against a bandwidth-minimizing reordering, on a model with an irregular +enough state space to make a difference. diff --git a/docs/source/guides/divorce_transition_without_lagged_state.md b/docs/source/guides/divorce_transition_without_lagged_state.md new file mode 100644 index 00000000..bbfadfd1 --- /dev/null +++ b/docs/source/guides/divorce_transition_without_lagged_state.md @@ -0,0 +1,243 @@ +# Implementing a divorce/marriage transition without a lagged partner state + +This note documents a working recipe, verified against an independent +hand-rolled EGM solver for both a two-period and a four-period model. The +full working code (model functions, reference solver, and the tests that +verify the equivalence claimed below) lives in +`tests/resources/divorce_model/` and `tests/test_divorce_toy_model.py`. + +## The problem + +You want: if the agent divorces between period `t` and `t+1`, they keep +half of what the household saved; if they marry, their new partner matches +their wealth and it doubles. That is a **transition-based** rule — it +depends on *both* `partner_state` at `t` and `partner_state` at `t+1`. + +There are two ways to get this into dcegm: + +1. **Add a `lagged_partner_state` to the state space** (a deterministic + state carrying last period's `partner_state` forward, alongside the + `lagged_choice` dcegm already tracks). With both `partner_state` + (current) and `lagged_partner_state` (previous) visible inside + `budget_constraint`, you can adjust the incoming asset directly and + literally, on the transition. This works, and is the more obvious route + -- at the cost of a genuine extra state dimension (bigger state space, + more batches, more memory). This note does not walk through that model. +2. **Keep the state space as-is and solve it via internal bookkeeping on + the individual level.** No extra state, no lagged variable. This is + what's implemented and verified in `tests/resources/divorce_model/`. + +Either way, the actual thing you need to get right is *how to correctly +rescale the asset given the transition*. Approach 2 answers that without +paying for a second state dimension, and that's what the rest of this note +covers. + +## Why the naive version of approach 2 breaks + +The first thing everyone tries: multiply the incoming asset by a +partner-status multiplier in `budget_constraint`, keyed off the *current* +`partner_state` (since that's all you have without the lagged state): + +```python +def budget_constraint(period, lagged_choice, partner_state, + asset_end_of_previous_period, income_shock_previous_period, + params, model_specs): + multiplier = 2.0 if partner_state == 1 else 1.0 + return multiplier * asset_end_of_previous_period * (1 + params["interest_rate"]) + income +``` + +This computes the **wealth level** correctly. It silently breaks the +**Euler equation**. dcegm's Euler-equation solver hardcodes the marginal +return on savings as `1 + params["interest_rate"]` +(`src/dcegm/egm/solve_euler_equation.py:166`, +`rhs_euler = marginal_utility_next * (1 + interest_rate) * discount_factor`). +It never differentiates `budget_constraint` — it just assumes +`d(wealth)/d(asset_end_of_previous_period) = 1 + interest_rate`, always. If +your `budget_constraint` actually implies `d(wealth)/d(a) = multiplier * +(1+r)`, dcegm computes the consumption-savings tradeoff as if a dollar +saved always earns `1+r`, when for a partnered agent it actually earns +`2*(1+r)`. The wealth level is right; the price of saving used in the +policy function is wrong. This is invisible unless you check against an +independent solve — the model still runs, converges, produces a +"reasonable-looking" policy, just the wrong one. + +## The recipe that works + +**Two changes, applied together, both keyed off the *current* period's own +`partner_state` only:** + +### 1. `budget_constraint`: double, then divide the whole thing by the same multiplier again + +```python +def budget_constraint(period, lagged_choice, partner_state, + asset_end_of_previous_period, income_shock_previous_period, + params, model_specs): + multiplier = jnp.where(partner_state == 1, 2.0, 1.0) + own_income = params["y_work"] * (lagged_choice == 0) + partner_income = params["y_partner"] * (partner_state == 1) + wealth = ( + multiplier * asset_end_of_previous_period * (1 + params["interest_rate"]) + + own_income + + partner_income + ) + return jnp.maximum(wealth, params["consumption_floor"]) / multiplier +``` + +For the asset term, `multiplier` cancels exactly: +`multiplier * a * (1+r) / multiplier = a * (1+r)`. So +`d(wealth)/d(a) = 1+r` regardless of `partner_state` — dcegm's hardcoded +assumption is now *actually true*, not just assumed. The income terms +still get divided by `multiplier`, i.e. partner income effectively gets +shared 50/50 through this division. + +### 2. Utility: scale the consumption argument by the same multiplier + +```python +def utility_func(consumption, choice, partner_state, params): + scale = jnp.where(partner_state == 1, 2.0, 1.0) + x = scale * consumption + felicity = ((x ** (1 - params["rho"]) - 1) / (1 - params["rho"])) + return felicity - (1 - choice) * params["delta"] +``` + +Consumption is tracked in *individual* terms (dcegm's own choice +variable), but if it's drawn from a jointly funded (pooled) account when +partnered, a dollar of individual consumption corresponds to two dollars of +joint spending — hence `scale * consumption` inside felicity. + +`marginal_utility_func` and `inverse_marginal_utility_func` must be +re-derived consistently from this, **not** just chain-ruled naively: + +```python +def marginal_utility_func(consumption, partner_state, params): + scale = jnp.where(partner_state == 1, 2.0, 1.0) + x = scale * consumption + return x ** (-params["rho"]) # NOT scale * x**(-rho) -- see below + +def inverse_marginal_utility_func(marginal_utility, partner_state, params): + scale = jnp.where(partner_state == 1, 2.0, 1.0) + return marginal_utility ** (-1 / params["rho"]) / scale +``` + +The naive chain rule of `d/dc [felicity(scale*c)]` is `scale * +felicity'(scale*c)`, i.e. *with* an extra outer `scale` factor. **Don't use +that version.** The one without the outer `scale` (just +`(scale*consumption)**(-rho)`) is the one that reproduces the correct +economics once combined with `budget_constraint` above — this is not an +approximation or a slip, it's required for the equivalence proved below. +If you're re-deriving this for a different utility function, treat "what +should `marginal_utility_func` return" as an equivalence to verify (as in +`tests/test_divorce_toy_model.py`), not as a pure calculus exercise on +`utility_func` in isolation. + +## Why this reproduces the transition-based rule exactly + +The natural worry: doesn't rescaling every period based on current status, +instead of only at the transition, change the economics? It doesn't. Write +out the transition-based resource function directly (this is the thing +approach 1, with the lagged state, would implement): + +```python +def resources_after_transition(a, partner_state_0, partner_state_1, income, r): + if partner_state_0 == 1 and partner_state_1 == 0: + a = a / 2 # divorce + elif partner_state_0 == 0 and partner_state_1 == 1: + a = a * 2 # marriage + return a * (1 + r) + income + partner_income(partner_state_1) +``` + +The following identity holds **exactly**, for every combination of +`partner_state_0`, `partner_state_1`: + +``` +resources_after_transition(scale(partner_state_0) * a, partner_state_0, partner_state_1, income) + == scale(partner_state_1) * a * (1 + r) + income + partner_income(partner_state_1) +``` + +The right-hand side depends only on `partner_state_1` (and the raw `a`) — +`partner_state_0` drops out entirely. In words: feeding the transition-based +formula an asset pre-scaled by *today's* multiplier makes it collapse to +exactly the per-period, current-state-only formula. These are not two +different economic models; they are the same model with two different unit +conventions for what the state variable "individual wealth" means: + +- **The per-period convention (recipe above):** `a` is always denominated + such that its real-dollar value is `scale(current partner_state) * a`. A + married person's `a=100` and a single person's `a=100` represent + different real dollar amounts. +- **The transition-based convention (what a `lagged_partner_state` model + would implement directly):** `a` is real dollars, always, and gets + explicitly converted (halved/doubled) exactly at the moment + `partner_state` changes. + +Converting between the two conventions is bookkeeping (multiply/divide by +`scale`), not a different model. `tests/test_divorce_toy_model.py`'s +`test_dcegm_policy_matches_hand_solved_reference` verifies this to +floating-point precision (~1e-13) for a two-period model by querying the +transition-based reference at `scale(partner_state_0) * a0_end` and +dividing the resulting policy back by `scale` — exactly the conversion +above. + +**Practical upshot:** the internal-bookkeeping recipe is not a compromise +relative to the lagged-state model; it's an exact reformulation of the same +economics, without the extra state dimension. + +## Extending to more than two periods + +Two additional, genuine (not conceptual) numerical issues show up once +intermediate periods are no longer analytic (see `reference.py`'s +`solve_reference` / `continuation_value_and_marg_util` for the fixes, and +`test_dcegm_policy_matches_hand_solved_reference_n_periods` for the +verification, ~4e-4 max relative error for a 4-period model): + +1. **Borrowing constraint, bottom of the grid.** A solved period's + endogenous grid only starts at its own `a_end=0` point's implied wealth + — it does not cover `[0, endog_grid[0])`. A continuation-wealth query + below that is not an edge case to `np.interp`-clamp away: the true + optimal policy there is "consume everything, save nothing" + (`policy = wealth`, exactly). dcegm handles this natively (its own + solved arrays carry a duplicate natural-borrowing-constraint point at + the bottom); if you hand-roll a reference solver for verification, you + have to add the same handling explicitly or your reference will be + silently wrong for a wide low-wealth range, not just approximately off. + +2. **Extrapolation, top of the grid.** A marriage transition can double + continuation wealth, which can push it *above* the range the + continuation period's own grid was solved over (built for un-doubled, + individual-scale wealth). There's no exact closed form here (unlike the + borrowing constraint); linear extrapolation from the top two grid points + is the standard, accurate-enough fix (`c(wealth)` is close to linear for + large wealth under CRRA). Alternatively, simply build the exogenous + asset grid wide enough that this rarely binds for your actual estimation + sample — but don't rely on that alone without checking, since the + multiplier can compound across consecutive partnered periods. + +Both issues are symptoms of the exact same root cause: interpolating a +solved policy/value function outside the domain it was actually solved +over. They're not specific to the divorce mechanism, but the `2x` (or +`0.5x`) rescaling here makes them bite much sooner than in a standard +model, so they need explicit handling rather than assuming the default grid +is "wide enough." + +## Checklist + +1. Add `partner_state`-conditional `scale = 2 if partnered else 1` (or the + real model's actual pooling assumption, if not a simple doubling) to + `budget_constraint`'s asset term, and divide the *whole* wealth + expression by the same `scale` again. +2. Re-derive `marginal_utility_func` / `inverse_marginal_utility_func` + consistently with the *no-extra-outer-scale* convention above, not via + naive chain rule. +3. Decide deliberately between this approach and adding + `lagged_partner_state` to the state space -- both are valid; this note's + recipe avoids the extra state dimension, at the cost of the + less-obvious derivation above. +4. If extending beyond two periods, check the asset grid's range against + the maximum multiplier compounding you expect given the estimated + marriage/divorce transition probabilities, and add explicit + borrowing-constraint / extrapolation handling to any independent + verification solver you write (dcegm handles both natively already). +5. Verify against an independent solve before trusting the result — this + class of bug (wealth level right, Euler equation silently wrong) + produces a model that runs and looks plausible without ever raising an + error. diff --git a/docs/source/index.rst b/docs/source/index.rst index e0e6874e..23253c5c 100644 --- a/docs/source/index.rst +++ b/docs/source/index.rst @@ -36,7 +36,6 @@ Check out our :ref:`guides` to find information on getting sta guides/two_occupation_model.ipynb - .. toctree:: :maxdepth: 2 :caption: Background @@ -57,6 +56,7 @@ Check out our :ref:`guides` to find information on getting sta development/team development/changes development/roadmap + development/internals/batching .. toctree:: diff --git a/src/dcegm/backward_induction.py b/src/dcegm/backward_induction.py index 8f9fc03c..426f9fd8 100644 --- a/src/dcegm/backward_induction.py +++ b/src/dcegm/backward_induction.py @@ -7,7 +7,6 @@ import jax.numpy as jnp from dcegm.final_periods import solve_last_two_periods -from dcegm.law_of_motion import calc_cont_grids_next_period from dcegm.pre_processing.sol_container import create_solution_container from dcegm.solve_single_period import solve_single_period @@ -40,20 +39,18 @@ def backward_induction( """ continuous_states_info = model_config["continuous_states_info"] - - # - calc_grids_jit = jax.jit( - lambda income_shock_draws, params_inner: calc_cont_grids_next_period( - model_structure=model_structure, - model_config=model_config, - income_shock_draws_unscaled=income_shock_draws, - params=params_inner, - model_funcs=model_funcs, - ) + skip_endog_grid_storage = model_config["upper_envelope"]["skip_endog_grid_storage"] + + # Scale income shock draws once. This is cheap (shape (n_quad,)) and shared by + # every batch/period; the actual (large) child continuous-state/wealth + # transitions are computed on demand per batch/period instead of upfront for the + # whole state space (see solve_single_period.py / final_periods.py). + income_shock_mean = model_funcs["read_funcs"]["income_shock_mean"](params) + income_shock_std = model_funcs["read_funcs"]["income_shock_std"](params) + income_shocks_scaled = ( + income_shock_draws_unscaled * income_shock_std + income_shock_mean ) - cont_grids_next_period = calc_grids_jit(income_shock_draws_unscaled, params) - # Infer n_continuous_state_combinations from model structure n_continuous_state_combinations = model_structure["continuous_state_space"][ next(iter(model_structure["continuous_state_space"])) @@ -68,18 +65,20 @@ def backward_induction( # Read out grid size n_total_wealth_grid=model_config["n_total_wealth_grid"], n_state_choices=model_structure["state_choice_space"].shape[0], + store_endog_grid=not skip_endog_grid_storage, ) # Solve the last two periods using lambda to capture static arguments solve_last_two_period_jit = jax.jit( - lambda params_inner, cont_grids, weights, val_solved, pol_solved, endog_solved: solve_last_two_periods( + lambda params_inner, shocks_scaled, weights, val_solved, pol_solved, endog_solved: solve_last_two_periods( params=params_inner, continuous_states_info=continuous_states_info, model_structure=model_structure, - cont_grids_next_period=cont_grids, + income_shocks_scaled=shocks_scaled, income_shock_weights=weights, model_funcs=model_funcs, upper_envelope_method=model_config["upper_envelope"]["method"], + skip_endog_grid_storage=skip_endog_grid_storage, last_two_period_batch_info=batch_info["last_two_period_info"], value_solved=val_solved, policy_solved=pol_solved, @@ -94,7 +93,7 @@ def backward_induction( endog_grid_solved, ) = solve_last_two_period_jit( params, - cont_grids_next_period, + income_shocks_scaled, income_shock_weights, value_solved, policy_solved, @@ -112,10 +111,11 @@ def backward_induction( params=params, continuous_grids_info=continuous_states_info, continuous_state_space=model_structure["continuous_state_space"], - cont_grids_next_period=cont_grids_next_period, + income_shocks_scaled=income_shocks_scaled, model_funcs=model_funcs, income_shock_weights=income_shock_weights, upper_envelope_method=model_config["upper_envelope"]["method"], + skip_endog_grid_storage=skip_endog_grid_storage, debug_info=None, ) diff --git a/src/dcegm/egm/interpolate_marginal_utility.py b/src/dcegm/egm/interpolate_marginal_utility.py index 078714dd..d3c90f94 100644 --- a/src/dcegm/egm/interpolate_marginal_utility.py +++ b/src/dcegm/egm/interpolate_marginal_utility.py @@ -1,4 +1,4 @@ -from typing import Any, Callable, Dict, Tuple, cast +from typing import Any, Callable, Dict, Tuple from jax import numpy as jnp from jax import vmap @@ -11,20 +11,21 @@ from dcegm.interpolation.interpnd_regular import ( interpnd_policy_and_value_for_child_states_on_regular_grids, ) +from dcegm.law_of_motion import calc_law_of_motion_for_state_choices def interpolate_value_and_marg_util( model_funcs, state_choice_vec: Dict[str, int], continuous_grids_info: Dict[str, Any], - cont_grids_next_period: Dict[str, jnp.ndarray], + income_shocks_scaled: jnp.ndarray, endog_grid_child_state_choice: jnp.ndarray, policy_child_state_choice: jnp.ndarray, value_child_state_choice: jnp.ndarray, - child_state_idxs: jnp.ndarray, continuous_state_space, params: Dict[str, float], upper_envelope_method: str, + skip_endog_grid_storage: bool = False, ) -> Tuple[jnp.ndarray, jnp.ndarray]: """Interpolate value and policy for all child states and compute marginal utility. @@ -59,25 +60,36 @@ def interpolate_value_and_marg_util( income shock. """ - wealth_child_states = cont_grids_next_period["assets_begin_of_period"][ - child_state_idxs - ] - compute_marginal_utility = model_funcs["compute_marginal_utility"] - compute_utility = model_funcs["compute_utility"] - discount_factor = model_funcs["read_funcs"]["discount_factor"](params) - # Check if interpolation needs to be multidimensional and irregular multi_dim = continuous_grids_info["has_additional_continuous_state"] irregular = upper_envelope_method == "fues" + # Compute the child continuous-state/wealth transitions on demand for exactly + # this batch's children, instead of reading from a precomputed whole-state-space + # structure (see law_of_motion.py). + law_of_motion = calc_law_of_motion_for_state_choices( + state_choice_vec=state_choice_vec, + continuous_state_space=continuous_state_space, + assets_grid_end_of_period=continuous_grids_info["assets_grid_end_of_period"], + income_shocks_scaled=income_shocks_scaled, + params=params, + model_funcs=model_funcs, + has_additional_continuous_states=multi_dim, + ) + wealth_child_states = law_of_motion["assets_begin_of_period"] + continuous_states_next = law_of_motion["continuous_states"] + + compute_marginal_utility = model_funcs["compute_marginal_utility"] + compute_utility = model_funcs["compute_utility"] + discount_factor = model_funcs["read_funcs"]["discount_factor"](params) + if multi_dim & irregular: return _interpolate_value_and_marg_util_2d_irregular( compute_marginal_utility=compute_marginal_utility, compute_utility=compute_utility, state_choice_vec=state_choice_vec, continuous_grids_info=continuous_grids_info, - cont_grids_next_period=cont_grids_next_period, - child_state_idxs=child_state_idxs, + continuous_states_next=continuous_states_next, wealth_child_states=wealth_child_states, endog_grid_child_state_choice=endog_grid_child_state_choice, policy_child_state_choice=policy_child_state_choice, @@ -92,20 +104,34 @@ def interpolate_value_and_marg_util( compute_utility=compute_utility, state_choice_vec=state_choice_vec, continuous_grids_info=continuous_grids_info, - cont_grids_next_period=cont_grids_next_period, - child_state_idxs=child_state_idxs, + continuous_states_next=continuous_states_next, wealth_child_states=wealth_child_states, endog_grid_child_state_choice=endog_grid_child_state_choice, policy_child_state_choice=policy_child_state_choice, value_child_state_choice=value_child_state_choice, params=params, discount_factor=discount_factor, + skip_endog_grid_storage=skip_endog_grid_storage, ) else: # Selects inside if jorgensen_druedahl or fues (different treatment of budget constraint) + # Under DJ, the wealth grid is a fixed constant (not batched over children), so + # it is passed with in_axes=None, broadcast only across the (small) continuous + # combo axis rather than the (potentially large) children axis. + if skip_endog_grid_storage: + dj_wealth_grid = continuous_grids_info["dj_wealth_grid"] + endog_grid_arg = jnp.broadcast_to( + dj_wealth_grid, + (policy_child_state_choice.shape[1], dj_wealth_grid.shape[0]), + ) + endog_grid_in_axes = None + else: + endog_grid_arg = endog_grid_child_state_choice + endog_grid_in_axes = 0 + interp_for_single_state_choice = vmap( interp1d_value_and_marg_util_for_state_choice, - in_axes=(None, None, 0, 0, 0, 0, 0, None, None, None), + in_axes=(None, None, 0, 0, endog_grid_in_axes, 0, 0, None, None, None), ) return interp_for_single_state_choice( @@ -113,7 +139,7 @@ def interpolate_value_and_marg_util( compute_utility, state_choice_vec, wealth_child_states, - endog_grid_child_state_choice, + endog_grid_arg, policy_child_state_choice, value_child_state_choice, params, @@ -222,8 +248,7 @@ def _interpolate_value_and_marg_util_2d_irregular( compute_utility: Callable, state_choice_vec: Dict[str, int], continuous_grids_info: Dict[str, Any], - cont_grids_next_period: Dict[str, jnp.ndarray], - child_state_idxs: jnp.ndarray, + continuous_states_next: Dict[str, jnp.ndarray], wealth_child_states: jnp.ndarray, endog_grid_child_state_choice: jnp.ndarray, policy_child_state_choice: jnp.ndarray, @@ -243,10 +268,7 @@ def _interpolate_value_and_marg_util_2d_irregular( continuous_state_name = continuous_grids_info["additional_continuous_state_names"][ 0 ] - continuous_states_next = _get_continuous_states_next(cont_grids_next_period) - continuous_state_child_states = continuous_states_next[continuous_state_name][ - child_state_idxs - ] + continuous_state_child_states = continuous_states_next[continuous_state_name] interp_for_single_state_choice = vmap( interp2d_value_and_marg_util_for_state_choice, @@ -285,14 +307,14 @@ def _interpolate_value_and_marg_util_nd_regular( compute_utility: Callable, state_choice_vec: Dict[str, int], continuous_grids_info: Dict[str, Any], - cont_grids_next_period: Dict[str, jnp.ndarray], - child_state_idxs: jnp.ndarray, + continuous_states_next: Dict[str, jnp.ndarray], wealth_child_states: jnp.ndarray, endog_grid_child_state_choice: jnp.ndarray, policy_child_state_choice: jnp.ndarray, value_child_state_choice: jnp.ndarray, params: Dict[str, float], discount_factor: float, + skip_endog_grid_storage: bool = False, ) -> Tuple[jnp.ndarray, jnp.ndarray]: """Interpolate value and marginal utility on the regular n-D grid. @@ -301,18 +323,25 @@ def _interpolate_value_and_marg_util_nd_regular( """ continuous_state_names = continuous_grids_info["additional_continuous_state_names"] - continuous_states_next = _get_continuous_states_next(cont_grids_next_period) continuous_state_child_states = { - name: continuous_states_next[name][child_state_idxs] - for name in continuous_state_names + name: continuous_states_next[name] for name in continuous_state_names } + # interpnd_policy_and_value_for_child_states_on_regular_grids already treats + # wealth_grid as a single unbatched 1D array (it is not itself a mapped vmap + # argument), so under DJ we can pass the fixed grid directly with no broadcast. + wealth_grid = ( + continuous_grids_info["dj_wealth_grid"] + if skip_endog_grid_storage + else endog_grid_child_state_choice[0, 0] + ) + policy_interp, value_interp = ( interpnd_policy_and_value_for_child_states_on_regular_grids( additional_continuous_state_grids=continuous_grids_info[ "additional_continuous_state_grids" ], - wealth_grid=endog_grid_child_state_choice[0, 0], + wealth_grid=wealth_grid, policy_grid_child_states=policy_child_state_choice, value_grid_child_states=value_child_state_choice, continuous_state_child_states=continuous_state_child_states, @@ -334,17 +363,6 @@ def _interpolate_value_and_marg_util_nd_regular( return value_interp, marg_util_interp -def _get_continuous_states_next( - cont_grids_next_period: Dict[str, jnp.ndarray], -) -> Dict[str, jnp.ndarray]: - if "continuous_states" not in cont_grids_next_period: - raise KeyError( - "Expected key 'continuous_states' in cont_grids_next_period. " - "This object should come from law_of_motion.calc_cont_grids_next_period()." - ) - return cast(Dict[str, jnp.ndarray], cont_grids_next_period["continuous_states"]) - - def _compute_nd_marginal_utility( compute_marginal_utility: Callable, policy_interp: jnp.ndarray, diff --git a/src/dcegm/final_periods.py b/src/dcegm/final_periods.py index 0d0ff7f4..dca3a2d1 100644 --- a/src/dcegm/final_periods.py +++ b/src/dcegm/final_periods.py @@ -8,6 +8,7 @@ from dcegm.check_func_outputs import ( check_budget_equation_and_return_wealth_plus_optional_aux, ) +from dcegm.law_of_motion import calc_law_of_motion_for_state_choices from dcegm.solve_single_period import solve_for_interpolated_values @@ -15,10 +16,11 @@ def solve_last_two_periods( params: Dict[str, float], continuous_states_info: Dict[str, Any], model_structure: Dict[str, Any], - cont_grids_next_period: Dict[str, Any], + income_shocks_scaled: jnp.ndarray, income_shock_weights: jnp.ndarray, model_funcs: Dict[str, Any], upper_envelope_method: str, + skip_endog_grid_storage: bool, last_two_period_batch_info, value_solved, policy_solved, @@ -59,11 +61,11 @@ def solve_last_two_periods( marginal_utility_final_last_period, ) = solve_final_period( idx_state_choices_final_period=batch_info["idx_state_choices_final_period"], - idx_parent_states_final_period=batch_info["idxs_parent_states_final_period"], state_choice_mat_final_period=batch_info["state_choice_mat_final_period"], - cont_grids_next_period=cont_grids_next_period, + income_shocks_scaled=income_shocks_scaled, continuous_states_info=continuous_states_info, upper_envelope_method=upper_envelope_method, + skip_endog_grid_storage=skip_endog_grid_storage, model_structure=model_structure, params=params, model_funcs=model_funcs, @@ -115,9 +117,10 @@ def solve_last_two_periods( policy_solved = policy_solved.at[idx_second_last, ...].set( out_dict_second_last["policy"] ) - endog_grid_solved = endog_grid_solved.at[idx_second_last, ...].set( - out_dict_second_last["endog_grid"] - ) + if not skip_endog_grid_storage: + endog_grid_solved = endog_grid_solved.at[idx_second_last, ...].set( + out_dict_second_last["endog_grid"] + ) # If we do not call the function in debug mode. Assign everything and return if debug_info is None: @@ -149,11 +152,11 @@ def solve_last_two_periods( def solve_final_period( idx_state_choices_final_period, - idx_parent_states_final_period, state_choice_mat_final_period, - cont_grids_next_period: Dict[str, Any], + income_shocks_scaled: jnp.ndarray, continuous_states_info: Dict[str, Any], upper_envelope_method: str, + skip_endog_grid_storage: bool, model_structure: Dict[str, Any], params: Dict[str, float], model_funcs: Dict[str, Any], @@ -187,15 +190,21 @@ def solve_final_period( compute_utility = model_funcs["compute_utility_final"] compute_marginal_utility = model_funcs["compute_marginal_utility_final"] - wealth_child_states_final_period = cont_grids_next_period["assets_begin_of_period"][ - idx_parent_states_final_period + law_of_motion_final_period = calc_law_of_motion_for_state_choices( + state_choice_vec=state_choice_mat_final_period, + continuous_state_space=model_structure["continuous_state_space"], + assets_grid_end_of_period=continuous_states_info["assets_grid_end_of_period"], + income_shocks_scaled=income_shocks_scaled, + params=params, + model_funcs=model_funcs, + has_additional_continuous_states=continuous_states_info[ + "has_additional_continuous_state" + ], + ) + wealth_child_states_final_period = law_of_motion_final_period[ + "assets_begin_of_period" ] - continuous_state_final = { - key: value[idx_parent_states_final_period] - for key, value in cont_grids_next_period["continuous_states"].items() - } - - n_assets = wealth_child_states_final_period.shape[-2] + continuous_state_final = law_of_motion_final_period["continuous_states"] value, marg_util = vmap( vmap( @@ -267,15 +276,22 @@ def solve_final_period( (zeros_to_append[..., None], wealth_sorted), axis=2 ) + # Width of the actual final-period wealth grid just built above -- for + # druedahl_jorgensen this is len(assets_begin_of_period) + 1 (matching + # n_total_wealth_grid, see check_model_config.py); for fues it's + # len(assets_grid_end_of_period) + 1. + n_wealth_final = values_with_zeros.shape[-1] + value_solved = value_solved.at[ - idx_state_choices_final_period, :, : n_assets + 1 + idx_state_choices_final_period, :, :n_wealth_final ].set(values_with_zeros) policy_solved = policy_solved.at[ - idx_state_choices_final_period, :, : n_assets + 1 - ].set(wealth_with_zeros) - endog_grid_solved = endog_grid_solved.at[ - idx_state_choices_final_period, :, : n_assets + 1 + idx_state_choices_final_period, :, :n_wealth_final ].set(wealth_with_zeros) + if not skip_endog_grid_storage: + endog_grid_solved = endog_grid_solved.at[ + idx_state_choices_final_period, :, :n_wealth_final + ].set(wealth_with_zeros) return ( value_solved, diff --git a/src/dcegm/interfaces/inspect_solution.py b/src/dcegm/interfaces/inspect_solution.py index 482c8442..463d12e4 100644 --- a/src/dcegm/interfaces/inspect_solution.py +++ b/src/dcegm/interfaces/inspect_solution.py @@ -5,7 +5,6 @@ import numpy as np from dcegm.final_periods import solve_last_two_periods -from dcegm.law_of_motion import calc_cont_grids_next_period from dcegm.pre_processing.sol_container import create_solution_container from dcegm.solve_single_period import solve_single_period @@ -37,13 +36,15 @@ def partially_solve( raise ValueError("You must at least solve for two periods.") continuous_states_info = model_config["continuous_states_info"] + skip_endog_grid_storage = model_config["upper_envelope"]["skip_endog_grid_storage"] - cont_grids_next_period = calc_cont_grids_next_period( - model_structure=model_structure, - model_config=model_config, - income_shock_draws_unscaled=income_shock_draws_unscaled, - params=params, - model_funcs=model_funcs, + # Scale income shock draws once (cheap, shared by every batch/period); the + # actual (large) child continuous-state/wealth transitions are computed on + # demand per batch/period instead of upfront for the whole state space. + income_shock_mean = model_funcs["read_funcs"]["income_shock_mean"](params) + income_shock_std = model_funcs["read_funcs"]["income_shock_std"](params) + income_shocks_scaled = ( + income_shock_draws_unscaled * income_shock_std + income_shock_mean ) # Determine the last period we need to solve for. last_relevant_period = model_config["n_periods"] - n_periods @@ -67,6 +68,7 @@ def partially_solve( n_total_wealth_grid=model_config["n_total_wealth_grid"], n_state_choices=relevant_state_choice_space.shape[0], n_continuous_state_combinations=n_continuous_state_combinations, + store_endog_grid=not skip_endog_grid_storage, ) if return_candidates: @@ -100,10 +102,11 @@ def partially_solve( params=params, continuous_states_info=continuous_states_info, model_structure=model_structure, - cont_grids_next_period=cont_grids_next_period, + income_shocks_scaled=income_shocks_scaled, income_shock_weights=income_shock_weights, model_funcs=model_funcs, upper_envelope_method=model_config["upper_envelope"]["method"], + skip_endog_grid_storage=skip_endog_grid_storage, last_two_period_batch_info=last_two_period_batch_info, value_solved=value_solved, policy_solved=policy_solved, @@ -217,10 +220,11 @@ def partially_solve( params=params, continuous_grids_info=continuous_states_info, continuous_state_space=model_structure["continuous_state_space"], - cont_grids_next_period=cont_grids_next_period, + income_shocks_scaled=income_shocks_scaled, model_funcs=model_funcs, income_shock_weights=income_shock_weights, upper_envelope_method=model_config["upper_envelope"]["method"], + skip_endog_grid_storage=skip_endog_grid_storage, debug_info=debug_info, ) diff --git a/src/dcegm/interfaces/interface.py b/src/dcegm/interfaces/interface.py index 23b96d31..97c3bf51 100644 --- a/src/dcegm/interfaces/interface.py +++ b/src/dcegm/interfaces/interface.py @@ -14,6 +14,7 @@ interpolate_policy_for_state_and_choice, interpolate_value_for_state_and_choice, ) +from dcegm.pre_processing.sol_container import broadcast_dj_wealth_grid def get_n_state_choice_period(model_structure): @@ -74,13 +75,6 @@ def policy_and_value_for_states_and_choices( state_choice_idx = get_state_choice_index_per_discrete_states_and_choices( states=state_choices, choices=choices, model_structure=model_structure ) - endog_grid_state_choice = jnp.take( - endog_grid_solved, - state_choice_idx, - axis=0, - mode="fill", - fill_value=jnp.nan, - ) value_grid_state_choice = jnp.take( value_solved, state_choice_idx, @@ -95,10 +89,29 @@ def policy_and_value_for_states_and_choices( mode="fill", fill_value=jnp.nan, ) + if model_config["upper_envelope"]["skip_endog_grid_storage"]: + # DJ-constant: never batch the wealth grid over the (possibly large) query + # dimension. Broadcast only across the (small) continuous combo axis and let + # vmap treat it as invariant via in_axes=None. + endog_grid_state_choice = broadcast_dj_wealth_grid( + model_config["continuous_states_info"], + (value_grid_state_choice.shape[1],) + + model_config["continuous_states_info"]["dj_wealth_grid"].shape, + ) + endog_grid_in_axes = None + else: + endog_grid_state_choice = jnp.take( + endog_grid_solved, + state_choice_idx, + axis=0, + mode="fill", + fill_value=jnp.nan, + ) + endog_grid_in_axes = 0 policy, value = jax.vmap( interpolate_policy_and_value_for_state_and_choice, - in_axes=(0, 0, 0, 0, None, None, None, None), + in_axes=(0, 0, endog_grid_in_axes, 0, None, None, None, None), )( value_grid_state_choice, policy_grid_state_choice, @@ -149,13 +162,6 @@ def value_for_state_and_choice( state_choice_idx = get_state_choice_index_per_discrete_states_and_choices( states=state_choices, choices=choices, model_structure=model_structure ) - endog_grid_state_choice = jnp.take( - endog_grid_solved, - state_choice_idx, - axis=0, - mode="fill", - fill_value=jnp.nan, - ) value_grid_state_choice = jnp.take( value_solved, state_choice_idx, @@ -163,10 +169,26 @@ def value_for_state_and_choice( mode="fill", fill_value=jnp.nan, ) + if model_config["upper_envelope"]["skip_endog_grid_storage"]: + endog_grid_state_choice = broadcast_dj_wealth_grid( + model_config["continuous_states_info"], + (value_grid_state_choice.shape[1],) + + model_config["continuous_states_info"]["dj_wealth_grid"].shape, + ) + endog_grid_in_axes = None + else: + endog_grid_state_choice = jnp.take( + endog_grid_solved, + state_choice_idx, + axis=0, + mode="fill", + fill_value=jnp.nan, + ) + endog_grid_in_axes = 0 value = jax.vmap( interpolate_value_for_state_and_choice, - in_axes=(0, 0, 0, None, None, None, None), + in_axes=(0, endog_grid_in_axes, 0, None, None, None, None), )( value_grid_state_choice, endog_grid_state_choice, @@ -213,13 +235,6 @@ def policy_for_state_choice_vec( state_choice_idx = get_state_choice_index_per_discrete_states_and_choices( states=state_choices, choices=choices, model_structure=model_structure ) - endog_grid_state_choice = jnp.take( - endog_grid_solved, - state_choice_idx, - axis=0, - mode="fill", - fill_value=jnp.nan, - ) policy_grid_state_choice = jnp.take( policy_solved, state_choice_idx, @@ -234,10 +249,26 @@ def policy_for_state_choice_vec( mode="fill", fill_value=jnp.nan, ) + if model_config["upper_envelope"]["skip_endog_grid_storage"]: + endog_grid_state_choice = broadcast_dj_wealth_grid( + model_config["continuous_states_info"], + (value_grid_state_choice.shape[1],) + + model_config["continuous_states_info"]["dj_wealth_grid"].shape, + ) + endog_grid_in_axes = None + else: + endog_grid_state_choice = jnp.take( + endog_grid_solved, + state_choice_idx, + axis=0, + mode="fill", + fill_value=jnp.nan, + ) + endog_grid_in_axes = 0 policy = jax.vmap( interpolate_policy_for_state_and_choice, - in_axes=(0, 0, 0, 0, None, None, None, None), + in_axes=(0, 0, endog_grid_in_axes, 0, None, None, None, None), )( policy_grid_state_choice, value_grid_state_choice, @@ -367,13 +398,22 @@ def choice_values_for_states( mode="fill", fill_value=jnp.nan, ) - endog_grid_states = jnp.take( - endog_grid_solved, - state_choice_indexes, - axis=0, - mode="fill", - fill_value=jnp.nan, - ) + if model_config["upper_envelope"]["skip_endog_grid_storage"]: + # Broadcast only across the combo axis, never across the (states, choices) + # batch axes that the double vmap below maps over. + endog_grid_states = broadcast_dj_wealth_grid( + model_config["continuous_states_info"], value_grid_states.shape[2:] + ) + endog_grid_in_axes = None + else: + endog_grid_states = jnp.take( + endog_grid_solved, + state_choice_indexes, + axis=0, + mode="fill", + fill_value=jnp.nan, + ) + endog_grid_in_axes = 0 def wrapper_interp_value_for_choice( state, @@ -399,9 +439,9 @@ def wrapper_interp_value_for_choice( choice_values_per_state = jax.vmap( jax.vmap( wrapper_interp_value_for_choice, - in_axes=(None, 0, 0, 0), + in_axes=(None, 0, endog_grid_in_axes, 0), ), - in_axes=(0, 0, 0, None), + in_axes=(0, 0, endog_grid_in_axes, None), )( states, value_grid_states, @@ -429,13 +469,6 @@ def choice_policies_for_states( mode="fill", fill_value=jnp.nan, ) - endog_grid_states = jnp.take( - endog_grid_solved, - state_choice_indexes, - axis=0, - mode="fill", - fill_value=jnp.nan, - ) value_grid_states = jnp.take( value_solved, state_choice_indexes, @@ -443,6 +476,20 @@ def choice_policies_for_states( mode="fill", fill_value=jnp.nan, ) + if model_config["upper_envelope"]["skip_endog_grid_storage"]: + endog_grid_states = broadcast_dj_wealth_grid( + model_config["continuous_states_info"], value_grid_states.shape[2:] + ) + endog_grid_in_axes = None + else: + endog_grid_states = jnp.take( + endog_grid_solved, + state_choice_indexes, + axis=0, + mode="fill", + fill_value=jnp.nan, + ) + endog_grid_in_axes = 0 def wrapper_interp_value_for_choice( state, @@ -470,9 +517,9 @@ def wrapper_interp_value_for_choice( choice_values_per_state = jax.vmap( jax.vmap( wrapper_interp_value_for_choice, - in_axes=(None, 0, 0, 0, 0), + in_axes=(None, 0, 0, endog_grid_in_axes, 0), ), - in_axes=(0, 0, 0, 0, None), + in_axes=(0, 0, 0, endog_grid_in_axes, None), )( states, policy_grid_states, diff --git a/src/dcegm/interfaces/sol_interface.py b/src/dcegm/interfaces/sol_interface.py index 4f496b08..66d9415e 100644 --- a/src/dcegm/interfaces/sol_interface.py +++ b/src/dcegm/interfaces/sol_interface.py @@ -20,6 +20,7 @@ generate_alternative_sim_functions, ) from dcegm.pre_processing.shared import try_jax_array +from dcegm.pre_processing.sol_container import broadcast_dj_wealth_grid from dcegm.simulation.sim_utils import create_simulation_df from dcegm.simulation.simulate import simulate_all_periods @@ -180,13 +181,6 @@ def get_solution_for_discrete_state_choice(self, states, choices): choices=state_choices["choice"], ) - endog_grid = jnp.take( - self.endog_grid, - state_choice_index, - axis=0, - mode="fill", - fill_value=jnp.nan, - ) value_grid = jnp.take( self.value, state_choice_index, @@ -201,6 +195,18 @@ def get_solution_for_discrete_state_choice(self, states, choices): mode="fill", fill_value=jnp.nan, ) + if self.model_config["upper_envelope"]["skip_endog_grid_storage"]: + endog_grid = broadcast_dj_wealth_grid( + self.model_config["continuous_states_info"], value_grid.shape + ) + else: + endog_grid = jnp.take( + self.endog_grid, + state_choice_index, + axis=0, + mode="fill", + fill_value=jnp.nan, + ) return endog_grid, value_grid, policy_grid diff --git a/src/dcegm/interpolation/simulation_interp.py b/src/dcegm/interpolation/simulation_interp.py index 6ff72bae..3dc3d994 100644 --- a/src/dcegm/interpolation/simulation_interp.py +++ b/src/dcegm/interpolation/simulation_interp.py @@ -29,6 +29,8 @@ def interpolate_policy_and_value_for_all_agents( upper_envelope_method, has_additional_continuous_state, discount_factor, + skip_endog_grid_storage=False, + dj_wealth_grid=None, ): # 1D interpolation path is independent of upper-envelope method and only @@ -54,13 +56,21 @@ def interpolate_policy_and_value_for_all_agents( mode="fill", fill_value=jnp.nan, )[:, :, 0, :] - endog_grid_agent = jnp.take( - endog_grid_solved, - discrete_state_choice_indexes, - axis=0, - mode="fill", - fill_value=jnp.nan, - )[:, :, 0, :] + if skip_endog_grid_storage: + # DJ-constant: never batch the wealth grid over agents/choices. Already + # matches the shape a single agent-choice item needs, so no broadcast at + # all is needed here. + endog_grid_agent = dj_wealth_grid + endog_grid_in_axes = None + else: + endog_grid_agent = jnp.take( + endog_grid_solved, + discrete_state_choice_indexes, + axis=0, + mode="fill", + fill_value=jnp.nan, + )[:, :, 0, :] + endog_grid_in_axes = 0 vectorized_interp = vmap( vmap( @@ -68,7 +78,7 @@ def interpolate_policy_and_value_for_all_agents( in_axes=( None, None, - 0, + endog_grid_in_axes, 0, 0, 0, @@ -78,7 +88,7 @@ def interpolate_policy_and_value_for_all_agents( None, ), ), - in_axes=(0, 0, 0, 0, 0, None, None, None, None, None), + in_axes=(0, 0, endog_grid_in_axes, 0, 0, None, None, None, None, None), ) policy_agent, value_agent = vectorized_interp( @@ -202,13 +212,22 @@ def interpolate_policy_and_value_for_all_agents( mode="fill", fill_value=jnp.nan, ) - endog_grid_agent = jnp.take( - endog_grid_solved, - discrete_state_choice_indexes, - axis=0, - mode="fill", - fill_value=jnp.nan, - ) + if skip_endog_grid_storage: + # DJ-constant: broadcast only across the combo axis, never across + # agents/choices. + endog_grid_agent = jnp.broadcast_to( + dj_wealth_grid, value_grid_agent.shape[2:] + ) + endog_grid_in_axes = None + else: + endog_grid_agent = jnp.take( + endog_grid_solved, + discrete_state_choice_indexes, + axis=0, + mode="fill", + fill_value=jnp.nan, + ) + endog_grid_in_axes = 0 additional_continuous_state_names = list(continuous_state_space.keys()) @@ -219,7 +238,7 @@ def interpolate_policy_and_value_for_all_agents( None, None, None, - 0, + endog_grid_in_axes, 0, 0, 0, @@ -235,7 +254,7 @@ def interpolate_policy_and_value_for_all_agents( 0, 0, 0, - 0, + endog_grid_in_axes, 0, 0, None, diff --git a/src/dcegm/law_of_motion.py b/src/dcegm/law_of_motion.py index 0f264cd6..50a76a3c 100644 --- a/src/dcegm/law_of_motion.py +++ b/src/dcegm/law_of_motion.py @@ -6,32 +6,30 @@ ) -def calc_cont_grids_next_period( +def calc_law_of_motion_for_state_choices( + state_choice_vec, + continuous_state_space, + assets_grid_end_of_period, + income_shocks_scaled, params, - income_shock_draws_unscaled, - model_structure, - model_config, model_funcs, + has_additional_continuous_states, ): + """Compute continuous-state and wealth transitions for a set of state-choices. - continuous_states_info = model_config["continuous_states_info"] - state_space_dict = model_structure["state_space_dict"] + ``state_choice_vec`` may or may not contain a ``"choice"`` key. It is dropped (via a + no-op-if-absent pop) before being passed to the user-supplied law-of-motion + functions, since the transition does not depend on it -- this is what lets + ``calc_cont_grids_next_period`` below reuse this function unchanged with the full + (choice-less) state space. - has_additional_continuous_states = continuous_states_info[ - "has_additional_continuous_state" - ] - - # Scale income shock draws - income_shock_mean = model_funcs["read_funcs"]["income_shock_mean"](params) - income_shock_std = model_funcs["read_funcs"]["income_shock_std"](params) - income_shocks_scaled = ( - income_shock_draws_unscaled * income_shock_std + income_shock_mean - ) + """ + state_vec = dict(state_choice_vec) + state_vec.pop("choice", None) - continuous_state_space = model_structure["continuous_state_space"] continuous_state_next_period = _get_continuous_state_next_period( has_additional_continuous_states=has_additional_continuous_states, - state_space_dict=state_space_dict, + state_space_dict=state_vec, continuous_state_space=continuous_state_space, params=params, model_funcs=model_funcs, @@ -69,9 +67,9 @@ def fix_assets_and_shocks_for_broadcast( ), in_axes=(0, 0, None, None), )( - state_space_dict, + state_vec, continuous_state_next_period, - continuous_states_info["assets_grid_end_of_period"], + assets_grid_end_of_period, income_shocks_scaled, ) @@ -82,6 +80,44 @@ def fix_assets_and_shocks_for_broadcast( } +def calc_cont_grids_next_period( + params, + income_shock_draws_unscaled, + model_structure, + model_config, + model_funcs, +): + """Compute continuous-state and wealth transitions for the entire state space. + + Thin wrapper around ``calc_law_of_motion_for_state_choices``. Kept only for the + debug/inspection entry points in ``interfaces/model_class.py`` that need the full + structure; the main solve path computes this on demand, per batch/period, instead + (see ``solve_single_period.py``/``final_periods.py``). + + """ + continuous_states_info = model_config["continuous_states_info"] + state_space_dict = model_structure["state_space_dict"] + + # Scale income shock draws + income_shock_mean = model_funcs["read_funcs"]["income_shock_mean"](params) + income_shock_std = model_funcs["read_funcs"]["income_shock_std"](params) + income_shocks_scaled = ( + income_shock_draws_unscaled * income_shock_std + income_shock_mean + ) + + return calc_law_of_motion_for_state_choices( + state_choice_vec=state_space_dict, + continuous_state_space=model_structure["continuous_state_space"], + assets_grid_end_of_period=continuous_states_info["assets_grid_end_of_period"], + income_shocks_scaled=income_shocks_scaled, + params=params, + model_funcs=model_funcs, + has_additional_continuous_states=continuous_states_info[ + "has_additional_continuous_state" + ], + ) + + def _get_continuous_state_next_period( has_additional_continuous_states, state_space_dict, diff --git a/src/dcegm/pre_processing/check_model_config.py b/src/dcegm/pre_processing/check_model_config.py index f8816d5a..798a819e 100644 --- a/src/dcegm/pre_processing/check_model_config.py +++ b/src/dcegm/pre_processing/check_model_config.py @@ -186,6 +186,20 @@ def check_model_config_and_process(model_config): processed_model_config["continuous_states_info"]["assets_begin_of_period"] = ( jnp.asarray(model_config["continuous_states"]["assets_begin_of_period"]) ) + # The Druedahl-Jorgensen upper envelope always evaluates on this fixed grid + # (see upper_envelope.jax.drued_jorg_jax), so the resulting "endogenous" grid + # is not actually endogenous. We precompute it once here so callers can reuse + # it instead of reading a stored (and redundant) endog_grid array. + processed_model_config["continuous_states_info"]["dj_wealth_grid"] = ( + jnp.concatenate( + ( + jnp.zeros(1), + processed_model_config["continuous_states_info"][ + "assets_begin_of_period" + ], + ) + ) + ) if upper_envelope["method"] == "fues": processed_model_config["n_total_wealth_grid"] = tuning_params[ @@ -199,6 +213,19 @@ def check_model_config_and_process(model_config): else: raise ValueError("Something wrong internally") + # With a single discrete choice, the upper envelope is skipped entirely (see + # create_upper_envelope_function), so the stored endog_grid is not the fixed + # Druedahl-Jorgensen grid in that case and must still be stored/read normally. + upper_envelope["skip_endog_grid_storage"] = ( + upper_envelope["method"] == "druedahl_jorgensen" + and len(processed_model_config["choices"]) >= 2 + ) + if upper_envelope["skip_endog_grid_storage"]: + assert ( + processed_model_config["continuous_states_info"]["dj_wealth_grid"].shape[0] + == processed_model_config["n_total_wealth_grid"] + ) + if "min_period_batch_segments" in model_config.keys(): processed_model_config["min_period_batch_segments"] = model_config[ "min_period_batch_segments" diff --git a/src/dcegm/pre_processing/sol_container.py b/src/dcegm/pre_processing/sol_container.py index b8bf2f3e..d8311024 100644 --- a/src/dcegm/pre_processing/sol_container.py +++ b/src/dcegm/pre_processing/sol_container.py @@ -7,8 +7,16 @@ def create_solution_container( n_total_wealth_grid: int, n_state_choices: int, n_continuous_state_combinations: int, + store_endog_grid: bool = True, ): - """Create solution containers for value, policy, and endog_grid.""" + """Create solution containers for value, policy, and endog_grid. + + endog_grid is only allocated when store_endog_grid is True. When the + Druedahl-Jorgensen upper envelope is used, the "endogenous" grid is actually the + fixed exogenous grid, so storing it is skipped and callers read the exogenous + grid directly instead (see model_config["continuous_states_info"]["dj_wealth_grid"]). + + """ value_solved = jnp.full( (n_state_choices, n_continuous_state_combinations, n_total_wealth_grid), dtype=jnp.float64, @@ -19,10 +27,25 @@ def create_solution_container( dtype=jnp.float64, fill_value=jnp.nan, ) - endog_grid_solved = jnp.full( - (n_state_choices, n_continuous_state_combinations, n_total_wealth_grid), - dtype=jnp.float64, - fill_value=jnp.nan, + endog_grid_solved = ( + jnp.full( + (n_state_choices, n_continuous_state_combinations, n_total_wealth_grid), + dtype=jnp.float64, + fill_value=jnp.nan, + ) + if store_endog_grid + else None ) return value_solved, policy_solved, endog_grid_solved + + +def broadcast_dj_wealth_grid(continuous_states_info: Dict[str, Any], shape): + """Broadcast the fixed Druedahl-Jorgensen wealth grid to the given shape. + + Used in place of reading a stored endog_grid when + model_config["upper_envelope"]["skip_endog_grid_storage"] is True. Expects + continuous_states_info = model_config["continuous_states_info"]. + + """ + return jnp.broadcast_to(continuous_states_info["dj_wealth_grid"], shape) diff --git a/src/dcegm/simulation/simulate.py b/src/dcegm/simulation/simulate.py index c3b199bf..e5bfefe2 100644 --- a/src/dcegm/simulation/simulate.py +++ b/src/dcegm/simulation/simulate.py @@ -194,6 +194,10 @@ def simulate_single_period( upper_envelope_method=model_config["upper_envelope"]["method"], has_additional_continuous_state=has_additional_continuous_state, discount_factor=discount_factor, + skip_endog_grid_storage=model_config["upper_envelope"][ + "skip_endog_grid_storage" + ], + dj_wealth_grid=continuous_states_info.get("dj_wealth_grid"), ) # Draw taste shocks and calculate final value. diff --git a/src/dcegm/solve_single_period.py b/src/dcegm/solve_single_period.py index 8eab9bda..e12a0d81 100644 --- a/src/dcegm/solve_single_period.py +++ b/src/dcegm/solve_single_period.py @@ -14,10 +14,11 @@ def solve_single_period( params, continuous_grids_info, continuous_state_space, - cont_grids_next_period, + income_shocks_scaled, model_funcs, income_shock_weights, upper_envelope_method, + skip_endog_grid_storage, debug_info, ): """Solve a single period of the model using DCEGM.""" @@ -33,21 +34,26 @@ def solve_single_period( state_choice_mat_child, ) = xs + policy_child_state_choice = policy_solved[child_state_choice_idxs_to_interp] + endog_grid_child_state_choice = ( + None + if skip_endog_grid_storage + else endog_grid_solved[child_state_choice_idxs_to_interp] + ) + # EGM step 1) value_interpolated, marginal_utility_interpolated = interpolate_value_and_marg_util( model_funcs=model_funcs, state_choice_vec=state_choice_mat_child, continuous_grids_info=continuous_grids_info, - cont_grids_next_period=cont_grids_next_period, - endog_grid_child_state_choice=endog_grid_solved[ - child_state_choice_idxs_to_interp - ], + income_shocks_scaled=income_shocks_scaled, + endog_grid_child_state_choice=endog_grid_child_state_choice, continuous_state_space=continuous_state_space, - policy_child_state_choice=policy_solved[child_state_choice_idxs_to_interp], + policy_child_state_choice=policy_child_state_choice, value_child_state_choice=value_solved[child_state_choice_idxs_to_interp], - child_state_idxs=child_state_idxs, params=params, upper_envelope_method=upper_envelope_method, + skip_endog_grid_storage=skip_endog_grid_storage, ) # Check if we have a scalar taste shock scale or state specific. Extract in each of the cases. @@ -85,9 +91,10 @@ def solve_single_period( policy_solved = policy_solved.at[state_choices_idxs, :].set( out_dict_period["policy"] ) - endog_grid_solved = endog_grid_solved.at[state_choices_idxs, :].set( - out_dict_period["endog_grid"] - ) + if not skip_endog_grid_storage: + endog_grid_solved = endog_grid_solved.at[state_choices_idxs, :].set( + out_dict_period["endog_grid"] + ) # If we are not in the debug mode, we only return the solution as a tuple and an empty tuple. if debug_info is None: diff --git a/tests/resources/divorce_model/__init__.py b/tests/resources/divorce_model/__init__.py new file mode 100644 index 00000000..ac63b780 --- /dev/null +++ b/tests/resources/divorce_model/__init__.py @@ -0,0 +1,3 @@ +import jax + +jax.config.update("jax_enable_x64", True) diff --git a/tests/resources/divorce_model/dcegm_functions.py b/tests/resources/divorce_model/dcegm_functions.py new file mode 100644 index 00000000..12518fd6 --- /dev/null +++ b/tests/resources/divorce_model/dcegm_functions.py @@ -0,0 +1,119 @@ +"""dcegm-facing model functions for the divorce toy model. + +`partner_state` is a genuine stochastic state (0 = single, 1 = married), +with per-period wealth and utility scaling keyed off the *current* period's +own `partner_state` only. See the dcegm guide "Implementing a +divorce/marriage transition without a lagged partner state" +(`docs/source/guides/`) for why the per-period convention used below is +nonetheless exactly equivalent to a transition-based ("halve on divorce, +double on marriage") rule. + +As in `reference.py`, nothing here is hardcoded: `params` and the asset +grid are passed in by the caller (the test file owns the actual numbers). +""" + +import jax +import jax.numpy as jnp +import numpy as np + +import dcegm + + +def utility_func(consumption, choice, partner_state, params): + # The 2 only applies if there is a partner: consumption is drawn from a + # jointly funded (single, pooled) account, so a given dollar of recorded + # spending only cost this individual half of it. + scale = jnp.where(partner_state == 1, 2.0, 1.0) + x = scale * consumption + felicity = jax.lax.select( + jnp.allclose(params["rho"], 1), + jnp.log(x), + (x ** (1 - params["rho"]) - 1) / (1 - params["rho"]), + ) + return felicity - (1 - choice) * params["delta"] + + +def marginal_utility_func(consumption, partner_state, params): + scale = jnp.where(partner_state == 1, 2.0, 1.0) + x = scale * consumption + du_dx = jax.lax.select(jnp.allclose(params["rho"], 1), 1 / x, x ** (-params["rho"])) + return du_dx + + +def inverse_marginal_utility_func(marginal_utility, partner_state, params): + # Inverts marginal_utility_func(c) = (scale*c)**(-rho): c = m**(-1/rho) / scale. + scale = jnp.where(partner_state == 1, 2.0, 1.0) + c_rho1 = (1 / scale) * (1 / marginal_utility) # log case: scale cancels, as always + c_general = marginal_utility ** (-1 / params["rho"]) / scale + return jax.lax.select(jnp.allclose(params["rho"], 1), c_rho1, c_general) + + +def utility_final(wealth, choice, partner_state, params): + return utility_func(wealth, choice, partner_state, params) + + +def marginal_utility_final(wealth, choice, partner_state, params): + return marginal_utility_func(wealth, partner_state, params) + + +def budget_constraint( + period, + lagged_choice, + partner_state, + asset_end_of_previous_period, + income_shock_previous_period, + params, + model_specs, +): + multiplier = jnp.where(partner_state == 1, 2.0, 1.0) + own_income = params["y_work"] * (lagged_choice == 0) + # Partner income depends on *this* period's own partner_state directly + # (no lagged_choice involved -- we don't model the partner's own labor + # supply), added unscaled just like own_income. + partner_income = params["y_partner"] * (partner_state == 1) + wealth = ( + multiplier * asset_end_of_previous_period * (1 + params["interest_rate"]) + + own_income + + partner_income + ) + return jnp.maximum(wealth, params["consumption_floor"]) / multiplier + + +def feasible_choice_set(lagged_choice, model_specs): + return np.arange(model_specs["n_choices"]) + + +def partner_transition(partner_state, params): + """Returns [P(single next), P(married next)].""" + prob_married_next = jnp.where( + partner_state == 1, params["persistence_married"], params["prob_marry"] + ) + return jnp.array([1 - prob_married_next, prob_married_next]) + + +def build_and_solve(params, n_periods, a_grid): + model_specs = {"n_periods": n_periods, "n_choices": 2} + model_config = { + "n_periods": n_periods, + "choices": np.arange(2), + "stochastic_states": {"partner_state": np.arange(2)}, + "continuous_states": {"assets_end_of_period": a_grid}, + "n_quad_points": 5, + } + model = dcegm.setup_model( + model_config=model_config, + model_specs=model_specs, + state_space_functions={"state_specific_choice_set": feasible_choice_set}, + stochastic_states_transitions={"partner_state": partner_transition}, + utility_functions={ + "utility": utility_func, + "marginal_utility": marginal_utility_func, + "inverse_marginal_utility": inverse_marginal_utility_func, + }, + utility_functions_final_period={ + "utility": utility_final, + "marginal_utility": marginal_utility_final, + }, + budget_constraint=budget_constraint, + ) + return model, model.solve(params) diff --git a/tests/resources/divorce_model/reference.py b/tests/resources/divorce_model/reference.py new file mode 100644 index 00000000..201c5f75 --- /dev/null +++ b/tests/resources/divorce_model/reference.py @@ -0,0 +1,296 @@ +"""Independent hand-rolled reference solver for the divorce toy model. + +Plain, explicit backward-induction EGM -- no upper envelope, no dcegm +machinery -- so it can serve as ground truth to check the dcegm-based model +in `dcegm_functions.py` against. State is always *individual* wealth; +resources are only rescaled at an actual period-to-period partner +transition (halved on divorce, doubled on marriage, unchanged otherwise), +and utility has no partner-status multiplier at all. See the dcegm guide +"Implementing a divorce/marriage transition without a lagged partner state" +(`docs/source/guides/`) for why the dcegm-side model uses a *different*, +per-period convention, and why the two are nonetheless the same underlying +economics. + +Every function here takes the model's numeric parametrization (`params`, a +dict) explicitly -- nothing is hardcoded at module level. The actual numbers +used for testing live in `tests/test_divorce_toy_model.py`, not here. + +Uses dcegm's timing convention throughout: income earned under a given +`choice` at period t is paid at the start of period t+1 (i.e. it enters the +budget via `lagged_choice`). + +""" + +import numpy as np + + +def partner_transition_np(partner_state, params): + """Returns [P(single next), P(married next)].""" + prob_married_next = ( + params["persistence_married"] if partner_state == 1 else params["prob_marry"] + ) + return np.array([1 - prob_married_next, prob_married_next]) + + +def resources_after_transition(a, partner_state_0, partner_state_1, own_income, params): + """Plain resources, with the individual's asset stock rescaled once for + the period-to-period partner transition: halved on divorce (lose the + ex-partner's share), doubled on marriage (new partner matches wealth), + unchanged otherwise. No scaling anywhere else -- not on staying + partnered, not on income, not in utility. Partner income is added + whenever partnered next period (partner_state_1), unscaled, exactly + like own_income. + """ + if partner_state_0 == 1 and partner_state_1 == 0: + a = a / 2 # divorce + elif partner_state_0 == 0 and partner_state_1 == 1: + a = a * 2 # marriage + + partner_income = params["y_partner"] if partner_state_1 == 1 else 0.0 + wealth = a * (1 + params["interest_rate"]) + own_income + partner_income + return max(wealth, params["consumption_floor"]) + + +def consumption_utility(consumption, choice, params): + mu, delta = params["rho"], params["delta"] + u = ( + np.log(consumption) + if abs(mu - 1) < 1e-12 + else (consumption ** (1 - mu) - 1) / (1 - mu) + ) + return u - (1 - choice) * delta + + +def marginal_utility_np(consumption, params): + mu = params["rho"] + return 1 / consumption if abs(mu - 1) < 1e-12 else consumption ** (-mu) + + +def inverse_marginal_utility_np(marg_util, params): + mu = params["rho"] + return 1 / marg_util if abs(mu - 1) < 1e-12 else marg_util ** (-1 / mu) + + +def solve_reference_at_point(a0_end, partner_state_0, work0, params): + """Manual EGM backward induction for a single exogenous end-of-period asset point, + two periods only: period 1 (analytic terminal period) then the period-0 Euler + equation. + + No grid needed for two periods -- the n-period version below (`solve_reference`) + generalizes this by interpolating a stored grid instead of evaluating period 1 + analytically. + + """ + beta, r = params["discount_factor"], params["interest_rate"] + taste_shock_scale = params["taste_shock_scale"] + transition_probs = partner_transition_np(partner_state_0, params) + + # work0's income arrives at the start of period 1, i.e. lagged_choice = + # work0 there. + income_period1 = params["y_work"] * (work0 == 0) + + expected_marg_util = 0.0 + expected_value = 0.0 + for partner_state_1, prob_partner in enumerate(transition_probs): + if prob_partner == 0.0: + continue + wealth_1 = resources_after_transition( + a0_end, partner_state_0, partner_state_1, income_period1, params + ) + + # Terminal period: consume everything regardless of choice1, only + # the disutility of choice1 differs. + choice_values = np.array( + [consumption_utility(wealth_1, c1, params) for c1 in (0, 1)] + ) + choice_marg_utils = np.array( + [marginal_utility_np(wealth_1, params) for _ in (0, 1)] + ) + + max_v = choice_values.max() + weights = np.exp((choice_values - max_v) / taste_shock_scale) + choice_probs = weights / weights.sum() + ev_choice = max_v + taste_shock_scale * np.log(weights.sum()) + + expected_marg_util += prob_partner * np.sum(choice_probs * choice_marg_utils) + expected_value += prob_partner * ev_choice + + rhs_euler = beta * (1 + r) * expected_marg_util + c0 = inverse_marginal_utility_np(rhs_euler, params) + + return { + "endog_grid": c0 + a0_end, + "policy": c0, + "value": consumption_utility(c0, work0, params) + beta * expected_value, + } + + +def solve_reference_backward_induction(params, a_grid): + """Two-period manual EGM backward induction over the whole exogenous grid. + + Returns a dict keyed by (partner_state_0, work0), each holding parallel arrays + `endog_grid`, `policy`, `value` -- one entry per point of `a_grid`, built by calling + `solve_reference_at_point` at each point. + + """ + solved = {} + + for partner_state_0 in (0, 1): + for work0 in (0, 1): + points = [ + solve_reference_at_point(a0_end, partner_state_0, work0, params) + for a0_end in a_grid + ] + solved[(partner_state_0, work0)] = { + key: np.array([p[key] for p in points]) + for key in ("endog_grid", "policy", "value") + } + + return solved + + +def solve_reference(n_periods, params, a_grid): + """General n-period manual EGM backward induction, no upper envelope. + + Only the terminal period is analytic (consume everything). Every earlier + period is solved on `a_grid` and stored; the period before it reads + (interpolates) that stored grid for its own continuation value and + marginal utility, exactly as dcegm does internally. + + Returns solved[period][(partner_state, choice)] = dict with parallel + arrays "endog_grid", "policy", "value", for period in + range(n_periods - 1). The terminal period (n_periods - 1) is not stored + -- it's cheaper and exact to just evaluate `consumption_utility` / + `marginal_utility_np` directly on any wealth level. + """ + beta, r = params["discount_factor"], params["interest_rate"] + taste_shock_scale = params["taste_shock_scale"] + last_period = n_periods - 1 + solved = {} + + def continuation_value_and_marg_util(period, wealth, partner_state, choice): + """Value and marginal utility of consumption in `period`, at a given + wealth level, either analytically (terminal period) or by + interpolating the already-solved grid for `period`. + + The endogenous grid for a solved period only starts at its own + a_end=0 point's implied wealth (`endog_grid[0] == policy[0]`, since + saving nothing means consuming everything) -- it does not cover + [0, endog_grid[0]). Querying a lower wealth there is *not* an + extrapolation edge case to clamp away: the true optimal policy is + the borrowing constraint binding, i.e. consume everything + (`policy = wealth`) and keep the same continuation choice of + a_end=0. Handling this explicitly (rather than letting `np.interp` + silently clamp to `policy[0]`) is exactly what dcegm's own natural + borrowing-constraint point does internally. + + Symmetrically, a wealth level *above* the grid's own top point can + arise here too (e.g. a marriage transition doubling wealth): there + is no exact closed form there (unlike the borrowing constraint), so + this uses standard linear extrapolation from the top two grid + points -- accurate since `policy(wealth)` is close to linear for + large wealth in a CRRA problem, and the alternative (letting + `np.interp` clamp to a constant) is not. + """ + if period == last_period: + return consumption_utility(wealth, choice, params), marginal_utility_np( + wealth, params + ) + arrs = solved[period][(partner_state, choice)] + endog_grid, policy_grid, value_grid = ( + arrs["endog_grid"], + arrs["policy"], + arrs["value"], + ) + if wealth <= endog_grid[0]: + policy = wealth + # Continuation value at a_end=0, backed out from the stored + # value at the grid's own lower endpoint (where policy == + # endog_grid[0] already, i.e. a_end=0 there too). + ev_at_zero_savings = ( + value_grid[0] - consumption_utility(policy_grid[0], choice, params) + ) / beta + value = ( + consumption_utility(policy, choice, params) + beta * ev_at_zero_savings + ) + elif wealth > endog_grid[-1]: + slope_c = (policy_grid[-1] - policy_grid[-2]) / ( + endog_grid[-1] - endog_grid[-2] + ) + slope_v = (value_grid[-1] - value_grid[-2]) / ( + endog_grid[-1] - endog_grid[-2] + ) + policy = policy_grid[-1] + slope_c * (wealth - endog_grid[-1]) + value = value_grid[-1] + slope_v * (wealth - endog_grid[-1]) + else: + policy = np.interp(wealth, endog_grid, policy_grid) + value = np.interp(wealth, endog_grid, value_grid) + return value, marginal_utility_np(policy, params) + + for period in range(last_period - 1, -1, -1): + solved[period] = {} + for partner_state_now in (0, 1): + transition_probs = partner_transition_np(partner_state_now, params) + + for choice_now in (0, 1): + endog_grid = np.empty_like(a_grid) + policy = np.empty_like(a_grid) + value = np.empty_like(a_grid) + + # Income from choice_now (work vs not) arrives at the start + # of period + 1, dcegm-style. + income_next = params["y_work"] * (choice_now == 0) + + for i, a_end in enumerate(a_grid): + expected_marg_util = 0.0 + expected_value = 0.0 + for partner_state_next, prob_partner in enumerate(transition_probs): + if prob_partner == 0.0: + continue + wealth_next = resources_after_transition( + a_end, + partner_state_now, + partner_state_next, + income_next, + params, + ) + + choice_values = np.empty(2) + choice_marg_utils = np.empty(2) + for choice_next in (0, 1): + v, mu = continuation_value_and_marg_util( + period + 1, + wealth_next, + partner_state_next, + choice_next, + ) + choice_values[choice_next] = v + choice_marg_utils[choice_next] = mu + + max_v = choice_values.max() + weights = np.exp((choice_values - max_v) / taste_shock_scale) + choice_probs = weights / weights.sum() + ev_choice = max_v + taste_shock_scale * np.log(weights.sum()) + + expected_marg_util += prob_partner * np.sum( + choice_probs * choice_marg_utils + ) + expected_value += prob_partner * ev_choice + + rhs_euler = beta * (1 + r) * expected_marg_util + c = inverse_marginal_utility_np(rhs_euler, params) + + endog_grid[i] = c + a_end + policy[i] = c + value[i] = ( + consumption_utility(c, choice_now, params) + + beta * expected_value + ) + + solved[period][(partner_state_now, choice_now)] = { + "endog_grid": endog_grid, + "policy": policy, + "value": value, + } + + return solved diff --git a/tests/test_divorce_toy_model.py b/tests/test_divorce_toy_model.py new file mode 100644 index 00000000..a0b33367 --- /dev/null +++ b/tests/test_divorce_toy_model.py @@ -0,0 +1,306 @@ +"""Two-period (and n-period) dcegm divorce toy model: tests. + +Model *mechanism* code lives in `tests/resources/divorce_model/` and takes +its numeric parametrization (`PARAMS`, grid sizes) as explicit arguments -- +none of it is hardcoded there. This file owns the actual numbers: + - `reference.py` -- independent hand-rolled backward-induction EGM solver + (no dcegm, no upper envelope). Transition-based wealth rescaling: halve + on divorce, double on marriage, unchanged otherwise; no multiplier in + utility. + - `dcegm_functions.py` -- the dcegm-facing model (utility, budget + constraint, model setup). Per-period wealth/utility rescaling keyed off + the *current* period's own partner_state only. + +See the dcegm guide "Implementing a divorce/marriage transition without a +lagged partner state" (`docs/source/guides/`) for why these two +different-looking mechanisms are nonetheless the same underlying economics, +and for the general recipe. + +Why `budget_constraint` (dcegm side) divides by the multiplier again +---------------------------------------------------------------------- +dcegm's Euler-equation solver hardcodes the marginal return on savings as +`1 + params["interest_rate"]` +(`src/dcegm/egm/solve_euler_equation.py:166`, +`rhs_euler = marginal_utility_next * (1 + interest_rate) * discount_factor`). +It never differentiates the user-supplied `budget_constraint`, so any extra +multiplier applied to the incoming asset would be silently ignored in the +Euler equation's *price* of savings even though it's correctly reflected in +the wealth *level*. The fix: `budget_constraint` doubles the individual +asset if partnered, lets it earn interest, *and then divides the whole +expression by the same multiplier again*. For the asset term this is just +`mult * a * (1+r) / mult = a * (1+r)` -- the multiplier cancels exactly, so +`d(wealth)/d(a)` is `1+r` regardless of partner_state, matching what dcegm +assumes. +""" + +import numpy as np +import pytest + +from .resources.divorce_model import dcegm_functions as dm +from .resources.divorce_model import reference as ref + +# --------------------------------------------------------------------------- +# Parametrization -- owned here, passed explicitly into every model/solver +# call. Nothing in tests/resources/divorce_model/ hardcodes any of this. +# --------------------------------------------------------------------------- + +RHO = 0.8 # estimated mu_low = mu_high in est_params_alg1_sparse.pkl +DELTA = 0.5 # disutility of work +BETA = 0.96 +R = 0.02 +TASTE_SHOCK_SCALE = 0.2 +Y_WORK = 30.0 +Y_PARTNER = 20.0 # added whenever partnered *this* period, no lagged_choice +A_GRID_MAX = 300.0 +A_GRID_POINTS = 200 +PROB_MARRY = 0.3 # P(married next | single now) +PERSISTENCE_MARRIED = 0.8 # P(married next | married now) + +PARAMS = { + "discount_factor": BETA, + "delta": DELTA, + "rho": RHO, + "interest_rate": R, + "taste_shock_scale": TASTE_SHOCK_SCALE, + "income_shock_std": 0.0, + "income_shock_mean": 0.0, + "consumption_floor": 1e-8, + "y_work": Y_WORK, + "y_partner": Y_PARTNER, + "prob_marry": PROB_MARRY, + "persistence_married": PERSISTENCE_MARRIED, +} + +A_GRID = np.linspace(0.0, A_GRID_MAX, A_GRID_POINTS) +N_PERIODS = 4 + + +# --------------------------------------------------------------------------- +# Shared helpers +# --------------------------------------------------------------------------- + + +def dcegm_raw_arrays(model, solved, period, work0, partner_state_0): + """Dcegm pads its stored arrays with extra points (e.g. a natural borrowing + constraint point below the exogenous grid), so array *index* does not line up 1:1 + with the exogenous grid index -- clean and sort by the endogenous wealth *value* + instead.""" + scs = model.model_structure["state_choice_space"] + # columns: period, lagged_choice, partner_state, ..., choice (last) + row = np.where( + (scs[:, 0] == period) + & (scs[:, 1] == 1) + & (scs[:, 2] == partner_state_0) + & (scs[:, -1] == work0) + )[0][0] + endog = np.asarray(solved.endog_grid[row, 0, :]) + policy = np.asarray(solved.policy[row, 0, :]) + value = np.asarray(solved.value[row, 0, :]) + mask = ~np.isnan(endog) + endog, policy, value = endog[mask], policy[mask], value[mask] + order = np.argsort(endog) + return endog[order], policy[order], value[order] + + +# --------------------------------------------------------------------------- +# Two-period tests +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("choice", [0, 1]) +@pytest.mark.parametrize("partner_state", [0, 1]) +def test_dcegm_final_period_matches_analytic_formula(partner_state, choice): + """Sanity check unaffected by the Euler-equation issue: the terminal + period should just consume everything, value = felicity(scale*wealth,choice).""" + model, solved = dm.build_and_solve(PARAMS, n_periods=2, a_grid=A_GRID) + scs = model.model_structure["state_choice_space"] + row_final = np.where( + (scs[:, 0] == 1) & (scs[:, 2] == partner_state) & (scs[:, -1] == choice) + )[0][0] + endog = np.asarray(solved.endog_grid[row_final, 0, :]) + policy = np.asarray(solved.policy[row_final, 0, :]) + value = np.asarray(solved.value[row_final, 0, :]) + mask = ~np.isnan(endog) & (endog > 0) + np.testing.assert_allclose(policy[mask], endog[mask]) + scale = 2.0 if partner_state == 1 else 1.0 + np.testing.assert_allclose( + value[mask], + ((scale * endog[mask]) ** (1 - RHO) - 1) / (1 - RHO) - (1 - choice) * DELTA, + ) + + +def test_budget_constraint_doubles_then_divides_so_the_asset_return_is_plain(): + """Assets get doubled (partner brings equal wealth), earn interest, then get divided + by the same multiplier again -- so the *asset* contribution to wealth is exactly + `a*(1+r)` regardless of partner_state. + + This is what makes dcegm's hardcoded `1+interest_rate` Euler-equation factor + correct. + + """ + model, _ = dm.build_and_solve(PARAMS, n_periods=2, a_grid=A_GRID) + compute_wealth = model.model_funcs["compute_assets_begin_of_period"] + + def wealth(asset, lagged_choice, partner_state): + return compute_wealth( + period=1, + lagged_choice=lagged_choice, + partner_state=partner_state, + asset_end_of_previous_period=asset, + income_shock_previous_period=0.0, + params=PARAMS, + ) + + # The asset *return* (the slope of wealth in `a`) is exactly 1+r for + # both partner states -- income (own and/or partner) is a constant + # w.r.t. `a`, so it drops out of the slope regardless of its value. + for partner_state in (0, 1): + for lagged_choice in (0, 1): + slope = ( + wealth(200.0, lagged_choice, partner_state) + - wealth(100.0, lagged_choice, partner_state) + ) / 100.0 + np.testing.assert_allclose(slope, 1 + R) + + # Own income (lagged_choice == 0) and partner income (partner_state == 1) + # both still get divided by the multiplier when partnered. + wealth_single_no_income = wealth(100.0, 1, 0) + wealth_married_no_income = wealth(100.0, 1, 1) + np.testing.assert_allclose( + wealth_married_no_income - 100.0 * (1 + R), Y_PARTNER / 2 + ) + np.testing.assert_allclose(wealth_single_no_income, 100.0 * (1 + R)) + + wealth_single_income = wealth(100.0, 0, 0) + wealth_married_income = wealth(100.0, 0, 1) + np.testing.assert_allclose(wealth_single_income - wealth_single_no_income, Y_WORK) + np.testing.assert_allclose( + wealth_married_income - wealth_married_no_income, Y_WORK / 2 + ) + + +@pytest.mark.parametrize("partner_state", [0, 1]) +@pytest.mark.parametrize("work0", [0, 1]) +def test_dcegm_policy_matches_hand_solved_reference(partner_state, work0): + """Cross-check dcegm's solved period-0 policy/value against the manual + backward-induction reference, evaluated at dcegm's own native points -- + no interpolation on either side. + + dcegm's EGM identity `endog_grid = policy + a0_end` always holds for its + own stored arrays, so `a0_end = endog_dcegm - policy_dcegm` recovers + exactly which exogenous asset point produced each entry. + + dcegm's own state/policy live in *individual* terms; `utility_func` + internally evaluates felicity at `scale * consumption`, i.e. at the + *joint* quantity. So the reference (which has no such scale) has to be + queried at `scale * a0_end` to land on the same joint quantity dcegm's + utility implicitly uses. The resulting reference `policy` is then + joint-scale and needs dividing by `scale` to compare against dcegm's + individual-scale `policy`; `value` needs no such correction. Verified to + match to floating-point precision (~1e-13) for all four + (work0, partner_state) combinations. + """ + model, solved = dm.build_and_solve(PARAMS, n_periods=2, a_grid=A_GRID) + endog_dcegm, policy_dcegm, value_dcegm = dcegm_raw_arrays( + model, solved, period=0, work0=work0, partner_state_0=partner_state + ) + scale = 2.0 if partner_state == 1 else 1.0 + + a0_end_dcegm = endog_dcegm - policy_dcegm + # Skip near-zero a0_end: a degenerate corner (consumption_floor binds) + # both sides handle slightly differently. + keep = a0_end_dcegm > 1.0 + + ref_policy = np.empty(keep.sum()) + ref_value = np.empty(keep.sum()) + for j, a0_end in enumerate(a0_end_dcegm[keep]): + point = ref.solve_reference_at_point( + scale * a0_end, partner_state, work0, PARAMS + ) + ref_policy[j] = point["policy"] / scale + ref_value[j] = point["value"] + + np.testing.assert_allclose(policy_dcegm[keep], ref_policy, rtol=1e-6) + np.testing.assert_allclose(value_dcegm[keep], ref_value, rtol=1e-6) + + +# --------------------------------------------------------------------------- +# Four-period tests +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("period", list(range(N_PERIODS))) +@pytest.mark.parametrize("choice", [0, 1]) +@pytest.mark.parametrize("partner_state", [0, 1]) +def test_dcegm_policy_is_finite_and_monotone_for_n_periods( + partner_state, choice, period +): + """Basic sanity check that the n-period model solves cleanly at every + period: policy should be finite, non-negative, and (weakly) increasing + in wealth (more resources -> (weakly) more consumption).""" + model, solved = dm.build_and_solve(PARAMS, n_periods=N_PERIODS, a_grid=A_GRID) + endog, policy, value = dcegm_raw_arrays( + model, solved, period=period, work0=choice, partner_state_0=partner_state + ) + assert np.all(np.isfinite(policy)) + assert np.all(policy >= 0) + assert np.all(np.diff(policy) >= -1e-8) # allow tiny numerical noise + assert np.all(np.isfinite(value[endog > 1.0])) + + +@pytest.mark.parametrize("partner_state", [0, 1]) +@pytest.mark.parametrize("work0", [0, 1]) +def test_dcegm_policy_matches_hand_solved_reference_n_periods(partner_state, work0): + """Same cross-check as the two-period test, but for the full N_PERIODS-period model, + comparing period 0's policy/value. + + Unlike the two-period case, periods before the terminal one are no + longer analytic on the reference side either -- `solve_reference` + interpolates its own stored grids for the continuation value, exactly + as dcegm does internally, including two edge cases that must be handled + explicitly rather than left to plain `np.interp` (which just clamps to + a constant outside its domain, silently wrong in both directions): + querying below the grid's own minimum (the borrowing constraint binds, + consume everything) and above its maximum (linear extrapolation, since + a marriage transition can double continuation wealth past what the + grid was built to cover). See `solve_reference`'s + `continuation_value_and_marg_util` for both. With those handled, this + is no longer an exact floating-point match (periods before the terminal + one are genuinely interpolated, on both sides, same as dcegm), but it's + close: max relative error ~4e-4. See + `test_dcegm_policy_matches_hand_solved_reference` above for the + interpolation-free two-period version. + + """ + model, solved = dm.build_and_solve(PARAMS, n_periods=N_PERIODS, a_grid=A_GRID) + endog_dcegm, policy_dcegm, value_dcegm = dcegm_raw_arrays( + model, solved, period=0, work0=work0, partner_state_0=partner_state + ) + scale = 2.0 if partner_state == 1 else 1.0 + + a0_end_dcegm = endog_dcegm - policy_dcegm + keep = a0_end_dcegm > 1.0 + + # A married agent is queried at 2*a0_end (see the docstring above), which + # can reach 2*A_GRID_MAX -- give the reference solve a wider grid than + # dcegm's own so that the outer lookup below doesn't clamp at its top edge. + wide_a_grid = np.linspace(0.0, 2 * A_GRID_MAX, 2 * A_GRID_POINTS) + ref_solved = ref.solve_reference(N_PERIODS, PARAMS, wide_a_grid) + ref_period0 = ref_solved[0][(partner_state, work0)] + + # ref_period0["policy"/"value"][i] is indexed by wide_a_grid[i] -- the + # *exogenous asset* grid, not the endogenous wealth grid -- since that's + # what solve_reference iterates over. Interpolate on that axis, same as + # solve_reference_at_point conceptually does exactly (just off a + # precomputed grid here instead of a fresh point solve). + ref_policy = ( + np.interp(scale * a0_end_dcegm[keep], wide_a_grid, ref_period0["policy"]) + / scale + ) + ref_value = np.interp(scale * a0_end_dcegm[keep], wide_a_grid, ref_period0["value"]) + + # Actual max relative error is ~4e-4 (median ~3e-6) with the borrowing + # constraint and top-of-grid extrapolation both handled explicitly in + # `solve_reference` -- this tolerance has headroom, not a fudge factor. + np.testing.assert_allclose(policy_dcegm[keep], ref_policy, rtol=1e-3) + np.testing.assert_allclose(value_dcegm[keep], ref_value, rtol=1e-3, atol=1e-3) diff --git a/tests/test_law_of_motion.py b/tests/test_law_of_motion.py index 3f1f9e8c..5854387f 100644 --- a/tests/test_law_of_motion.py +++ b/tests/test_law_of_motion.py @@ -10,8 +10,13 @@ from scipy.special import roots_sh_legendre from scipy.stats import norm +import dcegm import dcegm.toy_models as toy_models -from dcegm.law_of_motion import calculate_continuous_state +from dcegm.law_of_motion import ( + calc_cont_grids_next_period, + calc_law_of_motion_for_state_choices, + calculate_continuous_state, +) from dcegm.pre_processing.check_params import process_params from dcegm.toy_models.cons_ret_model_dcegm_paper import budget_constraint @@ -197,3 +202,108 @@ def test_wealth_and_second_continuous_state(model_name, max_wealth, n_grid_point ) aaae(exp_next, experience_next) + + +# ===================================================================================== +# On-demand vs. full-state-space law of motion equivalence +# ===================================================================================== + + +def _check_subset_matches_full(model, params): + """calc_law_of_motion_for_state_choices on a state-choice subset must match + calc_cont_grids_next_period's full-state-space computation for the same underlying + states -- and must be identical across different choices sharing the same state + (confirming the choice-drop / no-dedup behavior is correct).""" + model_structure = model.model_structure + model_config = model.model_config + model_funcs = model.model_funcs + continuous_states_info = model_config["continuous_states_info"] + + full = calc_cont_grids_next_period( + params=params, + income_shock_draws_unscaled=model.income_shock_draws_unscaled, + model_structure=model_structure, + model_config=model_config, + model_funcs=model_funcs, + ) + + state_choice_space_dict = model_structure["state_choice_space_dict"] + map_state_choice_to_parent_state = model_structure[ + "map_state_choice_to_parent_state" + ] + + # Pick >=2 state-choice indices sharing the same parent state (different + # choices, same state) plus a few others, to test the no-dedup property. + values, counts = np.unique(map_state_choice_to_parent_state, return_counts=True) + shared_parent_state = values[counts >= 2][0] + idx_sharing_state = np.where( + map_state_choice_to_parent_state == shared_parent_state + )[0][:2] + other_idx = np.where(map_state_choice_to_parent_state != shared_parent_state)[0][:3] + test_idx = np.concatenate([idx_sharing_state, other_idx]) + + state_choice_subset = { + key: jnp.asarray(var[test_idx]) for key, var in state_choice_space_dict.items() + } + + income_shock_std = model_funcs["read_funcs"]["income_shock_std"](params) + income_shock_mean = model_funcs["read_funcs"]["income_shock_mean"](params) + income_shocks_scaled = ( + model.income_shock_draws_unscaled * income_shock_std + income_shock_mean + ) + + subset_result = calc_law_of_motion_for_state_choices( + state_choice_vec=state_choice_subset, + continuous_state_space=model_structure["continuous_state_space"], + assets_grid_end_of_period=continuous_states_info["assets_grid_end_of_period"], + income_shocks_scaled=income_shocks_scaled, + params=params, + model_funcs=model_funcs, + has_additional_continuous_states=continuous_states_info[ + "has_additional_continuous_state" + ], + ) + + expected_parent_states = map_state_choice_to_parent_state[test_idx] + expected = full["assets_begin_of_period"][expected_parent_states] + np.testing.assert_allclose( + np.asarray(subset_result["assets_begin_of_period"]), np.asarray(expected) + ) + if continuous_states_info["has_additional_continuous_state"]: + for key, expected_cont in full["continuous_states"].items(): + np.testing.assert_allclose( + np.asarray(subset_result["continuous_states"][key]), + np.asarray(expected_cont[expected_parent_states]), + ) + + # Two indices sharing the same parent state (different choices) must give + # IDENTICAL wealth transitions -- confirms "choice" has no effect, as intended. + np.testing.assert_array_equal( + np.asarray(subset_result["assets_begin_of_period"][0]), + np.asarray(subset_result["assets_begin_of_period"][1]), + ) + + +def _build_model(model_name): + model_funcs = toy_models.load_example_model_functions(model_name) + params, model_specs, model_config = ( + toy_models.load_example_params_model_specs_and_config(model_name) + ) + model = dcegm.setup_model( + model_config=model_config, + model_specs=model_specs, + **model_funcs, + ) + return model, params + + +def test_law_of_motion_subset_matches_full_discrete(): + # Retirement model: >=2 choices per state, no additional continuous state. + model, params = _build_model("dcegm_paper_retirement_no_shocks") + _check_subset_matches_full(model, params) + + +def test_law_of_motion_subset_matches_full_cont_exp(): + # >=2 choices per state, plus an additional continuous state ("experience"). + model, params = _build_model("with_cont_exp") + _check_subset_matches_full(model, params) diff --git a/tests/test_model_config.py b/tests/test_model_config.py index c496ff2c..5e2c2f15 100644 --- a/tests/test_model_config.py +++ b/tests/test_model_config.py @@ -123,3 +123,37 @@ def test_upper_envelope_method_default(valid_model_config): check_model_config_and_process(valid_model_config)["upper_envelope"]["method"] == "fues" ) + + +def test_skip_endog_grid_storage_false_for_fues(valid_model_config): + options = check_model_config_and_process(valid_model_config) + assert options["upper_envelope"]["skip_endog_grid_storage"] is False + + +def test_dj_wealth_grid_and_skip_flag(valid_model_config): + valid_model_config["upper_envelope"] = {"method": "druedahl_jorgensen"} + valid_model_config["continuous_states"]["assets_begin_of_period"] = np.linspace( + 0, 10, 11 + ) + options = check_model_config_and_process(valid_model_config) + + dj_wealth_grid = options["continuous_states_info"]["dj_wealth_grid"] + expected = np.concatenate( + ([0.0], valid_model_config["continuous_states"]["assets_begin_of_period"]) + ) + assert_array_equal(np.asarray(dj_wealth_grid), expected) + assert dj_wealth_grid.shape[0] == options["n_total_wealth_grid"] + assert options["upper_envelope"]["skip_endog_grid_storage"] is True + + +def test_skip_endog_grid_storage_false_for_single_choice_dj(valid_model_config): + # With a single discrete choice, the upper envelope is skipped entirely, so the + # stored endog_grid is not the fixed Druedahl-Jorgensen grid and must still be + # stored/read normally, even though the method is "druedahl_jorgensen". + valid_model_config["choices"] = [1] + valid_model_config["upper_envelope"] = {"method": "druedahl_jorgensen"} + valid_model_config["continuous_states"]["assets_begin_of_period"] = np.linspace( + 0, 10, 11 + ) + options = check_model_config_and_process(valid_model_config) + assert options["upper_envelope"]["skip_endog_grid_storage"] is False diff --git a/tests/test_two_occupation_model.py b/tests/test_two_occupation_model.py index 97ff2e43..d07db308 100644 --- a/tests/test_two_occupation_model.py +++ b/tests/test_two_occupation_model.py @@ -566,6 +566,14 @@ def aligned_states(): # ==================================================================================== +def test_dj_models_do_not_store_endog_grid(solved_discrete, solved_cont_exp): + """Druedahl-Jorgensen solutions skip storing the (redundant) endog_grid array.""" + for solved in (solved_discrete, solved_cont_exp): + assert solved.endog_grid is None + assert solved.value.shape == solved.policy.shape + assert jnp.isfinite(solved.value).any() + + def test_discrete_interface_joint_vs_separate(solved_discrete): """Joint policy+value query matches individual queries for the discrete model.""" states_eval = { diff --git a/tests/test_two_period_continuous_experience.py b/tests/test_two_period_continuous_experience.py index d4f15423..0a2389d3 100644 --- a/tests/test_two_period_continuous_experience.py +++ b/tests/test_two_period_continuous_experience.py @@ -10,7 +10,6 @@ import dcegm import dcegm.toy_models as toy_models from dcegm.final_periods import solve_final_period -from dcegm.law_of_motion import calc_cont_grids_next_period from dcegm.numerical_integration import quadrature_legendre from dcegm.pre_processing.sol_container import create_solution_container from dcegm.solve_single_period import solve_for_interpolated_values @@ -262,7 +261,7 @@ def create_test_inputs(): model_config = model.model_config ( - cont_grids_next_period, + income_shocks_scaled, income_shock_draws_unscaled, income_shock_weights, taste_shock_scale, @@ -285,17 +284,17 @@ def create_test_inputs(): idx_state_choices_final_period=last_two_period_batch_info_cont[ "idx_state_choices_final_period" ], - idx_parent_states_final_period=last_two_period_batch_info_cont[ - "idxs_parent_states_final_period" - ], state_choice_mat_final_period=last_two_period_batch_info_cont[ "state_choice_mat_final_period" ], - cont_grids_next_period=cont_grids_next_period, + income_shocks_scaled=income_shocks_scaled, continuous_states_info=model_config["continuous_states_info"], model_structure=model.model_structure, params=params, upper_envelope_method=model_config["upper_envelope"]["method"], + skip_endog_grid_storage=model_config["upper_envelope"][ + "skip_endog_grid_storage" + ], model_funcs=model_funcs_cont, value_solved=value_solved, policy_solved=policy_solved, @@ -452,12 +451,10 @@ def _get_solve_last_two_periods_args(model, params, has_second_continuous_state) model_structure = model.model_structure model_funcs = model.model_funcs - cont_grids_next_period = calc_cont_grids_next_period( - params=params, - income_shock_draws_unscaled=income_shock_draws_unscaled, - model_structure=model_structure, - model_config=model_config, - model_funcs=model_funcs, + income_shock_mean = model_funcs["read_funcs"]["income_shock_mean"](params) + income_shock_std = model_funcs["read_funcs"]["income_shock_std"](params) + income_shocks_scaled = ( + income_shock_draws_unscaled * income_shock_std + income_shock_mean ) n_continuous_state_combinations = model_structure["continuous_state_space"][ @@ -475,7 +472,7 @@ def _get_solve_last_two_periods_args(model, params, has_second_continuous_state) ) return ( - cont_grids_next_period, + income_shocks_scaled, income_shock_draws_unscaled, income_shock_weights, taste_shock_scale,