From cc2fdab63feeed5d8ee391d2a56fa8648066474c Mon Sep 17 00:00:00 2001 From: Max Ghenis Date: Thu, 1 Oct 2026 14:10:57 -0400 Subject: [PATCH 1/7] Calculate formula branches under their overrides; limit the CTC by actual liability A policyengine-core branch starts as a copy of every array its parent has cached, and set_input on it clears none of them, so a value the parent calculated from the overridden input answers for the branch. - ctc_limiting_tax_liability no longer recomputes the liability without SALT in a "no_salt" branch. The branch usually inherited the liability with SALT, so the CTC limit depended on which variable was calculated first. 26 U.S.C. 26(a) limits the credit by the actual tax liability, which reflects the SALT deduction. - get_override_branch (tools/override_branch.py) creates a comparison branch once per period and, when the parent has already calculated an overridden input, drops the arrays the branch copied except inputs. The itemization, Delaware and Virginia EITC refundability, Idaho aged or disabled, Missouri TANF caretaker and Medicaid SSI-supplement branches use it; the Alabama 2020-IRC and New York pinned-parameter branches are created per period. Co-Authored-By: Claude Opus 5.5 --- .../fix-branch-override-shadowing.fixed.md | 1 + .../tests/core/test_override_branches.py | 347 ++++++++++++++++++ .../tests/test_ctc_itemizing_branch_cycle.py | 6 +- policyengine_us/tools/general.py | 4 + policyengine_us/tools/override_branch.py | 152 ++++++++ ...icaid_enrolled_for_ssi_state_supplement.py | 15 +- .../refundable/ctc_limiting_tax_liability.py | 36 +- .../deductions/tax_liability_if_itemizing.py | 8 +- .../tax_liability_if_not_itemizing.py | 8 +- .../al_federal_income_tax_deduction.py | 2 +- ...ome_tax_if_claiming_non_refundable_eitc.py | 8 +- ..._income_tax_if_claiming_refundable_eitc.py | 8 +- ...ax_if_receiving_aged_or_disabled_credit.py | 12 +- ...if_receiving_aged_or_disabled_deduction.py | 12 +- ...o_tanf_if_non_parent_caretaker_excluded.py | 12 +- ...o_tanf_if_non_parent_caretaker_included.py | 8 +- .../tax/income/credits/ctc/ny_ctc_pre_2024.py | 2 +- .../credits/ctc/ny_ctc_pre_2024_eligible.py | 2 +- .../states/ny/tax/income/credits/ny_eitc.py | 2 +- ...ome_tax_if_claiming_non_refundable_eitc.py | 8 +- ..._income_tax_if_claiming_refundable_eitc.py | 8 +- 21 files changed, 593 insertions(+), 68 deletions(-) create mode 100644 changelog.d/fix-branch-override-shadowing.fixed.md create mode 100644 policyengine_us/tests/core/test_override_branches.py create mode 100644 policyengine_us/tools/override_branch.py diff --git a/changelog.d/fix-branch-override-shadowing.fixed.md b/changelog.d/fix-branch-override-shadowing.fixed.md new file mode 100644 index 00000000000..b183cc39fd7 --- /dev/null +++ b/changelog.d/fix-branch-override-shadowing.fixed.md @@ -0,0 +1 @@ +Limit the non-refundable Child Tax Credit by the actual tax liability, SALT deduction included (26 U.S.C. 26(a)), instead of a recomputation without SALT that applied or not depending on which variables were calculated first; and make the itemization, Delaware and Virginia EITC, Idaho aged or disabled, Missouri TANF caretaker and Medicaid SSI-supplement comparison branches calculate under their overridden inputs even when the simulation has already calculated those inputs, and again for each period. diff --git a/policyengine_us/tests/core/test_override_branches.py b/policyengine_us/tests/core/test_override_branches.py new file mode 100644 index 00000000000..07434effd78 --- /dev/null +++ b/policyengine_us/tests/core/test_override_branches.py @@ -0,0 +1,347 @@ +"""Formula branches calculate under their overrides, whatever was cached first. + +Several formulas compare a tax unit's liability under alternative choices by +calculating it in a branch with one input overridden: itemizing or not +(``tax_unit_itemizes``), Delaware and Virginia EITC refundability, the Idaho +aged or disabled credit or deduction, Missouri TANF caretaker inclusion and +Medicaid for SSI state supplements. A policyengine-core branch starts as a copy +of every array its parent has cached, and setting an input on it clears none +of them, so a value the parent calculated from the old input answers for the +branch. ``get_override_branch`` drops the copied values when the parent has +already calculated the overridden input, and creates the branch again for each +period. + +The non-refundable CTC used to be limited by the tax liability recomputed +without the SALT deduction, in a "no_salt" branch. The branch usually inherited +the liability with SALT, so the CTC depended on which variables a caller asked +for first. 26 U.S.C. 26(a) limits the credit by the actual tax liability, SALT +deduction included, and ``ctc_limiting_tax_liability`` now reads it directly. + +A seeded sample of households shares one simulation; the tests check: + +1. The CTC-limiting liability is income tax before credits less the other + non-refundable credits, and the non-refundable and refundable parts add up + to the credit. +2. Every reported variable is the same whichever variable is calculated first. +3. Each comparison branch equals a fresh simulation that sets the overridden + input before anything is calculated (a differential test of the branch + against the reference path). +4. The same holds when the parent has already calculated the overridden + input, and when an earlier year was calculated first. +""" + +import numpy as np +import pytest + +from policyengine_us import Simulation +from policyengine_us.tools.override_branch import ( + drop_inherited_values, + get_override_branch, +) + +YEAR = 2026 +SEED = 20261001 +N = 48 +STATES = ["CA", "NY", "VA", "DE", "ID", "NJ", "MA", "IL", "TX", "OR"] +REPORTED = [ + "income_tax", + "income_tax_before_credits", + "ctc_limiting_tax_liability", + "non_refundable_ctc", + "refundable_ctc", + "tax_unit_itemizes", + "state_income_tax", + "household_net_income", +] + + +def _sample(): + rng = np.random.default_rng(SEED) + households = [] + for i in range(N): + married = bool(rng.random() < 0.6) + children = int(rng.integers(0, 4)) + earnings = float(np.round(np.exp(rng.uniform(np.log(8_000), np.log(300_000))))) + households.append( + dict( + state=STATES[i % len(STATES)], + married=married, + children=children, + earnings=earnings, + spouse_share=float(rng.choice([0, 0.2, 0.5])) if married else 0, + mortgage=float(rng.choice([0, 8_000, 18_000, 30_000])), + property_tax=float(rng.choice([0, 3_000, 9_000, 16_000])), + charity=float(rng.choice([0, 2_000, 9_000])), + aged_parent=bool(rng.random() < 0.15), + ) + ) + return households + + +def _situation(households, years=(YEAR,), itemizes=None): + people, units, marital = {}, {}, {} + for i, h in enumerate(households): + members = [] + + def add(name, **inputs): + people[name] = {k: {y: v for y in years} for k, v in inputs.items()} + members.append(name) + + head = f"h{i}" + earnings = h["earnings"] + add( + head, + age=40, + employment_income=earnings * (1 - h["spouse_share"]), + deductible_mortgage_interest=h["mortgage"], + real_estate_taxes=h["property_tax"], + charitable_cash_donations=h["charity"], + ) + marital[f"mu_{head}"] = {"members": [head]} + if h["married"]: + add(f"s{i}", age=38, employment_income=earnings * h["spouse_share"]) + marital[f"mu_{head}"]["members"].append(f"s{i}") + for c in range(h["children"]): + add(f"c{i}_{c}", age=3 + 4 * c, is_tax_unit_dependent=True) + marital[f"mu_c{i}_{c}"] = {"members": [f"c{i}_{c}"]} + if h["aged_parent"]: + add( + f"p{i}", + age=74, + is_tax_unit_dependent=True, + share_of_care_and_support_costs_paid_by_tax_filer=1.0, + ) + marital[f"mu_p{i}"] = {"members": [f"p{i}"]} + units[i] = members + tax_units = {f"tu{i}": {"members": m} for i, m in units.items()} + if itemizes is not None: + for i, unit in enumerate(tax_units.values()): + unit["tax_unit_itemizes"] = {y: bool(itemizes[i]) for y in years} + return { + "people": people, + "tax_units": tax_units, + "spm_units": {f"spm{i}": {"members": m} for i, m in units.items()}, + "families": {f"fam{i}": {"members": m} for i, m in units.items()}, + "marital_units": marital, + "households": { + f"hh{i}": { + "members": m, + "state_code": {y: households[i]["state"] for y in years}, + } + for i, m in units.items() + }, + } + + +@pytest.fixture(scope="module") +def households(): + return _sample() + + +@pytest.fixture(scope="module") +def baseline(households): + simulation = Simulation(situation=_situation(households)) + return {v: simulation.calculate(v, YEAR) for v in REPORTED} + + +def test_ctc_limit_is_actual_liability_less_other_credits(households): + simulation = Simulation(situation=_situation(households)) + limit = simulation.calculate("ctc_limiting_tax_liability", YEAR) + before_credits = simulation.calculate("income_tax_before_credits", YEAR) + credits = simulation.tax_benefit_system.parameters(YEAR).gov.irs.credits + other = sum( + simulation.calculate(credit, YEAR) + for credit in credits.non_refundable + if credit != "non_refundable_ctc" + ) + np.testing.assert_allclose(limit, np.maximum(0, before_credits - other)) + ctc = simulation.calculate("ctc", YEAR) + non_refundable = simulation.calculate("non_refundable_ctc", YEAR) + refundable = simulation.calculate("refundable_ctc", YEAR) + # non_refundable_ctc is the credit less its refundable part; the tax + # limit applies when non-refundable credits are capped. + assert (refundable <= ctc + 0.01).all() + np.testing.assert_allclose(non_refundable + refundable, ctc, atol=0.01) + # The sample includes itemizers with SALT whose CTC the limit binds. + salt = simulation.calculate("salt_deduction", YEAR) + itemizes = simulation.calculate("tax_unit_itemizes", YEAR) + assert (itemizes & (salt > 0) & (limit < ctc)).any() + + +@pytest.mark.parametrize( + "first", + [ + "income_tax", + "refundable_ctc", + "ctc_value", + "state_income_tax", + "tax_liability_if_itemizing", + "spm_unit_net_income", + ], +) +def test_results_do_not_depend_on_calculation_order(households, baseline, first): + simulation = Simulation(situation=_situation(households)) + simulation.calculate(first, YEAR) + for variable in REPORTED: + np.testing.assert_array_equal( + simulation.calculate(variable, YEAR), + baseline[variable], + err_msg=f"{variable} changes when {first} is calculated first", + ) + + +def _fresh(households, variable, overrides, year=YEAR): + simulation = Simulation(situation=_situation(households, years=(year,))) + for name, value in overrides.items(): + simulation.set_input(name, year, value) + return simulation.calculate(variable, year) + + +def test_itemization_branches_match_fresh_simulations(households): + simulation = Simulation(situation=_situation(households)) + simulation.calculate("household_net_income", YEAR) + n = len(households) + for comparison, itemizes in ( + ("tax_liability_if_itemizing", True), + ("tax_liability_if_not_itemizing", False), + ): + np.testing.assert_allclose( + simulation.calculate(comparison, YEAR), + _fresh(households, "income_tax", {"tax_unit_itemizes": [itemizes] * n}), + atol=0.01, + err_msg=comparison, + ) + + +@pytest.mark.parametrize( + "comparison,variable,override,value", + [ + ( + "de_income_tax_if_claiming_refundable_eitc", + "de_income_tax", + "de_claims_refundable_eitc", + True, + ), + ( + "de_income_tax_if_claiming_non_refundable_eitc", + "de_income_tax", + "de_claims_refundable_eitc", + False, + ), + ( + "va_income_tax_if_claiming_refundable_eitc", + "va_income_tax", + "va_claims_refundable_eitc", + True, + ), + ( + "va_income_tax_if_claiming_non_refundable_eitc", + "va_income_tax", + "va_claims_refundable_eitc", + False, + ), + ( + "id_income_tax_if_receiving_aged_or_disabled_credit", + "id_income_tax", + "id_receives_aged_or_disabled_credit", + True, + ), + ( + "id_income_tax_if_receiving_aged_or_disabled_deduction", + "id_income_tax", + "id_receives_aged_or_disabled_credit", + False, + ), + ], +) +def test_state_choice_branches_match_fresh_simulations( + households, comparison, variable, override, value +): + simulation = Simulation(situation=_situation(households)) + simulation.calculate("household_net_income", YEAR) + np.testing.assert_allclose( + simulation.calculate(comparison, YEAR), + _fresh(households, variable, {override: [value] * len(households)}), + atol=0.01, + ) + + +def test_branch_after_parent_calculated_the_overridden_input(households): + # tax_unit_itemizes is an input here, so the parent calculates income tax + # without the itemizing branch; the branch is created afterwards and must + # not answer with the parent's income tax. + n = len(households) + situation = _situation(households, itemizes=[False] * n) + simulation = Simulation(situation=situation) + not_itemizing = simulation.calculate("income_tax", YEAR) + itemizing = simulation.calculate("tax_liability_if_itemizing", YEAR) + expected = _fresh(households, "income_tax", {"tax_unit_itemizes": [True] * n}) + np.testing.assert_allclose(itemizing, expected, atol=0.01) + assert not np.allclose(itemizing, not_itemizing) + + +def test_later_year_branches_match_single_year_simulation(households): + years = (2025, YEAR) + simulation = Simulation(situation=_situation(households, years=years)) + simulation.calculate("tax_liability_if_itemizing", 2025) + simulation.calculate("income_tax", 2025) + fresh = Simulation(situation=_situation(households, years=years)) + for variable in REPORTED + ["tax_liability_if_itemizing"]: + np.testing.assert_array_equal( + simulation.calculate(variable, YEAR), + fresh.calculate(variable, YEAR), + err_msg=f"{variable} for {YEAR} depends on 2025 being calculated first", + ) + assert simulation.branches["itemizing"].branch_period.start.year == YEAR + + +def test_get_override_branch_reuses_within_period_and_recreates_otherwise( + households, +): + simulation = Simulation(situation=_situation(households, years=(2025, YEAR))) + n = len(households) + ones = np.ones(n, dtype=bool) + first = get_override_branch( + simulation, "test_branch", YEAR, {"tax_unit_itemizes": ones} + ) + assert ( + get_override_branch( + simulation, "test_branch", YEAR, {"tax_unit_itemizes": ones} + ) + is first + ) + other_inputs = get_override_branch( + simulation, "test_branch", YEAR, {"tax_unit_itemizes": ~ones} + ) + assert other_inputs is not first + other_period = get_override_branch( + simulation, "test_branch", 2025, {"tax_unit_itemizes": ones} + ) + assert other_period is not other_inputs + assert simulation.branches["test_branch"] is other_period + + +def test_drop_inherited_values_keeps_inputs_only(households): + n = len(households) + itemizes = np.ones(n, dtype=bool) + simulation = Simulation(situation=_situation(households)) + simulation.calculate("income_tax", YEAR) + # A plain branch with an input of its own, and a branch nested in it. + parent = simulation.get_branch("parent_with_input") + parent.set_input("tax_unit_itemizes", YEAR, itemizes) + child = parent.get_branch("child") + drop_inherited_values(child) + # Inputs survive: the situation's and the one set on the parent branch. + np.testing.assert_array_equal( + child.get_array("employment_income", YEAR), + simulation.get_array("employment_income", YEAR), + ) + np.testing.assert_array_equal(child.get_array("tax_unit_itemizes", YEAR), itemizes) + # Calculated values are gone, and are calculated again from the inputs. + assert child.get_array("income_tax", YEAR) is None + assert child.get_array("taxable_income", YEAR) is None + np.testing.assert_allclose( + child.calculate("taxable_income", YEAR), + _fresh(households, "taxable_income", {"tax_unit_itemizes": itemizes}), + atol=0.01, + ) diff --git a/policyengine_us/tests/test_ctc_itemizing_branch_cycle.py b/policyengine_us/tests/test_ctc_itemizing_branch_cycle.py index 9d797ce7b43..94ac4a349da 100644 --- a/policyengine_us/tests/test_ctc_itemizing_branch_cycle.py +++ b/policyengine_us/tests/test_ctc_itemizing_branch_cycle.py @@ -15,9 +15,9 @@ consumers (e.g. policyengine.py's household-impact integration tests) on `policyengine-core >= 3.24`. -The fix in `ctc_limiting_tax_liability.py` propagates the parent's -`tax_unit_itemizes` value to the no_salt child branch so the -`tax_unit_itemizes` formula is never re-entered there. +`ctc_limiting_tax_liability` no longer creates the no_salt branch: it reads +the branch's own `income_tax_before_credits`, which the itemizing branch +calculates with `tax_unit_itemizes` set, so the chain cannot re-enter it. """ import numpy as np diff --git a/policyengine_us/tools/general.py b/policyengine_us/tools/general.py index 7ad450690e3..0696875e350 100644 --- a/policyengine_us/tools/general.py +++ b/policyengine_us/tools/general.py @@ -1,6 +1,10 @@ from policyengine_core.model_api import * from policyengine_us.entities import * from policyengine_us.tools.branched_simulation import BranchedSimulation +from policyengine_us.tools.override_branch import ( + get_branch_for_period, + get_override_branch, +) from pathlib import Path import pandas as pd from policyengine_us.typing import Formula diff --git a/policyengine_us/tools/override_branch.py b/policyengine_us/tools/override_branch.py new file mode 100644 index 00000000000..d48368c3647 --- /dev/null +++ b/policyengine_us/tools/override_branch.py @@ -0,0 +1,152 @@ +"""Branches that calculate one period under overridden inputs. + +A policyengine-core branch starts as a copy of every array its parent has +cached, and ``set_input`` on the branch stores the new value without clearing +anything that was calculated from the old one. A branch therefore answers any +variable its parent has already calculated with the parent's value, whatever +the branch's inputs say. + +``get_override_branch`` makes the override reach every value the branch +calculates: + +- A branch serves one period. Asked for another period, it is created again + from the parent as the parent stands then, as a simulation calculating only + that period would create it. +- When the branch is created, it keeps the parent's cache only if the parent + has no value yet for any overridden variable and period. A cached value + cannot have been calculated from a value that did not exist, so the + parent's cache is then safe to share; this is the usual case, where a + formula branches while its parent is still calculating the variable the + branch overrides. Otherwise the branch drops every array it copied except + inputs, and calculates the rest itself. +- A branch reused within its period with different inputs is created again. +""" + +from typing import Dict, Tuple, Union + +import numpy as np +from policyengine_core.periods import Period +from policyengine_core.periods import period as to_period +from policyengine_core.simulations import Simulation + +Override = Union[np.ndarray, Tuple[Period, np.ndarray]] + + +def _overrides(period: Period, inputs: Dict[str, Override]): + for variable, value in inputs.items(): + if isinstance(value, tuple): + input_period, value = value + yield variable, to_period(input_period), np.asarray(value) + else: + yield variable, period, np.asarray(value) + + +def _is_known(simulation: Simulation, variable: str, period: Period) -> bool: + holder = simulation.get_holder(variable) + if holder.variable.is_neutralized: + return False + return holder.get_array(period, simulation.branch_name) is not None + + +def drop_inherited_values(branch: Simulation) -> None: + """Delete every array ``branch`` holds except inputs. + + Inputs are the variables the simulation was built with and each value set + through ``set_input`` on a branch this one reads. + """ + input_variables = set(branch.input_variables) + user_input_keys = getattr(branch, "_user_input_keys", set()) + visible_branches = set(branch._get_visible_branch_names()) + for population in branch.populations.values(): + for name, holder in population._holders.items(): + if name in input_variables: + continue + for branch_name, known_period in holder.get_known_branch_periods(): + if ( + branch_name in visible_branches + and (name, branch_name, known_period) in user_input_keys + ): + continue + # Exact key: ``Holder.delete_arrays`` would also delete any + # input stored at a sub-period of ``known_period``. + key = f"{branch_name}:{known_period}" + holder._memory_storage._arrays.pop(key, None) + if holder._disk_storage is not None: + holder._disk_storage._files.pop( + f"{branch_name}_{known_period}", None + ) + branch._fast_cache = {} + + +def get_override_branch( + simulation: Simulation, + name: str, + period: Period, + inputs: Dict[str, Override], +) -> Simulation: + """Return ``simulation``'s branch ``name`` for ``period``, with ``inputs`` set. + + ``inputs`` maps each overridden variable to its value for ``period``, or + to a ``(period, value)`` pair for another period (e.g. a month). + """ + overrides = list(_overrides(period, inputs)) + if name == simulation.branch_name: + # Already inside this branch (core's get_branch returns the + # simulation itself): only the inputs need setting. + for variable, input_period, value in overrides: + simulation.set_input(variable, input_period, value) + return simulation + branch = simulation.branches.get(name) + if branch is not None and ( + getattr(branch, "branch_period", None) != period + or not _same_overrides(branch, overrides) + ): + del simulation.branches[name] + branch = None + if branch is None: + parent_knows_override = any( + _is_known(simulation, variable, input_period) + for variable, input_period, _ in overrides + ) + branch = simulation.get_branch(name) + branch.branch_period = period + if parent_knows_override: + drop_inherited_values(branch) + for variable, input_period, value in overrides: + branch.set_input(variable, input_period, value) + branch.branch_overrides = overrides + return branch + + +def _same_overrides(branch: Simulation, overrides) -> bool: + previous = getattr(branch, "branch_overrides", None) + if previous is None or len(previous) != len(overrides): + return False + return all( + variable == previous_variable + and input_period == previous_period + and np.array_equal(value, previous_value) + for (variable, input_period, value), ( + previous_variable, + previous_period, + previous_value, + ) in zip(overrides, previous) + ) + + +def get_branch_for_period( + simulation: Simulation, name: str, period: Period +) -> Simulation: + """Return ``simulation``'s branch ``name`` for ``period``, created per period. + + For branches that change the tax-benefit system rather than inputs: the + caller swaps the system and deletes the variables it recalculates. + """ + if name == simulation.branch_name: + return simulation + branch = simulation.branches.get(name) + if branch is not None and getattr(branch, "branch_period", None) != period: + del simulation.branches[name] + branch = simulation.get_branch(name) + branch.branch_period = period + return branch diff --git a/policyengine_us/variables/gov/hhs/medicaid/medicaid_enrolled_for_ssi_state_supplement.py b/policyengine_us/variables/gov/hhs/medicaid/medicaid_enrolled_for_ssi_state_supplement.py index e9256eca27c..d02f1625c4e 100644 --- a/policyengine_us/variables/gov/hhs/medicaid/medicaid_enrolled_for_ssi_state_supplement.py +++ b/policyengine_us/variables/gov/hhs/medicaid/medicaid_enrolled_for_ssi_state_supplement.py @@ -35,12 +35,17 @@ def formula(person, period, parameters): # age/blindness/disability independently of SNAP/TANF. Keep all other # Medicaid conditions, take-up and supplied inputs in a private branch. branch_name = f"{simulation.branch_name}_ssi_state_supplement_medicaid_{period}" - branch = simulation.get_branch(branch_name) try: - branch.set_input( - "medicaid_community_engagement_pass_through_eligible", - period.first_month, - np.zeros(person.count, dtype=bool), + branch = get_override_branch( + simulation, + branch_name, + period, + { + "medicaid_community_engagement_pass_through_eligible": ( + period.first_month, + np.zeros(person.count, dtype=bool), + ) + }, ) return branch.calculate("medicaid_enrolled", period) finally: diff --git a/policyengine_us/variables/gov/irs/credits/ctc/refundable/ctc_limiting_tax_liability.py b/policyengine_us/variables/gov/irs/credits/ctc/refundable/ctc_limiting_tax_liability.py index 796ce992c1d..afd5737e9a4 100644 --- a/policyengine_us/variables/gov/irs/credits/ctc/refundable/ctc_limiting_tax_liability.py +++ b/policyengine_us/variables/gov/irs/credits/ctc/refundable/ctc_limiting_tax_liability.py @@ -6,28 +6,26 @@ class ctc_limiting_tax_liability(Variable): entity = TaxUnit label = "CTC-limiting tax liability" unit = USD - documentation = "The tax liability used to determine the maximum amount of the non-refundable CTC. Excludes SALT from all calculations (this is an inaccuracy required to avoid circular dependencies)." + documentation = ( + "The tax liability that limits the non-refundable Child Tax Credit: " + "income tax before credits (regular tax plus alternative minimum tax), " + "less the other non-refundable credits." + ) definition_period = YEAR + reference = ( + # 26 U.S.C. 26(a): credits in this subpart are limited to regular tax + # liability (26(b)(1): the tax imposed by chapter 1) plus the tax + # imposed by section 55(a). + "https://www.law.cornell.edu/uscode/text/26/26#a", + # 2025 Schedule 8812 instructions, Credit Limit Worksheet A: line 1 is + # the amount from Form 1040 line 18. + "https://www.irs.gov/instructions/i1040s8", + ) def formula(tax_unit, period, parameters): - simulation = tax_unit.simulation - no_salt_branch = simulation.get_branch("no_salt") - no_salt_branch.set_input("salt_deduction", period, np.zeros(tax_unit.count)) - # Propagate the parent's itemization determination so the - # no_salt branch doesn't re-enter - # `tax_unit_itemizes` -> `tax_liability_if_itemizing` -> - # `income_tax` -> `refundable_ctc`, which forms a cycle - # (issue #8059). The parent's value has already been computed - # by the time we get here: either set as input on the - # itemizing / not_itemizing branch, or computed and cached on - # the top-level sim before `refundable_ctc` was reached (the - # `income_tax_before_credits` branch of - # `income_tax_before_refundable_credits` runs first). - itemizes = tax_unit("tax_unit_itemizes", period) - no_salt_branch.set_input("tax_unit_itemizes", period, itemizes) - tax_liability_before_credits = no_salt_branch.calculate( - "income_tax_before_credits", period - ) + # The tax on taxable income reflects every itemized deduction, + # including state and local taxes. + tax_liability_before_credits = tax_unit("income_tax_before_credits", period) non_refundable_credits = parameters(period).gov.irs.credits.non_refundable non_refundable_credits_ex_ctc = [ x for x in non_refundable_credits if x != "non_refundable_ctc" diff --git a/policyengine_us/variables/gov/irs/income/taxable_income/deductions/tax_liability_if_itemizing.py b/policyengine_us/variables/gov/irs/income/taxable_income/deductions/tax_liability_if_itemizing.py index 775823dd9c3..bbfb5c097ba 100644 --- a/policyengine_us/variables/gov/irs/income/taxable_income/deductions/tax_liability_if_itemizing.py +++ b/policyengine_us/variables/gov/irs/income/taxable_income/deductions/tax_liability_if_itemizing.py @@ -10,8 +10,10 @@ class tax_liability_if_itemizing(Variable): def formula(tax_unit, period, parameters): simulation = tax_unit.simulation - itemized_branch = simulation.get_branch("itemizing") - itemized_branch.set_input( - "tax_unit_itemizes", period, np.ones((tax_unit.count,), dtype=bool) + itemized_branch = get_override_branch( + simulation, + "itemizing", + period, + {"tax_unit_itemizes": np.ones((tax_unit.count,), dtype=bool)}, ) return itemized_branch.calculate("income_tax", period) diff --git a/policyengine_us/variables/gov/irs/income/taxable_income/deductions/tax_liability_if_not_itemizing.py b/policyengine_us/variables/gov/irs/income/taxable_income/deductions/tax_liability_if_not_itemizing.py index eca0caec429..07279c5639b 100644 --- a/policyengine_us/variables/gov/irs/income/taxable_income/deductions/tax_liability_if_not_itemizing.py +++ b/policyengine_us/variables/gov/irs/income/taxable_income/deductions/tax_liability_if_not_itemizing.py @@ -11,10 +11,10 @@ class tax_liability_if_not_itemizing(Variable): def formula(tax_unit, period, parameters): simulation = tax_unit.simulation - non_itemized_branch = simulation.get_branch("not_itemizing") - non_itemized_branch.set_input( - "tax_unit_itemizes", + non_itemized_branch = get_override_branch( + simulation, + "not_itemizing", period, - np.zeros((tax_unit.count,), dtype=bool), + {"tax_unit_itemizes": np.zeros((tax_unit.count,), dtype=bool)}, ) return non_itemized_branch.calculate("income_tax", period) diff --git a/policyengine_us/variables/gov/states/al/tax/income/deductions/federal_income_tax/al_federal_income_tax_deduction.py b/policyengine_us/variables/gov/states/al/tax/income/deductions/federal_income_tax/al_federal_income_tax_deduction.py index e8e93ffa639..46a0b93b1c8 100644 --- a/policyengine_us/variables/gov/states/al/tax/income/deductions/federal_income_tax/al_federal_income_tax_deduction.py +++ b/policyengine_us/variables/gov/states/al/tax/income/deductions/federal_income_tax/al_federal_income_tax_deduction.py @@ -52,7 +52,7 @@ def formula(tax_unit, period, parameters): actual_non_refundable_cdcc = tax_unit("cdcc", period) * (not cdcc_is_refundable) simulation = tax_unit.simulation - branch = simulation.get_branch("al_2020_irc") + branch = get_branch_for_period(simulation, "al_2020_irc", period) branch.tax_benefit_system = get_2020_irc_tbs(simulation.tax_benefit_system) for variable in branch.tax_benefit_system.variables: if any(key in variable for key in ("ctc", "cdcc", "eitc")): diff --git a/policyengine_us/variables/gov/states/de/tax/income/credits/eitc/refundability_calculation/de_income_tax_if_claiming_non_refundable_eitc.py b/policyengine_us/variables/gov/states/de/tax/income/credits/eitc/refundability_calculation/de_income_tax_if_claiming_non_refundable_eitc.py index 94782f0ddac..cccd12d0820 100644 --- a/policyengine_us/variables/gov/states/de/tax/income/credits/eitc/refundability_calculation/de_income_tax_if_claiming_non_refundable_eitc.py +++ b/policyengine_us/variables/gov/states/de/tax/income/credits/eitc/refundability_calculation/de_income_tax_if_claiming_non_refundable_eitc.py @@ -11,10 +11,10 @@ class de_income_tax_if_claiming_non_refundable_eitc(Variable): def formula(tax_unit, period, parameters): simulation = tax_unit.simulation - non_refundable_branch = simulation.get_branch("de_non_refundable_eitc") - non_refundable_branch.set_input( - "de_claims_refundable_eitc", + non_refundable_branch = get_override_branch( + simulation, + "de_non_refundable_eitc", period, - np.zeros((tax_unit.count,), dtype=bool), + {"de_claims_refundable_eitc": np.zeros((tax_unit.count,), dtype=bool)}, ) return non_refundable_branch.calculate("de_income_tax", period) diff --git a/policyengine_us/variables/gov/states/de/tax/income/credits/eitc/refundability_calculation/de_income_tax_if_claiming_refundable_eitc.py b/policyengine_us/variables/gov/states/de/tax/income/credits/eitc/refundability_calculation/de_income_tax_if_claiming_refundable_eitc.py index cec2dfa4451..df3161e0ded 100644 --- a/policyengine_us/variables/gov/states/de/tax/income/credits/eitc/refundability_calculation/de_income_tax_if_claiming_refundable_eitc.py +++ b/policyengine_us/variables/gov/states/de/tax/income/credits/eitc/refundability_calculation/de_income_tax_if_claiming_refundable_eitc.py @@ -11,10 +11,10 @@ class de_income_tax_if_claiming_refundable_eitc(Variable): def formula(tax_unit, period, parameters): simulation = tax_unit.simulation - refundable_branch = simulation.get_branch("de_refundable_eitc") - refundable_branch.set_input( - "de_claims_refundable_eitc", + refundable_branch = get_override_branch( + simulation, + "de_refundable_eitc", period, - np.ones((tax_unit.count,), dtype=bool), + {"de_claims_refundable_eitc": np.ones((tax_unit.count,), dtype=bool)}, ) return refundable_branch.calculate("de_income_tax", period) diff --git a/policyengine_us/variables/gov/states/id/tax/income/id_income_tax_if_receiving_aged_or_disabled_credit.py b/policyengine_us/variables/gov/states/id/tax/income/id_income_tax_if_receiving_aged_or_disabled_credit.py index c024963bf58..cc8c3b9f9d4 100644 --- a/policyengine_us/variables/gov/states/id/tax/income/id_income_tax_if_receiving_aged_or_disabled_credit.py +++ b/policyengine_us/variables/gov/states/id/tax/income/id_income_tax_if_receiving_aged_or_disabled_credit.py @@ -13,10 +13,14 @@ class id_income_tax_if_receiving_aged_or_disabled_credit(Variable): def formula(tax_unit, period, parameters): simulation = tax_unit.simulation - branch = simulation.get_branch("id_receives_aged_or_disabled_credit_branch") - branch.set_input( - "id_receives_aged_or_disabled_credit", + branch = get_override_branch( + simulation, + "id_receives_aged_or_disabled_credit_branch", period, - np.ones((tax_unit.count,), dtype=bool), + { + "id_receives_aged_or_disabled_credit": np.ones( + (tax_unit.count,), dtype=bool + ) + }, ) return branch.calculate("id_income_tax", period) diff --git a/policyengine_us/variables/gov/states/id/tax/income/id_income_tax_if_receiving_aged_or_disabled_deduction.py b/policyengine_us/variables/gov/states/id/tax/income/id_income_tax_if_receiving_aged_or_disabled_deduction.py index 259be96edbe..8f55c47de8a 100644 --- a/policyengine_us/variables/gov/states/id/tax/income/id_income_tax_if_receiving_aged_or_disabled_deduction.py +++ b/policyengine_us/variables/gov/states/id/tax/income/id_income_tax_if_receiving_aged_or_disabled_deduction.py @@ -13,10 +13,14 @@ class id_income_tax_if_receiving_aged_or_disabled_deduction(Variable): def formula(tax_unit, period, parameters): simulation = tax_unit.simulation - branch = simulation.get_branch("id_receives_aged_or_disabled_deduction_branch") - branch.set_input( - "id_receives_aged_or_disabled_credit", + branch = get_override_branch( + simulation, + "id_receives_aged_or_disabled_deduction_branch", period, - np.zeros((tax_unit.count,), dtype=bool), + { + "id_receives_aged_or_disabled_credit": np.zeros( + (tax_unit.count,), dtype=bool + ) + }, ) return branch.calculate("id_income_tax", period) diff --git a/policyengine_us/variables/gov/states/mo/dss/tanf/assistance_unit/non_parent_caretaker/mo_tanf_if_non_parent_caretaker_excluded.py b/policyengine_us/variables/gov/states/mo/dss/tanf/assistance_unit/non_parent_caretaker/mo_tanf_if_non_parent_caretaker_excluded.py index c1e0900e7c5..a2cfb749a93 100644 --- a/policyengine_us/variables/gov/states/mo/dss/tanf/assistance_unit/non_parent_caretaker/mo_tanf_if_non_parent_caretaker_excluded.py +++ b/policyengine_us/variables/gov/states/mo/dss/tanf/assistance_unit/non_parent_caretaker/mo_tanf_if_non_parent_caretaker_excluded.py @@ -19,12 +19,16 @@ def formula(spm_unit, period, parameters): # mo_tanf_if_non_parent_caretaker_included. simulation = spm_unit.simulation branch_name = f"{simulation.branch_name}_mo_tanf_npcr_excluded_{period}" - branch = simulation.get_branch(branch_name) try: - branch.set_input( - "mo_tanf_non_parent_caretaker_included", + branch = get_override_branch( + simulation, + branch_name, period, - np.zeros(spm_unit.count, dtype=bool), + { + "mo_tanf_non_parent_caretaker_included": np.zeros( + spm_unit.count, dtype=bool + ) + }, ) return branch.calculate("mo_tanf", period) finally: diff --git a/policyengine_us/variables/gov/states/mo/dss/tanf/assistance_unit/non_parent_caretaker/mo_tanf_if_non_parent_caretaker_included.py b/policyengine_us/variables/gov/states/mo/dss/tanf/assistance_unit/non_parent_caretaker/mo_tanf_if_non_parent_caretaker_included.py index 75d99502077..3ffd42c9ac6 100644 --- a/policyengine_us/variables/gov/states/mo/dss/tanf/assistance_unit/non_parent_caretaker/mo_tanf_if_non_parent_caretaker_included.py +++ b/policyengine_us/variables/gov/states/mo/dss/tanf/assistance_unit/non_parent_caretaker/mo_tanf_if_non_parent_caretaker_included.py @@ -20,9 +20,13 @@ def formula(spm_unit, period, parameters): needy = spm_unit("mo_tanf_non_parent_caretaker_needy", period) simulation = spm_unit.simulation branch_name = f"{simulation.branch_name}_mo_tanf_npcr_included_{period}" - branch = simulation.get_branch(branch_name) try: - branch.set_input("mo_tanf_non_parent_caretaker_included", period, needy) + branch = get_override_branch( + simulation, + branch_name, + period, + {"mo_tanf_non_parent_caretaker_included": needy}, + ) return branch.calculate("mo_tanf", period) finally: # A branch clones cached arrays; drop it once the grant is read. diff --git a/policyengine_us/variables/gov/states/ny/tax/income/credits/ctc/ny_ctc_pre_2024.py b/policyengine_us/variables/gov/states/ny/tax/income/credits/ctc/ny_ctc_pre_2024.py index 0688b6389bb..c331a3336b8 100644 --- a/policyengine_us/variables/gov/states/ny/tax/income/credits/ctc/ny_ctc_pre_2024.py +++ b/policyengine_us/variables/gov/states/ny/tax/income/credits/ctc/ny_ctc_pre_2024.py @@ -29,7 +29,7 @@ def formula(tax_unit, period, parameters): # Initialize pre-TCJA CTC branch with cached pinned parameters # (one clone per process; see tools/pinned_tbs.py, issue #8114). simulation = tax_unit.simulation - pre_tcja_ctc = simulation.get_branch("pre_tcja_ctc") + pre_tcja_ctc = get_branch_for_period(simulation, "pre_tcja_ctc", period) pre_tcja_ctc.tax_benefit_system = get_pre_tcja_ctc_tbs( simulation.tax_benefit_system ) diff --git a/policyengine_us/variables/gov/states/ny/tax/income/credits/ctc/ny_ctc_pre_2024_eligible.py b/policyengine_us/variables/gov/states/ny/tax/income/credits/ctc/ny_ctc_pre_2024_eligible.py index 3364c58fd18..c7c3b398e65 100644 --- a/policyengine_us/variables/gov/states/ny/tax/income/credits/ctc/ny_ctc_pre_2024_eligible.py +++ b/policyengine_us/variables/gov/states/ny/tax/income/credits/ctc/ny_ctc_pre_2024_eligible.py @@ -28,7 +28,7 @@ def formula(tax_unit, period, parameters): # Initialize pre-TCJA CTC branch for eligibility check with # cached pinned parameters (see tools/pinned_tbs.py, issue #8114). simulation = tax_unit.simulation - pre_tcja_ctc = simulation.get_branch("pre_tcja_ctc") + pre_tcja_ctc = get_branch_for_period(simulation, "pre_tcja_ctc", period) pre_tcja_ctc.tax_benefit_system = get_pre_tcja_ctc_tbs( simulation.tax_benefit_system ) diff --git a/policyengine_us/variables/gov/states/ny/tax/income/credits/ny_eitc.py b/policyengine_us/variables/gov/states/ny/tax/income/credits/ny_eitc.py index a45b7d86eb2..20c38b23cdb 100644 --- a/policyengine_us/variables/gov/states/ny/tax/income/credits/ny_eitc.py +++ b/policyengine_us/variables/gov/states/ny/tax/income/credits/ny_eitc.py @@ -23,7 +23,7 @@ def formula(tax_unit, period, parameters): # does not conform. Recompute the federal EITC with # pre-ARPA (2020) parameter values. simulation = tax_unit.simulation - branch = simulation.get_branch("ny_pre_arpa_eitc") + branch = get_branch_for_period(simulation, "ny_pre_arpa_eitc", period) branch.tax_benefit_system = get_pre_arpa_eitc_tbs( simulation.tax_benefit_system ) diff --git a/policyengine_us/variables/gov/states/va/tax/income/credits/eitc/refundability_calculation/va_income_tax_if_claiming_non_refundable_eitc.py b/policyengine_us/variables/gov/states/va/tax/income/credits/eitc/refundability_calculation/va_income_tax_if_claiming_non_refundable_eitc.py index cef69a3c26c..2de32bddeb0 100644 --- a/policyengine_us/variables/gov/states/va/tax/income/credits/eitc/refundability_calculation/va_income_tax_if_claiming_non_refundable_eitc.py +++ b/policyengine_us/variables/gov/states/va/tax/income/credits/eitc/refundability_calculation/va_income_tax_if_claiming_non_refundable_eitc.py @@ -11,10 +11,10 @@ class va_income_tax_if_claiming_non_refundable_eitc(Variable): def formula(tax_unit, period, parameters): simulation = tax_unit.simulation - non_refundable_branch = simulation.get_branch("va_non_refundable_eitc") - non_refundable_branch.set_input( - "va_claims_refundable_eitc", + non_refundable_branch = get_override_branch( + simulation, + "va_non_refundable_eitc", period, - np.zeros((tax_unit.count,), dtype=bool), + {"va_claims_refundable_eitc": np.zeros((tax_unit.count,), dtype=bool)}, ) return non_refundable_branch.calculate("va_income_tax", period) diff --git a/policyengine_us/variables/gov/states/va/tax/income/credits/eitc/refundability_calculation/va_income_tax_if_claiming_refundable_eitc.py b/policyengine_us/variables/gov/states/va/tax/income/credits/eitc/refundability_calculation/va_income_tax_if_claiming_refundable_eitc.py index b510db847a2..ff7ea5ee356 100644 --- a/policyengine_us/variables/gov/states/va/tax/income/credits/eitc/refundability_calculation/va_income_tax_if_claiming_refundable_eitc.py +++ b/policyengine_us/variables/gov/states/va/tax/income/credits/eitc/refundability_calculation/va_income_tax_if_claiming_refundable_eitc.py @@ -11,10 +11,10 @@ class va_income_tax_if_claiming_refundable_eitc(Variable): def formula(tax_unit, period, parameters): simulation = tax_unit.simulation - refundable_branch = simulation.get_branch("va_refundable_eitc") - refundable_branch.set_input( - "va_claims_refundable_eitc", + refundable_branch = get_override_branch( + simulation, + "va_refundable_eitc", period, - np.ones((tax_unit.count,), dtype=bool), + {"va_claims_refundable_eitc": np.ones((tax_unit.count,), dtype=bool)}, ) return refundable_branch.calculate("va_income_tax", period) From acb44bd5360a32eb80b8b6225e2ae105027d20ac Mon Sep 17 00:00:00 2001 From: Max Ghenis Date: Mon, 5 Oct 2026 14:02:24 -0400 Subject: [PATCH 2/7] Make comparison branches calculate under their overrides (get_override_branch) Fold #9741's get_override_branch into the period_branch module #9738 added, so there is one branch helper: get_override_branch keeps one branch per name and period, creates it again for another period or other override values, and drops the arrays a branch copied (keeping inputs) when the parent already has a value for an overridden input. get_branch_for_period becomes the same call with no inputs, for the Alabama 2020-IRC and New York pinned-parameter branches. Itemizing/not itemizing, Delaware and Virginia EITC refundability, Idaho aged or disabled credit/deduction, Missouri TANF caretaker and the Medicaid SSI state-supplement branch now set their inputs through get_override_branch. Co-Authored-By: Claude Opus 5.5 --- policyengine_us/tools/general.py | 5 +- policyengine_us/tools/period_branch.py | 161 ++++++++++++++++-- ...icaid_enrolled_for_ssi_state_supplement.py | 15 +- .../deductions/tax_liability_if_itemizing.py | 8 +- .../tax_liability_if_not_itemizing.py | 8 +- ...ome_tax_if_claiming_non_refundable_eitc.py | 10 +- ..._income_tax_if_claiming_refundable_eitc.py | 10 +- ...ax_if_receiving_aged_or_disabled_credit.py | 14 +- ...if_receiving_aged_or_disabled_deduction.py | 14 +- ...o_tanf_if_non_parent_caretaker_excluded.py | 12 +- ...o_tanf_if_non_parent_caretaker_included.py | 8 +- ...ome_tax_if_claiming_non_refundable_eitc.py | 10 +- ..._income_tax_if_claiming_refundable_eitc.py | 10 +- 13 files changed, 214 insertions(+), 71 deletions(-) diff --git a/policyengine_us/tools/general.py b/policyengine_us/tools/general.py index 4bf2cf005d8..ccbaa8ab152 100644 --- a/policyengine_us/tools/general.py +++ b/policyengine_us/tools/general.py @@ -1,7 +1,10 @@ from policyengine_core.model_api import * from policyengine_us.entities import * from policyengine_us.tools.branched_simulation import BranchedSimulation -from policyengine_us.tools.period_branch import get_branch_for_period +from policyengine_us.tools.period_branch import ( + get_branch_for_period, + get_override_branch, +) from pathlib import Path import pandas as pd from policyengine_us.typing import Formula diff --git a/policyengine_us/tools/period_branch.py b/policyengine_us/tools/period_branch.py index fa516601f3b..439e35028aa 100644 --- a/policyengine_us/tools/period_branch.py +++ b/policyengine_us/tools/period_branch.py @@ -1,27 +1,156 @@ +"""Formula branches that calculate one period, under overridden inputs. + +A policyengine-core branch starts from every array its parent has cached when +the branch is created and never sees the parent's later calculations, and +``set_input`` on the branch stores the new value without clearing anything +that was calculated from the old one. A branch therefore answers any variable +its parent had already calculated with the parent's value, whatever the +branch's inputs say. A branch kept from another period also answers this +period from a copy taken before any of this period was calculated, unlike the +branch a simulation calculating only this period would create, so a later +year came out differently when an earlier year had been calculated first. + +``get_override_branch`` makes the override reach every value the branch +calculates: + +- A branch serves one period. Asked for another period, it is created again + from the parent as the parent stands then, as a simulation calculating only + that period would create it. +- When the branch is created, it keeps the parent's cache only if the parent + has no value yet for any overridden variable and period. A cached value + cannot have been calculated from a value that did not exist, so the + parent's cache is then safe to share; this is the usual case, where a + formula branches while its parent is still calculating the variable the + branch overrides. Otherwise the branch drops every array it copied except + inputs, and calculates the rest itself. +- A branch reused within its period with different inputs is created again. + +``get_branch_for_period`` is the same with no inputs, for branches that +change the tax-benefit system rather than inputs: the caller swaps the system +and deletes the variables it recalculates. +""" + +from typing import Dict, Tuple, Union + +import numpy as np from policyengine_core.periods import Period +from policyengine_core.periods import period as to_period from policyengine_core.simulations import Simulation +Override = Union[np.ndarray, Tuple[Period, np.ndarray]] + + +def _overrides(period: Period, inputs: Dict[str, Override]): + for variable, value in inputs.items(): + if isinstance(value, tuple): + input_period, value = value + yield variable, to_period(input_period), np.asarray(value) + else: + yield variable, period, np.asarray(value) + + +def _is_known(simulation: Simulation, variable: str, period: Period) -> bool: + holder = simulation.get_holder(variable) + if holder.variable.is_neutralized: + return False + return holder.get_array(period, simulation.branch_name) is not None -def get_branch_for_period(simulation: Simulation, name: str, period: Period): - """Return ``simulation``'s branch ``name`` for calculating ``period``. - A branch copies its parent's cached arrays when it is created and never - sees the parent's later calculations, and an input set on it does not - clear values it already copied. What a branch computes for a period - therefore depends on what its parent had calculated when the branch was - created. A branch kept from another period answers this period from a copy - taken before any of this period was calculated, unlike the branch a - simulation calculating only this period would create, so a later year came - out differently when an earlier year had been calculated first. +def drop_inherited_values(branch: Simulation) -> None: + """Delete every array ``branch`` holds except inputs. - Within a period the branch is shared as before. A branch left from another - period is dropped, and the branch is created again from the parent as it - stands now. + Inputs are the variables the simulation was built with and each value set + through ``set_input`` on a branch this one reads. """ + input_variables = set(branch.input_variables) + user_input_keys = getattr(branch, "_user_input_keys", set()) + visible_branches = set(branch._get_visible_branch_names()) + for population in branch.populations.values(): + for name, holder in population._holders.items(): + if name in input_variables: + continue + for branch_name, known_period in holder.get_known_branch_periods(): + if ( + branch_name in visible_branches + and (name, branch_name, known_period) in user_input_keys + ): + continue + # Exact key: ``Holder.delete_arrays`` would also delete any + # input stored at a sub-period of ``known_period``. + key = f"{branch_name}:{known_period}" + holder._memory_storage._arrays.pop(key, None) + if holder._disk_storage is not None: + holder._disk_storage._files.pop( + f"{branch_name}_{known_period}", None + ) + branch._fast_cache = {} + + +def get_override_branch( + simulation: Simulation, + name: str, + period: Period, + inputs: Dict[str, Override], +) -> Simulation: + """Return ``simulation``'s branch ``name`` for ``period``, with ``inputs`` set. + + ``inputs`` maps each overridden variable to its value for ``period``, or + to a ``(period, value)`` pair for another period (e.g. a month). + """ + period = to_period(period) + overrides = list(_overrides(period, inputs)) + if name == simulation.branch_name: + # Already inside this branch (core's get_branch returns the + # simulation itself): only the inputs need setting. + for variable, input_period, value in overrides: + simulation.set_input(variable, input_period, value) + return simulation branch = simulation.branches.get(name) - if branch is not None and getattr(branch, "branch_period", period) != period: + if branch is not None and ( + getattr(branch, "branch_period", None) != period + or not _same_overrides(branch, overrides) + ): del simulation.branches[name] - branch = simulation.get_branch(name) - if branch is not simulation: + branch = None + if branch is None: + parent_knows_override = any( + _is_known(simulation, variable, input_period) + for variable, input_period, _ in overrides + ) + branch = simulation.get_branch(name) branch.branch_period = period + if parent_knows_override: + drop_inherited_values(branch) + for variable, input_period, value in overrides: + branch.set_input(variable, input_period, value) + branch.branch_overrides = overrides return branch + + +def _same_overrides(branch: Simulation, overrides) -> bool: + previous = getattr(branch, "branch_overrides", None) + if previous is None or len(previous) != len(overrides): + return False + return all( + variable == previous_variable + and input_period == previous_period + and np.array_equal(value, previous_value) + for (variable, input_period, value), ( + previous_variable, + previous_period, + previous_value, + ) in zip(overrides, previous) + ) + + +def get_branch_for_period( + simulation: Simulation, name: str, period: Period +) -> Simulation: + """Return ``simulation``'s branch ``name`` for ``period``, created per period. + + For branches that change the tax-benefit system rather than inputs: the + caller swaps the system and deletes the variables it recalculates. Within + a period the branch is shared; a branch left from another period is + dropped and created again from the parent as it stands now. + """ + return get_override_branch(simulation, name, period, {}) diff --git a/policyengine_us/variables/gov/hhs/medicaid/medicaid_enrolled_for_ssi_state_supplement.py b/policyengine_us/variables/gov/hhs/medicaid/medicaid_enrolled_for_ssi_state_supplement.py index e9256eca27c..d02f1625c4e 100644 --- a/policyengine_us/variables/gov/hhs/medicaid/medicaid_enrolled_for_ssi_state_supplement.py +++ b/policyengine_us/variables/gov/hhs/medicaid/medicaid_enrolled_for_ssi_state_supplement.py @@ -35,12 +35,17 @@ def formula(person, period, parameters): # age/blindness/disability independently of SNAP/TANF. Keep all other # Medicaid conditions, take-up and supplied inputs in a private branch. branch_name = f"{simulation.branch_name}_ssi_state_supplement_medicaid_{period}" - branch = simulation.get_branch(branch_name) try: - branch.set_input( - "medicaid_community_engagement_pass_through_eligible", - period.first_month, - np.zeros(person.count, dtype=bool), + branch = get_override_branch( + simulation, + branch_name, + period, + { + "medicaid_community_engagement_pass_through_eligible": ( + period.first_month, + np.zeros(person.count, dtype=bool), + ) + }, ) return branch.calculate("medicaid_enrolled", period) finally: diff --git a/policyengine_us/variables/gov/irs/income/taxable_income/deductions/tax_liability_if_itemizing.py b/policyengine_us/variables/gov/irs/income/taxable_income/deductions/tax_liability_if_itemizing.py index db3bc7182f3..bbfb5c097ba 100644 --- a/policyengine_us/variables/gov/irs/income/taxable_income/deductions/tax_liability_if_itemizing.py +++ b/policyengine_us/variables/gov/irs/income/taxable_income/deductions/tax_liability_if_itemizing.py @@ -10,8 +10,10 @@ class tax_liability_if_itemizing(Variable): def formula(tax_unit, period, parameters): simulation = tax_unit.simulation - itemized_branch = get_branch_for_period(simulation, "itemizing", period) - itemized_branch.set_input( - "tax_unit_itemizes", period, np.ones((tax_unit.count,), dtype=bool) + itemized_branch = get_override_branch( + simulation, + "itemizing", + period, + {"tax_unit_itemizes": np.ones((tax_unit.count,), dtype=bool)}, ) return itemized_branch.calculate("income_tax", period) diff --git a/policyengine_us/variables/gov/irs/income/taxable_income/deductions/tax_liability_if_not_itemizing.py b/policyengine_us/variables/gov/irs/income/taxable_income/deductions/tax_liability_if_not_itemizing.py index cc3738fd771..07279c5639b 100644 --- a/policyengine_us/variables/gov/irs/income/taxable_income/deductions/tax_liability_if_not_itemizing.py +++ b/policyengine_us/variables/gov/irs/income/taxable_income/deductions/tax_liability_if_not_itemizing.py @@ -11,10 +11,10 @@ class tax_liability_if_not_itemizing(Variable): def formula(tax_unit, period, parameters): simulation = tax_unit.simulation - non_itemized_branch = get_branch_for_period(simulation, "not_itemizing", period) - non_itemized_branch.set_input( - "tax_unit_itemizes", + non_itemized_branch = get_override_branch( + simulation, + "not_itemizing", period, - np.zeros((tax_unit.count,), dtype=bool), + {"tax_unit_itemizes": np.zeros((tax_unit.count,), dtype=bool)}, ) return non_itemized_branch.calculate("income_tax", period) diff --git a/policyengine_us/variables/gov/states/de/tax/income/credits/eitc/refundability_calculation/de_income_tax_if_claiming_non_refundable_eitc.py b/policyengine_us/variables/gov/states/de/tax/income/credits/eitc/refundability_calculation/de_income_tax_if_claiming_non_refundable_eitc.py index 36dccc9258e..cccd12d0820 100644 --- a/policyengine_us/variables/gov/states/de/tax/income/credits/eitc/refundability_calculation/de_income_tax_if_claiming_non_refundable_eitc.py +++ b/policyengine_us/variables/gov/states/de/tax/income/credits/eitc/refundability_calculation/de_income_tax_if_claiming_non_refundable_eitc.py @@ -11,12 +11,10 @@ class de_income_tax_if_claiming_non_refundable_eitc(Variable): def formula(tax_unit, period, parameters): simulation = tax_unit.simulation - non_refundable_branch = get_branch_for_period( - simulation, "de_non_refundable_eitc", period - ) - non_refundable_branch.set_input( - "de_claims_refundable_eitc", + non_refundable_branch = get_override_branch( + simulation, + "de_non_refundable_eitc", period, - np.zeros((tax_unit.count,), dtype=bool), + {"de_claims_refundable_eitc": np.zeros((tax_unit.count,), dtype=bool)}, ) return non_refundable_branch.calculate("de_income_tax", period) diff --git a/policyengine_us/variables/gov/states/de/tax/income/credits/eitc/refundability_calculation/de_income_tax_if_claiming_refundable_eitc.py b/policyengine_us/variables/gov/states/de/tax/income/credits/eitc/refundability_calculation/de_income_tax_if_claiming_refundable_eitc.py index ad4437284ea..df3161e0ded 100644 --- a/policyengine_us/variables/gov/states/de/tax/income/credits/eitc/refundability_calculation/de_income_tax_if_claiming_refundable_eitc.py +++ b/policyengine_us/variables/gov/states/de/tax/income/credits/eitc/refundability_calculation/de_income_tax_if_claiming_refundable_eitc.py @@ -11,12 +11,10 @@ class de_income_tax_if_claiming_refundable_eitc(Variable): def formula(tax_unit, period, parameters): simulation = tax_unit.simulation - refundable_branch = get_branch_for_period( - simulation, "de_refundable_eitc", period - ) - refundable_branch.set_input( - "de_claims_refundable_eitc", + refundable_branch = get_override_branch( + simulation, + "de_refundable_eitc", period, - np.ones((tax_unit.count,), dtype=bool), + {"de_claims_refundable_eitc": np.ones((tax_unit.count,), dtype=bool)}, ) return refundable_branch.calculate("de_income_tax", period) diff --git a/policyengine_us/variables/gov/states/id/tax/income/id_income_tax_if_receiving_aged_or_disabled_credit.py b/policyengine_us/variables/gov/states/id/tax/income/id_income_tax_if_receiving_aged_or_disabled_credit.py index a838a5d8a5d..cc8c3b9f9d4 100644 --- a/policyengine_us/variables/gov/states/id/tax/income/id_income_tax_if_receiving_aged_or_disabled_credit.py +++ b/policyengine_us/variables/gov/states/id/tax/income/id_income_tax_if_receiving_aged_or_disabled_credit.py @@ -13,12 +13,14 @@ class id_income_tax_if_receiving_aged_or_disabled_credit(Variable): def formula(tax_unit, period, parameters): simulation = tax_unit.simulation - branch = get_branch_for_period( - simulation, "id_receives_aged_or_disabled_credit_branch", period - ) - branch.set_input( - "id_receives_aged_or_disabled_credit", + branch = get_override_branch( + simulation, + "id_receives_aged_or_disabled_credit_branch", period, - np.ones((tax_unit.count,), dtype=bool), + { + "id_receives_aged_or_disabled_credit": np.ones( + (tax_unit.count,), dtype=bool + ) + }, ) return branch.calculate("id_income_tax", period) diff --git a/policyengine_us/variables/gov/states/id/tax/income/id_income_tax_if_receiving_aged_or_disabled_deduction.py b/policyengine_us/variables/gov/states/id/tax/income/id_income_tax_if_receiving_aged_or_disabled_deduction.py index e38bf9414bf..8f55c47de8a 100644 --- a/policyengine_us/variables/gov/states/id/tax/income/id_income_tax_if_receiving_aged_or_disabled_deduction.py +++ b/policyengine_us/variables/gov/states/id/tax/income/id_income_tax_if_receiving_aged_or_disabled_deduction.py @@ -13,12 +13,14 @@ class id_income_tax_if_receiving_aged_or_disabled_deduction(Variable): def formula(tax_unit, period, parameters): simulation = tax_unit.simulation - branch = get_branch_for_period( - simulation, "id_receives_aged_or_disabled_deduction_branch", period - ) - branch.set_input( - "id_receives_aged_or_disabled_credit", + branch = get_override_branch( + simulation, + "id_receives_aged_or_disabled_deduction_branch", period, - np.zeros((tax_unit.count,), dtype=bool), + { + "id_receives_aged_or_disabled_credit": np.zeros( + (tax_unit.count,), dtype=bool + ) + }, ) return branch.calculate("id_income_tax", period) diff --git a/policyengine_us/variables/gov/states/mo/dss/tanf/assistance_unit/non_parent_caretaker/mo_tanf_if_non_parent_caretaker_excluded.py b/policyengine_us/variables/gov/states/mo/dss/tanf/assistance_unit/non_parent_caretaker/mo_tanf_if_non_parent_caretaker_excluded.py index c1e0900e7c5..a2cfb749a93 100644 --- a/policyengine_us/variables/gov/states/mo/dss/tanf/assistance_unit/non_parent_caretaker/mo_tanf_if_non_parent_caretaker_excluded.py +++ b/policyengine_us/variables/gov/states/mo/dss/tanf/assistance_unit/non_parent_caretaker/mo_tanf_if_non_parent_caretaker_excluded.py @@ -19,12 +19,16 @@ def formula(spm_unit, period, parameters): # mo_tanf_if_non_parent_caretaker_included. simulation = spm_unit.simulation branch_name = f"{simulation.branch_name}_mo_tanf_npcr_excluded_{period}" - branch = simulation.get_branch(branch_name) try: - branch.set_input( - "mo_tanf_non_parent_caretaker_included", + branch = get_override_branch( + simulation, + branch_name, period, - np.zeros(spm_unit.count, dtype=bool), + { + "mo_tanf_non_parent_caretaker_included": np.zeros( + spm_unit.count, dtype=bool + ) + }, ) return branch.calculate("mo_tanf", period) finally: diff --git a/policyengine_us/variables/gov/states/mo/dss/tanf/assistance_unit/non_parent_caretaker/mo_tanf_if_non_parent_caretaker_included.py b/policyengine_us/variables/gov/states/mo/dss/tanf/assistance_unit/non_parent_caretaker/mo_tanf_if_non_parent_caretaker_included.py index 75d99502077..3ffd42c9ac6 100644 --- a/policyengine_us/variables/gov/states/mo/dss/tanf/assistance_unit/non_parent_caretaker/mo_tanf_if_non_parent_caretaker_included.py +++ b/policyengine_us/variables/gov/states/mo/dss/tanf/assistance_unit/non_parent_caretaker/mo_tanf_if_non_parent_caretaker_included.py @@ -20,9 +20,13 @@ def formula(spm_unit, period, parameters): needy = spm_unit("mo_tanf_non_parent_caretaker_needy", period) simulation = spm_unit.simulation branch_name = f"{simulation.branch_name}_mo_tanf_npcr_included_{period}" - branch = simulation.get_branch(branch_name) try: - branch.set_input("mo_tanf_non_parent_caretaker_included", period, needy) + branch = get_override_branch( + simulation, + branch_name, + period, + {"mo_tanf_non_parent_caretaker_included": needy}, + ) return branch.calculate("mo_tanf", period) finally: # A branch clones cached arrays; drop it once the grant is read. diff --git a/policyengine_us/variables/gov/states/va/tax/income/credits/eitc/refundability_calculation/va_income_tax_if_claiming_non_refundable_eitc.py b/policyengine_us/variables/gov/states/va/tax/income/credits/eitc/refundability_calculation/va_income_tax_if_claiming_non_refundable_eitc.py index b53d63491c1..2de32bddeb0 100644 --- a/policyengine_us/variables/gov/states/va/tax/income/credits/eitc/refundability_calculation/va_income_tax_if_claiming_non_refundable_eitc.py +++ b/policyengine_us/variables/gov/states/va/tax/income/credits/eitc/refundability_calculation/va_income_tax_if_claiming_non_refundable_eitc.py @@ -11,12 +11,10 @@ class va_income_tax_if_claiming_non_refundable_eitc(Variable): def formula(tax_unit, period, parameters): simulation = tax_unit.simulation - non_refundable_branch = get_branch_for_period( - simulation, "va_non_refundable_eitc", period - ) - non_refundable_branch.set_input( - "va_claims_refundable_eitc", + non_refundable_branch = get_override_branch( + simulation, + "va_non_refundable_eitc", period, - np.zeros((tax_unit.count,), dtype=bool), + {"va_claims_refundable_eitc": np.zeros((tax_unit.count,), dtype=bool)}, ) return non_refundable_branch.calculate("va_income_tax", period) diff --git a/policyengine_us/variables/gov/states/va/tax/income/credits/eitc/refundability_calculation/va_income_tax_if_claiming_refundable_eitc.py b/policyengine_us/variables/gov/states/va/tax/income/credits/eitc/refundability_calculation/va_income_tax_if_claiming_refundable_eitc.py index 1b363d70c96..ff7ea5ee356 100644 --- a/policyengine_us/variables/gov/states/va/tax/income/credits/eitc/refundability_calculation/va_income_tax_if_claiming_refundable_eitc.py +++ b/policyengine_us/variables/gov/states/va/tax/income/credits/eitc/refundability_calculation/va_income_tax_if_claiming_refundable_eitc.py @@ -11,12 +11,10 @@ class va_income_tax_if_claiming_refundable_eitc(Variable): def formula(tax_unit, period, parameters): simulation = tax_unit.simulation - refundable_branch = get_branch_for_period( - simulation, "va_refundable_eitc", period - ) - refundable_branch.set_input( - "va_claims_refundable_eitc", + refundable_branch = get_override_branch( + simulation, + "va_refundable_eitc", period, - np.ones((tax_unit.count,), dtype=bool), + {"va_claims_refundable_eitc": np.ones((tax_unit.count,), dtype=bool)}, ) return refundable_branch.calculate("va_income_tax", period) From 12a69e32c6aa06bf1e124f4d2281c48fe42f2f83 Mon Sep 17 00:00:00 2001 From: Max Ghenis Date: Mon, 5 Oct 2026 14:02:53 -0400 Subject: [PATCH 3/7] Start Credit Limit Worksheet A from the actual tax liability, SALT included ctc_tax_liability_after_preceding_credits (Worksheet A line 3, from #9742) read income tax before credits from a no_salt branch, whose override was shadowed or not depending on which variables were calculated first. 26 U.S.C. 26(a) limits the credit by regular tax liability (the chapter 1 tax, on taxable income after itemized deductions) plus the 55(a) tax, and Worksheet A line 1 is Form 1040 line 18. Read income_tax_before_credits directly. No cycle: the SALT deduction counts state income tax through state_withheld_income_tax, withholding estimated from each person's AGI. Adds a hand-worked 2025 itemizer whose limit binds; the multi-year test no longer expects a no_salt branch. Co-Authored-By: Claude Opus 5.5 --- .../tests/core/test_multi_year_simulation.py | 7 +-- ...tax_liability_after_preceding_credits.yaml | 50 +++++++++++++++++++ .../tests/test_ctc_credit_limit_worksheets.py | 2 +- .../tests/test_ctc_itemizing_branch_cycle.py | 7 +-- .../refundable/ctc_limiting_tax_liability.py | 3 +- ...c_tax_liability_after_preceding_credits.py | 35 ++++++------- 6 files changed, 74 insertions(+), 30 deletions(-) diff --git a/policyengine_us/tests/core/test_multi_year_simulation.py b/policyengine_us/tests/core/test_multi_year_simulation.py index af43f9fade7..1df53d21061 100644 --- a/policyengine_us/tests/core/test_multi_year_simulation.py +++ b/policyengine_us/tests/core/test_multi_year_simulation.py @@ -14,7 +14,8 @@ periods. A branch copies its parent's cached arrays when it is created, so a branch created in 2024 answered 2025 from a copy without any of the parent's 2025 values, unlike the branch a 2025-only simulation creates. - ``get_branch_for_period`` creates them again for each period. + ``get_override_branch`` and ``get_branch_for_period`` create them again for + each period. (The CTC limit no longer uses a SALT branch.) Each test compares a simulation that calculates the base year first with a fresh simulation that calculates only the later year. The branch test gives @@ -163,7 +164,7 @@ def test_formula_branches_are_created_again_for_a_later_year(year): simulation = Simulation(situation=_situation(age_every_year=True)) _later_year_values(simulation, BASE_YEAR) - branches = ("itemizing", "not_itemizing", "no_salt") + branches = ("itemizing", "not_itemizing") base_year_branches = {name: simulation.branches[name] for name in branches} _assert_same(_later_year_values(simulation, year), fresh, year) @@ -191,7 +192,7 @@ def test_later_year_matches_single_year_simulation(year): simulation = Simulation(situation=_situation()) _later_year_values(simulation, BASE_YEAR) - BRANCHES = ("itemizing", "not_itemizing", "no_salt") + BRANCHES = ("itemizing", "not_itemizing") base_year_branches = {name: simulation.branches[name] for name in BRANCHES} # Request every month of the base year too, as monthly programs do for # people they cover. diff --git a/policyengine_us/tests/policy/baseline/gov/irs/credits/ctc/refundable/ctc_tax_liability_after_preceding_credits.yaml b/policyengine_us/tests/policy/baseline/gov/irs/credits/ctc/refundable/ctc_tax_liability_after_preceding_credits.yaml index 2304ab2ae24..2af75e6e088 100644 --- a/policyengine_us/tests/policy/baseline/gov/irs/credits/ctc/refundable/ctc_tax_liability_after_preceding_credits.yaml +++ b/policyengine_us/tests/policy/baseline/gov/irs/credits/ctc/refundable/ctc_tax_liability_after_preceding_credits.yaml @@ -50,3 +50,53 @@ cdcc: 500 output: ctc_tax_liability_after_preceding_credits: 3_000 + +- name: An itemizer's line 1 is the tax after the SALT deduction, which limits the non-refundable CTC. + # Line 1 is Form 1040 line 18 (2025 Schedule 8812 instructions, Credit Limit + # Worksheet A; 26 U.S.C. 26(a)), the tax on taxable income after every + # itemized deduction. A married couple in Texas with three children, $90,000 + # of wages, $40,000 of mortgage interest and $15,000 of real estate tax: + # itemized deductions 40,000 + 15,000 (under the $40,000 SALT cap) = 55,000 + # > the 31,500 standard deduction; taxable income 90,000 - 55,000 = 35,000; + # tax = 10% x 23,850 + 12% x 11,150 = 3,723. + # Without SALT, taxable income would be 50,000 and the tax 5,523. + period: 2025 + absolute_error_margin: 0.01 + input: + people: + head: + age: 40 + employment_income: 90_000 + deductible_mortgage_interest: 40_000 + real_estate_taxes: 15_000 + spouse: + age: 40 + child1: + age: 4 + child2: + age: 8 + child3: + age: 12 + tax_units: + tax_unit: + members: [head, spouse, child1, child2, child3] + # SALT is the real estate tax alone. + state_sales_tax: 0 + local_sales_tax: 0 + households: + household: + members: [head, spouse, child1, child2, child3] + state_code: TX + output: + salt_deduction: 15_000 + tax_unit_itemizes: true + income_tax_before_credits: 3_723 + # No line 2 credits. + ctc_tax_liability_after_preceding_credits: 3_723 + # No residential clean energy credit, so line 5 = line 3. + ctc_limiting_tax_liability: 3_723 + # Schedule 8812: CTC 3 x 2,200 = 6,600; line 14 = min(6,600, 3,723); + # line 27 = min(6,600 - 3,723, 3 x 1,700, 15% x (90,000 - 2,500)) = 2,877. + non_refundable_ctc: 3_723 + refundable_ctc: 2_877 + income_tax: -2_877 diff --git a/policyengine_us/tests/test_ctc_credit_limit_worksheets.py b/policyengine_us/tests/test_ctc_credit_limit_worksheets.py index 5285ed1a2e2..f477077746a 100644 --- a/policyengine_us/tests/test_ctc_credit_limit_worksheets.py +++ b/policyengine_us/tests/test_ctc_credit_limit_worksheets.py @@ -14,7 +14,7 @@ other subpart A credit. Households are married or single filers in Texas who take the standard -deduction, so the CTC limit's no-SALT liability equals actual liability. +deduction. """ import numpy as np diff --git a/policyengine_us/tests/test_ctc_itemizing_branch_cycle.py b/policyengine_us/tests/test_ctc_itemizing_branch_cycle.py index 9ba59171672..2198de58e74 100644 --- a/policyengine_us/tests/test_ctc_itemizing_branch_cycle.py +++ b/policyengine_us/tests/test_ctc_itemizing_branch_cycle.py @@ -16,9 +16,10 @@ consumers (e.g. policyengine.py's household-impact integration tests) on `policyengine-core >= 3.24`. -The fix in `ctc_tax_liability_after_preceding_credits.py` propagates the parent's -`tax_unit_itemizes` value to the no_salt child branch so the -`tax_unit_itemizes` formula is never re-entered there. +`ctc_tax_liability_after_preceding_credits` no longer creates the no_salt +branch: it reads the branch's own `income_tax_before_credits`, which the +itemizing branch calculates with `tax_unit_itemizes` set, so the chain cannot +re-enter it. """ import numpy as np diff --git a/policyengine_us/variables/gov/irs/credits/ctc/refundable/ctc_limiting_tax_liability.py b/policyengine_us/variables/gov/irs/credits/ctc/refundable/ctc_limiting_tax_liability.py index 85d6ccccbdb..88692127de9 100644 --- a/policyengine_us/variables/gov/irs/credits/ctc/refundable/ctc_limiting_tax_liability.py +++ b/policyengine_us/variables/gov/irs/credits/ctc/refundable/ctc_limiting_tax_liability.py @@ -11,8 +11,7 @@ class ctc_limiting_tax_liability(Variable): "non-refundable CTC (Schedule 8812 Credit Limit Worksheet A, line 5): " "income tax before credits less the credits that precede the CTC, " "less the residential clean energy credit when Credit Limit Worksheet " - "B applies. Excludes SALT from income tax before credits (this is an " - "inaccuracy required to avoid circular dependencies)." + "B applies." ) definition_period = YEAR reference = ( diff --git a/policyengine_us/variables/gov/irs/credits/ctc/refundable/ctc_tax_liability_after_preceding_credits.py b/policyengine_us/variables/gov/irs/credits/ctc/refundable/ctc_tax_liability_after_preceding_credits.py index d950fd27e22..67091a03289 100644 --- a/policyengine_us/variables/gov/irs/credits/ctc/refundable/ctc_tax_liability_after_preceding_credits.py +++ b/policyengine_us/variables/gov/irs/credits/ctc/refundable/ctc_tax_liability_after_preceding_credits.py @@ -9,35 +9,28 @@ class ctc_tax_liability_after_preceding_credits(Variable): documentation = ( "Income tax before credits less the non-refundable credits that " "precede the Child Tax Credit (Schedule 8812 Credit Limit Worksheet " - "A, line 3). Excludes SALT from income tax before credits (this is an " - "inaccuracy required to avoid circular dependencies)." + "A, line 3). Income tax before credits is the actual liability, on " + "taxable income after every itemized deduction, state and local taxes " + "included." ) definition_period = YEAR reference = ( + # 26 U.S.C. 26(a): credits in this subpart are limited to regular tax + # liability (26(b)(1): the tax imposed by chapter 1) plus the tax + # imposed by section 55(a). "https://www.law.cornell.edu/uscode/text/26/26#a", - # 2025 Instructions for Schedule 8812, Credit Limit Worksheet A. + # 2025 Instructions for Schedule 8812, Credit Limit Worksheet A: line 1 + # is the amount from Form 1040 line 18. "https://www.irs.gov/pub/irs-pdf/i1040s8.pdf#page=4", ) def formula(tax_unit, period, parameters): - simulation = tax_unit.simulation - no_salt_branch = get_branch_for_period(simulation, "no_salt", period) - no_salt_branch.set_input("salt_deduction", period, np.zeros(tax_unit.count)) - # Propagate the parent's itemization determination so the - # no_salt branch doesn't re-enter - # `tax_unit_itemizes` -> `tax_liability_if_itemizing` -> - # `income_tax` -> `refundable_ctc`, which forms a cycle - # (issue #8059). The parent's value has already been computed - # by the time we get here: either set as input on the - # itemizing / not_itemizing branch, or computed and cached on - # the top-level sim before `refundable_ctc` was reached (the - # `income_tax_before_credits` branch of - # `income_tax_before_refundable_credits` runs first). - itemizes = tax_unit("tax_unit_itemizes", period) - no_salt_branch.set_input("tax_unit_itemizes", period, itemizes) - tax_liability_before_credits = no_salt_branch.calculate( - "income_tax_before_credits", period - ) + # Line 1: the tax on taxable income reflects every itemized deduction, + # including state and local taxes. Reading it here forms no cycle: + # the SALT deduction counts state income tax through + # state_withheld_income_tax, withholding estimated from income, which + # reads neither federal tax nor credits. + tax_liability_before_credits = tax_unit("income_tax_before_credits", period) p = parameters(period).gov.irs.credits.ctc_tax_liability_limit # add() returns None for an empty list, which a reform may set. preceding_credits = ( From ff3e8b552c98da05feeb7780c1ef6b8b11da6893 Mon Sep 17 00:00:00 2001 From: Max Ghenis Date: Mon, 5 Oct 2026 14:03:01 -0400 Subject: [PATCH 4/7] Test that branches calculate under their overrides in any order Order-independence (seeded sample and Hypothesis), differential tests of each comparison branch against a fresh simulation (also when the parent already has the input and after an earlier year), the Worksheet A identities from actual liability, and the helper contracts. Co-Authored-By: Claude Opus 5.5 --- .../fix-branch-override-shadowing.fixed.md | 1 + .../tests/core/test_override_branches.py | 433 ++++++++++++++++++ 2 files changed, 434 insertions(+) create mode 100644 changelog.d/fix-branch-override-shadowing.fixed.md create mode 100644 policyengine_us/tests/core/test_override_branches.py diff --git a/changelog.d/fix-branch-override-shadowing.fixed.md b/changelog.d/fix-branch-override-shadowing.fixed.md new file mode 100644 index 00000000000..68081829e55 --- /dev/null +++ b/changelog.d/fix-branch-override-shadowing.fixed.md @@ -0,0 +1 @@ +Limit the non-refundable Child Tax Credit by the actual tax liability, SALT deduction included (26 U.S.C. 26(a); Schedule 8812 Credit Limit Worksheet A, line 1), instead of a recomputation without SALT that applied or not depending on which variables were calculated first; and make the itemization, Delaware and Virginia EITC, Idaho aged or disabled, Missouri TANF caretaker and Medicaid SSI-supplement comparison branches calculate under their overridden inputs even when the simulation has already calculated those inputs. diff --git a/policyengine_us/tests/core/test_override_branches.py b/policyengine_us/tests/core/test_override_branches.py new file mode 100644 index 00000000000..f155cd6b6ce --- /dev/null +++ b/policyengine_us/tests/core/test_override_branches.py @@ -0,0 +1,433 @@ +"""Formula branches calculate under their overrides, whatever was cached first. + +Several formulas compare a tax unit's liability under alternative choices by +calculating it in a branch with one input overridden: itemizing or not +(``tax_unit_itemizes``), Delaware and Virginia EITC refundability, the Idaho +aged or disabled credit or deduction, Missouri TANF caretaker inclusion and +Medicaid for SSI state supplements. A policyengine-core branch starts as a copy +of every array its parent has cached, and setting an input on it clears none +of them, so a value the parent calculated from the old input answers for the +branch. ``get_override_branch`` drops the copied values when the parent has +already calculated the overridden input, and creates the branch again for each +period; ``get_branch_for_period`` is the same without inputs. + +The non-refundable CTC used to be limited by the tax liability recomputed +without the SALT deduction, in a "no_salt" branch. The branch usually inherited +the liability with SALT, so the CTC depended on which variables a caller asked +for first. 26 U.S.C. 26(a) limits the credit by the actual tax liability, SALT +deduction included, and Schedule 8812 Credit Limit Worksheet A starts from it +(line 1 is Form 1040 line 18); ``ctc_tax_liability_after_preceding_credits`` +(line 3) now reads it directly. + +A seeded sample of households shares one simulation; the tests check: + +1. The CTC-limiting liability is Worksheet A line 5 computed from income tax + before credits: less the credits that precede the CTC (line 3), then less + the credits that follow it when Worksheet B applies; and the non-refundable + and refundable parts add up to the credit. +2. Every reported variable is the same whichever variable is calculated first, + on the seeded sample and on random households (Hypothesis). +3. Each comparison branch equals a fresh simulation that sets the overridden + input before anything is calculated (a differential test of the branch + against the reference path). +4. The same holds when the parent has already calculated the overridden + input, and when an earlier year was calculated first. +""" + +import numpy as np +import pytest +from hypothesis import HealthCheck, given, settings +from hypothesis import strategies as st +from policyengine_core.periods import period + +from policyengine_us import Simulation +from policyengine_us.tools.period_branch import ( + drop_inherited_values, + get_branch_for_period, + get_override_branch, +) + +YEAR = 2026 +SEED = 20261001 +N = 48 +STATES = ["CA", "NY", "VA", "DE", "ID", "NJ", "MA", "IL", "TX", "OR"] +REPORTED = [ + "income_tax", + "income_tax_before_credits", + "ctc_limiting_tax_liability", + "non_refundable_ctc", + "refundable_ctc", + "tax_unit_itemizes", + "state_income_tax", + "household_net_income", +] + + +def _sample(): + rng = np.random.default_rng(SEED) + households = [] + for i in range(N): + married = bool(rng.random() < 0.6) + children = int(rng.integers(0, 4)) + earnings = float(np.round(np.exp(rng.uniform(np.log(8_000), np.log(300_000))))) + households.append( + dict( + state=STATES[i % len(STATES)], + married=married, + children=children, + earnings=earnings, + spouse_share=float(rng.choice([0, 0.2, 0.5])) if married else 0, + mortgage=float(rng.choice([0, 8_000, 18_000, 30_000])), + property_tax=float(rng.choice([0, 3_000, 9_000, 16_000])), + charity=float(rng.choice([0, 2_000, 9_000])), + aged_parent=bool(rng.random() < 0.15), + ) + ) + return households + + +def _situation(households, years=(YEAR,), itemizes=None): + people, units, marital = {}, {}, {} + for i, h in enumerate(households): + members = [] + + def add(name, **inputs): + people[name] = {k: {y: v for y in years} for k, v in inputs.items()} + members.append(name) + + head = f"h{i}" + earnings = h["earnings"] + add( + head, + age=40, + employment_income=earnings * (1 - h["spouse_share"]), + deductible_mortgage_interest=h["mortgage"], + real_estate_taxes=h["property_tax"], + charitable_cash_donations=h["charity"], + ) + marital[f"mu_{head}"] = {"members": [head]} + if h["married"]: + add(f"s{i}", age=38, employment_income=earnings * h["spouse_share"]) + marital[f"mu_{head}"]["members"].append(f"s{i}") + for c in range(h["children"]): + add(f"c{i}_{c}", age=3 + 4 * c, is_tax_unit_dependent=True) + marital[f"mu_c{i}_{c}"] = {"members": [f"c{i}_{c}"]} + if h["aged_parent"]: + add( + f"p{i}", + age=74, + is_tax_unit_dependent=True, + share_of_care_and_support_costs_paid_by_tax_filer=1.0, + ) + marital[f"mu_p{i}"] = {"members": [f"p{i}"]} + units[i] = members + tax_units = {f"tu{i}": {"members": m} for i, m in units.items()} + if itemizes is not None: + for i, unit in enumerate(tax_units.values()): + unit["tax_unit_itemizes"] = {y: bool(itemizes[i]) for y in years} + return { + "people": people, + "tax_units": tax_units, + "spm_units": {f"spm{i}": {"members": m} for i, m in units.items()}, + "families": {f"fam{i}": {"members": m} for i, m in units.items()}, + "marital_units": marital, + "households": { + f"hh{i}": { + "members": m, + "state_code": {y: households[i]["state"] for y in years}, + } + for i, m in units.items() + }, + } + + +@pytest.fixture(scope="module") +def households(): + return _sample() + + +@pytest.fixture(scope="module") +def baseline(households): + simulation = Simulation(situation=_situation(households)) + return {v: simulation.calculate(v, YEAR) for v in REPORTED} + + +def test_ctc_limit_starts_from_actual_liability(households): + simulation = Simulation(situation=_situation(households)) + limit = simulation.calculate("ctc_limiting_tax_liability", YEAR) + before_credits = simulation.calculate("income_tax_before_credits", YEAR) + p = simulation.tax_benefit_system.parameters(YEAR).gov.irs.credits + preceding = sum( + simulation.calculate(credit, YEAR) + for credit in p.ctc_tax_liability_limit.preceding_credits + ) + subsequent = sum( + simulation.calculate(credit, YEAR) + for credit in p.ctc_tax_liability_limit.subsequent_credits + ) + worksheet_b = simulation.calculate("ctc_credit_limit_worksheet_b_applies", YEAR) + line_3 = np.maximum(0, before_credits - preceding) + np.testing.assert_allclose( + simulation.calculate("ctc_tax_liability_after_preceding_credits", YEAR), + line_3, + ) + np.testing.assert_allclose( + limit, np.maximum(0, line_3 - np.where(worksheet_b, subsequent, 0)) + ) + ctc = simulation.calculate("ctc", YEAR) + non_refundable = simulation.calculate("non_refundable_ctc", YEAR) + refundable = simulation.calculate("refundable_ctc", YEAR) + # non_refundable_ctc is the credit less its refundable part; the tax + # limit applies when non-refundable credits are capped. + assert (refundable <= ctc + 0.01).all() + np.testing.assert_allclose(non_refundable + refundable, ctc, atol=0.01) + # The sample includes itemizers with SALT whose CTC the limit binds. + salt = simulation.calculate("salt_deduction", YEAR) + itemizes = simulation.calculate("tax_unit_itemizes", YEAR) + assert (itemizes & (salt > 0) & (limit < ctc)).any() + + +@pytest.mark.parametrize( + "first", + [ + "income_tax", + "refundable_ctc", + "ctc_value", + "state_income_tax", + "tax_liability_if_itemizing", + "spm_unit_net_income", + ], +) +def test_results_do_not_depend_on_calculation_order(households, baseline, first): + simulation = Simulation(situation=_situation(households)) + simulation.calculate(first, YEAR) + for variable in REPORTED: + np.testing.assert_array_equal( + simulation.calculate(variable, YEAR), + baseline[variable], + err_msg=f"{variable} changes when {first} is calculated first", + ) + + +household_strategy = st.fixed_dictionaries( + { + "state": st.sampled_from(STATES), + "married": st.booleans(), + "children": st.integers(0, 3), + "earnings": st.integers(8_000, 300_000), + "spouse_share": st.sampled_from([0, 0.2, 0.5]), + "mortgage": st.integers(0, 40_000), + "property_tax": st.integers(0, 20_000), + "charity": st.integers(0, 10_000), + "aged_parent": st.booleans(), + } +) + + +@settings( + max_examples=4, + deadline=None, + suppress_health_check=[HealthCheck.too_slow, HealthCheck.data_too_large], +) +@given( + st.lists(household_strategy, min_size=1, max_size=8), + st.sampled_from( + ["income_tax", "refundable_ctc", "tax_liability_if_itemizing", "eitc"] + ), +) +def test_random_households_do_not_depend_on_calculation_order(households, first): + reference = Simulation(situation=_situation(households)) + reference.calculate("household_net_income", YEAR) + simulation = Simulation(situation=_situation(households)) + simulation.calculate(first, YEAR) + for variable in REPORTED: + np.testing.assert_array_equal( + simulation.calculate(variable, YEAR), + reference.calculate(variable, YEAR), + err_msg=f"{variable} changes when {first} is calculated first", + ) + # Credit Limit Worksheet A, line 3 starts from the actual liability. + p = simulation.tax_benefit_system.parameters(YEAR).gov.irs.credits + preceding = sum( + simulation.calculate(credit, YEAR) + for credit in p.ctc_tax_liability_limit.preceding_credits + ) + np.testing.assert_allclose( + simulation.calculate("ctc_tax_liability_after_preceding_credits", YEAR), + np.maximum( + 0, simulation.calculate("income_tax_before_credits", YEAR) - preceding + ), + ) + + +def _fresh(households, variable, overrides, year=YEAR): + simulation = Simulation(situation=_situation(households, years=(year,))) + for name, value in overrides.items(): + simulation.set_input(name, year, value) + return simulation.calculate(variable, year) + + +def test_itemization_branches_match_fresh_simulations(households): + simulation = Simulation(situation=_situation(households)) + simulation.calculate("household_net_income", YEAR) + n = len(households) + for comparison, itemizes in ( + ("tax_liability_if_itemizing", True), + ("tax_liability_if_not_itemizing", False), + ): + np.testing.assert_allclose( + simulation.calculate(comparison, YEAR), + _fresh(households, "income_tax", {"tax_unit_itemizes": [itemizes] * n}), + atol=0.01, + err_msg=comparison, + ) + + +@pytest.mark.parametrize( + "comparison,variable,override,value", + [ + ( + "de_income_tax_if_claiming_refundable_eitc", + "de_income_tax", + "de_claims_refundable_eitc", + True, + ), + ( + "de_income_tax_if_claiming_non_refundable_eitc", + "de_income_tax", + "de_claims_refundable_eitc", + False, + ), + ( + "va_income_tax_if_claiming_refundable_eitc", + "va_income_tax", + "va_claims_refundable_eitc", + True, + ), + ( + "va_income_tax_if_claiming_non_refundable_eitc", + "va_income_tax", + "va_claims_refundable_eitc", + False, + ), + ( + "id_income_tax_if_receiving_aged_or_disabled_credit", + "id_income_tax", + "id_receives_aged_or_disabled_credit", + True, + ), + ( + "id_income_tax_if_receiving_aged_or_disabled_deduction", + "id_income_tax", + "id_receives_aged_or_disabled_credit", + False, + ), + ], +) +def test_state_choice_branches_match_fresh_simulations( + households, comparison, variable, override, value +): + simulation = Simulation(situation=_situation(households)) + simulation.calculate("household_net_income", YEAR) + np.testing.assert_allclose( + simulation.calculate(comparison, YEAR), + _fresh(households, variable, {override: [value] * len(households)}), + atol=0.01, + ) + + +def test_branch_after_parent_calculated_the_overridden_input(households): + # tax_unit_itemizes is an input here, so the parent calculates income tax + # without the itemizing branch; the branch is created afterwards and must + # not answer with the parent's income tax. + n = len(households) + situation = _situation(households, itemizes=[False] * n) + simulation = Simulation(situation=situation) + not_itemizing = simulation.calculate("income_tax", YEAR) + itemizing = simulation.calculate("tax_liability_if_itemizing", YEAR) + expected = _fresh(households, "income_tax", {"tax_unit_itemizes": [True] * n}) + np.testing.assert_allclose(itemizing, expected, atol=0.01) + assert not np.allclose(itemizing, not_itemizing) + + +def test_later_year_branches_match_single_year_simulation(households): + years = (2025, YEAR) + simulation = Simulation(situation=_situation(households, years=years)) + simulation.calculate("tax_liability_if_itemizing", 2025) + simulation.calculate("income_tax", 2025) + fresh = Simulation(situation=_situation(households, years=years)) + for variable in REPORTED + ["tax_liability_if_itemizing"]: + np.testing.assert_array_equal( + simulation.calculate(variable, YEAR), + fresh.calculate(variable, YEAR), + err_msg=f"{variable} for {YEAR} depends on 2025 being calculated first", + ) + assert simulation.branches["itemizing"].branch_period.start.year == YEAR + + +def test_get_override_branch_reuses_within_period_and_recreates_otherwise( + households, +): + simulation = Simulation(situation=_situation(households, years=(2025, YEAR))) + n = len(households) + ones = np.ones(n, dtype=bool) + first = get_override_branch( + simulation, "test_branch", YEAR, {"tax_unit_itemizes": ones} + ) + assert ( + get_override_branch( + simulation, "test_branch", YEAR, {"tax_unit_itemizes": ones} + ) + is first + ) + other_inputs = get_override_branch( + simulation, "test_branch", YEAR, {"tax_unit_itemizes": ~ones} + ) + assert other_inputs is not first + other_period = get_override_branch( + simulation, "test_branch", 2025, {"tax_unit_itemizes": ones} + ) + assert other_period is not other_inputs + assert simulation.branches["test_branch"] is other_period + + +def test_get_branch_for_period_is_the_branch_without_inputs(households): + simulation = Simulation(situation=_situation(households, years=(2025, YEAR))) + # A branch of the same name that the helpers did not create is replaced. + plain = simulation.get_branch("pinned_branch") + first = get_branch_for_period(simulation, "pinned_branch", YEAR) + assert first is not plain + assert first.branch_period == period(YEAR) + assert get_branch_for_period(simulation, "pinned_branch", YEAR) is first + assert get_override_branch(simulation, "pinned_branch", YEAR, {}) is first + later = get_branch_for_period(simulation, "pinned_branch", 2025) + assert later is not first + assert later.branch_period == period(2025) + # Inside the branch itself, core's get_branch returns the simulation. + assert get_branch_for_period(later, "pinned_branch", 2025) is later + + +def test_drop_inherited_values_keeps_inputs_only(households): + n = len(households) + itemizes = np.ones(n, dtype=bool) + simulation = Simulation(situation=_situation(households)) + simulation.calculate("income_tax", YEAR) + # A plain branch with an input of its own, and a branch nested in it. + parent = simulation.get_branch("parent_with_input") + parent.set_input("tax_unit_itemizes", YEAR, itemizes) + child = parent.get_branch("child") + drop_inherited_values(child) + # Inputs survive: the situation's and the one set on the parent branch. + np.testing.assert_array_equal( + child.get_array("employment_income", YEAR), + simulation.get_array("employment_income", YEAR), + ) + np.testing.assert_array_equal(child.get_array("tax_unit_itemizes", YEAR), itemizes) + # Calculated values are gone, and are calculated again from the inputs. + assert child.get_array("income_tax", YEAR) is None + assert child.get_array("taxable_income", YEAR) is None + np.testing.assert_allclose( + child.calculate("taxable_income", YEAR), + _fresh(households, "taxable_income", {"tax_unit_itemizes": itemizes}), + atol=0.01, + ) From fb97c4e590a28fe592d8180d030e91c5e0d3a304 Mon Sep 17 00:00:00 2001 From: Max Ghenis Date: Mon, 5 Oct 2026 16:23:32 -0400 Subject: [PATCH 5/7] Drop the no-SALT note from the credit-order test docstring The CTC limit reads the actual liability now, whatever the deduction. Co-Authored-By: Claude Opus 5.5 --- policyengine_us/tests/test_federal_credit_limit_order.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/policyengine_us/tests/test_federal_credit_limit_order.py b/policyengine_us/tests/test_federal_credit_limit_order.py index bb05a1a844c..0444e200574 100644 --- a/policyengine_us/tests/test_federal_credit_limit_order.py +++ b/policyengine_us/tests/test_federal_credit_limit_order.py @@ -17,8 +17,7 @@ These tests check that the parameters encode that order, and that the model's credits match a direct sequential application of 26 U.S.C. 26(a) in that -order. The households take the standard deduction, so the CTC limit's no-SALT -liability equals actual liability. +order. The households take the standard deduction. """ import numpy as np From 2d8202a9782c5da969aa576ad324e02cb102d757 Mon Sep 17 00:00:00 2001 From: Max Ghenis Date: Mon, 5 Oct 2026 19:55:12 -0400 Subject: [PATCH 6/7] Limit the Oklahoma test's federal CTC by actual liability The case's tax before credits is $2,799.38 after its itemized SALT; the no-SALT limit this PR removes allowed $4,400 of non-refundable credit. Co-Authored-By: Claude Opus 5.5 --- .../states/ok/tax/income/ok_child_care_child_tax_credit.yaml | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/policyengine_us/tests/policy/baseline/gov/states/ok/tax/income/ok_child_care_child_tax_credit.yaml b/policyengine_us/tests/policy/baseline/gov/states/ok/tax/income/ok_child_care_child_tax_credit.yaml index bd3e96c010f..1f7b77e70cc 100644 --- a/policyengine_us/tests/policy/baseline/gov/states/ok/tax/income/ok_child_care_child_tax_credit.yaml +++ b/policyengine_us/tests/policy/baseline/gov/states/ok/tax/income/ok_child_care_child_tax_credit.yaml @@ -167,7 +167,10 @@ members: [head, child1, child2] state_code: OK output: - ctc_value: 4_400 + # Non-refundable CTC is limited by actual tax liability (26 U.S.C. 26(a)): + # income_tax_before_credits is $2,799.38 after $33,284 of SALT among the + # itemized deductions. The earlier no-SALT limit ($6,793.40) allowed $4,400. + ctc_value: 2_799 ok_federal_ctc: 2_799 ok_child_care_child_tax_credit: 140 taxsim_ok_child_tax_credit_component: 140 From 53cdd80bbd421545699202f03de36864826821dd Mon Sep 17 00:00:00 2001 From: Max Ghenis Date: Mon, 5 Oct 2026 23:06:58 -0400 Subject: [PATCH 7/7] Keep branch inputs by key and period, not by variable name drop_inherited_values kept every array of any variable in the simulation's input_variables, whatever its period. A formula variable supplied as an input for one year (e.g. taxable_income for 2025) then kept the parent's calculated value for another year in the branch: the 2026 itemizing branch answered with taxable income calculated without itemizing ($13,140 of income tax instead of $8,191.05, review r2). Keep an array only if set_input stored it under that variable, branch and period, on the branch or an ancestor; match eternal variables, which store every period under one key, by branch. This also stops dropping an eternal input set on an ancestor branch. Adds review r2's case as a regression, a Hypothesis differential over mixed input and formula years, and a Hypothesis property that the helper keeps exactly the input keys. All three fail on 39124edc77. Co-Authored-By: Claude Opus 5.5 --- .../fix-branch-override-shadowing.fixed.md | 2 +- .../tests/core/test_override_branches.py | 200 +++++++++++++++++- policyengine_us/tools/period_branch.py | 43 ++-- 3 files changed, 228 insertions(+), 17 deletions(-) diff --git a/changelog.d/fix-branch-override-shadowing.fixed.md b/changelog.d/fix-branch-override-shadowing.fixed.md index 68081829e55..a433672f230 100644 --- a/changelog.d/fix-branch-override-shadowing.fixed.md +++ b/changelog.d/fix-branch-override-shadowing.fixed.md @@ -1 +1 @@ -Limit the non-refundable Child Tax Credit by the actual tax liability, SALT deduction included (26 U.S.C. 26(a); Schedule 8812 Credit Limit Worksheet A, line 1), instead of a recomputation without SALT that applied or not depending on which variables were calculated first; and make the itemization, Delaware and Virginia EITC, Idaho aged or disabled, Missouri TANF caretaker and Medicaid SSI-supplement comparison branches calculate under their overridden inputs even when the simulation has already calculated those inputs. +Limit the non-refundable Child Tax Credit by the actual tax liability, SALT deduction included (26 U.S.C. 26(a); Schedule 8812 Credit Limit Worksheet A, line 1), instead of a recomputation without SALT that applied or not depending on which variables were calculated first; and make the itemization, Delaware and Virginia EITC, Idaho aged or disabled, Missouri TANF caretaker and Medicaid SSI-supplement comparison branches calculate under their overridden inputs even when the simulation has already calculated those inputs, including variables it was given as inputs only for another year. diff --git a/policyengine_us/tests/core/test_override_branches.py b/policyengine_us/tests/core/test_override_branches.py index f155cd6b6ce..9e46a63644e 100644 --- a/policyengine_us/tests/core/test_override_branches.py +++ b/policyengine_us/tests/core/test_override_branches.py @@ -32,11 +32,15 @@ against the reference path). 4. The same holds when the parent has already calculated the overridden input, and when an earlier year was calculated first. +5. When the branch drops what it copied, it keeps exactly the values set as + inputs, each for the period it was set for: a variable that is an input in + one year is calculated again in the others (Hypothesis, over random mixes + of input years and formula years). """ import numpy as np import pytest -from hypothesis import HealthCheck, given, settings +from hypothesis import HealthCheck, example, given, settings from hypothesis import strategies as st from policyengine_core.periods import period @@ -350,6 +354,193 @@ def test_branch_after_parent_calculated_the_overridden_input(households): assert not np.allclose(itemizing, not_itemizing) +YEARS = (2025, YEAR) +# Formula variables between itemization and income tax that a situation can +# also supply as inputs, here for one year only. +MIXED_YEAR_INPUTS = [ + "taxable_income", + "income_tax_main_rates", + "adjusted_gross_income", + "salt_deduction", + "ctc", +] +# The household from review r2 of PolicyEngine/policyengine-us#9741. +REVIEW_HOUSEHOLD = dict( + state="CA", + married=True, + children=2, + earnings=160_000, + spouse_share=0, + mortgage=30_000, + property_tax=14_000, + charity=0, + aged_parent=False, +) + + +def _with_tax_unit_inputs(situation, inputs): + """Add each ``variable: (year, value)`` to every tax unit, for that year.""" + for unit in situation["tax_units"].values(): + for variable, (year, value) in inputs.items(): + unit[variable] = {year: value} + return situation + + +def test_branch_recalculates_a_variable_that_is_an_input_in_another_year(): + # taxable_income is an input for 2025 only, so its 2026 value is + # calculated: by the parent, without itemizing. The itemizing branch has + # to calculate its own. Before the branch kept inputs by key and period, + # it answered with the parent's ($13,140 of income tax instead of the + # $8,191.05 an itemizing simulation gives). + situation = _with_tax_unit_inputs( + _situation([REVIEW_HOUSEHOLD], years=YEARS, itemizes=[False]), + {"taxable_income": (2025, 1)}, + ) + simulation = Simulation(situation=situation) + not_itemizing = simulation.calculate("income_tax", YEAR) + itemizing = simulation.calculate("tax_liability_if_itemizing", YEAR) + fresh = Simulation(situation=situation) + fresh.set_input("tax_unit_itemizes", YEAR, np.array([True])) + np.testing.assert_allclose( + itemizing, fresh.calculate("income_tax", YEAR), atol=0.01 + ) + assert not np.allclose(itemizing, not_itemizing) + + +@settings( + max_examples=4, + deadline=None, + suppress_health_check=[HealthCheck.too_slow, HealthCheck.data_too_large], +) +@given( + households=st.lists(household_strategy, min_size=1, max_size=4), + inputs=st.dictionaries( + st.sampled_from(MIXED_YEAR_INPUTS), + st.tuples(st.sampled_from(YEARS), st.integers(0, 100_000)), + min_size=1, + ), + year=st.sampled_from(YEARS), + first=st.sampled_from(["income_tax", "taxable_income", "household_net_income"]), + first_year=st.sampled_from(YEARS), + parent_itemizes=st.sampled_from([None, False, True]), +) +@example( + households=[REVIEW_HOUSEHOLD], + inputs={"income_tax_main_rates": (2025, 1)}, + year=YEAR, + first="income_tax", + first_year=YEAR, + parent_itemizes=False, +) +def test_branches_with_inputs_in_other_years_match_fresh_simulations( + households, inputs, year, first, first_year, parent_itemizes +): + # Each variable in ``inputs`` is an input in one year and a formula in + # the other; tax_unit_itemizes is an input too unless parent_itemizes is + # None. Both comparison branches equal a simulation with the same inputs + # that sets itemization before calculating anything. + n = len(households) + itemizes = None if parent_itemizes is None else [parent_itemizes] * n + situation = _with_tax_unit_inputs( + _situation(households, years=YEARS, itemizes=itemizes), inputs + ) + simulation = Simulation(situation=situation) + simulation.calculate(first, first_year) + for comparison, itemizing in ( + ("tax_liability_if_itemizing", True), + ("tax_liability_if_not_itemizing", False), + ): + fresh = Simulation(situation=situation) + fresh.set_input("tax_unit_itemizes", year, np.full(n, itemizing)) + np.testing.assert_allclose( + simulation.calculate(comparison, year), + fresh.calculate("income_tax", year), + atol=0.01, + err_msg=f"{comparison} for {year} after {first} for {first_year}", + ) + + +@settings( + max_examples=6, + deadline=None, + suppress_health_check=[HealthCheck.too_slow, HealthCheck.data_too_large], +) +@given( + households=st.lists(household_strategy, min_size=1, max_size=3), + situation_inputs=st.dictionaries( + st.sampled_from(MIXED_YEAR_INPUTS), st.sampled_from(YEARS) + ), + branch_inputs=st.dictionaries( + st.sampled_from(MIXED_YEAR_INPUTS + ["tax_unit_itemizes"]), + st.sampled_from(YEARS), + ), + calculated=st.lists( + st.tuples( + st.sampled_from(["income_tax", "taxable_income", "tax_unit_itemizes"]), + st.sampled_from(YEARS), + ), + min_size=1, + max_size=3, + ), + head_input=st.sampled_from([None, "situation", "branch"]), +) +def test_drop_inherited_values_keeps_exactly_the_input_keys( + households, situation_inputs, branch_inputs, calculated, head_input +): + # Inputs in random years, on the simulation (from the situation) and on + # a branch of it; values calculated in random years on both. A branch of + # that branch, after drop_inherited_values, holds each input for the + # period it was set for, the nearer branch's first, and nothing else. + n = len(households) + situation = _with_tax_unit_inputs( + _situation(households, years=YEARS), + {variable: (year, 1_000) for variable, year in situation_inputs.items()}, + ) + heads = np.array([name.startswith("h") for name in situation["people"]]) + if head_input == "situation": + for name in situation["people"]: + if name.startswith("h"): + situation["people"][name]["is_household_head"] = {YEAR: True} + simulation = Simulation(situation=situation) + for variable, year in calculated: + simulation.calculate(variable, year) + parent = simulation.get_branch("parent") + expected = { + (variable, year): np.full(n, 1_000.0) + for variable, year in situation_inputs.items() + } + for variable, year in branch_inputs.items(): + value = np.ones(n, dtype=bool) if variable == "tax_unit_itemizes" else 2_000.0 + parent.set_input(variable, year, np.broadcast_to(value, (n,)).copy()) + expected[(variable, year)] = np.broadcast_to(value, (n,)) + if head_input == "branch": + parent.set_input("is_household_head", YEAR, heads) + for variable, year in calculated: + parent.calculate(variable, year) + child = parent.get_branch("child") + drop_inherited_values(child) + for variable in MIXED_YEAR_INPUTS + [ + "tax_unit_itemizes", + "income_tax", + "income_tax_before_credits", + ]: + for year in YEARS: + kept = child.get_array(variable, year) + if (variable, year) in expected: + np.testing.assert_array_equal( + kept, expected[(variable, year)], err_msg=f"{variable} {year}" + ) + else: + assert kept is None, f"calculated {variable} {year} was kept" + # An eternal variable is stored once for every period. + for year in YEARS: + head = child.get_array("is_household_head", year) + if head_input is None: + assert head is None + else: + np.testing.assert_array_equal(head, heads) + + def test_later_year_branches_match_single_year_simulation(households): years = (2025, YEAR) simulation = Simulation(situation=_situation(households, years=years)) @@ -417,10 +608,11 @@ def test_drop_inherited_values_keeps_inputs_only(households): parent.set_input("tax_unit_itemizes", YEAR, itemizes) child = parent.get_branch("child") drop_inherited_values(child) - # Inputs survive: the situation's and the one set on the parent branch. + # Inputs survive: the situation's (its employment income is stored as + # employment_income_before_lsr) and the one set on the parent branch. np.testing.assert_array_equal( - child.get_array("employment_income", YEAR), - simulation.get_array("employment_income", YEAR), + child.get_array("employment_income_before_lsr", YEAR), + simulation.get_array("employment_income_before_lsr", YEAR), ) np.testing.assert_array_equal(child.get_array("tax_unit_itemizes", YEAR), itemizes) # Calculated values are gone, and are calculated again from the inputs. diff --git a/policyengine_us/tools/period_branch.py b/policyengine_us/tools/period_branch.py index 439e35028aa..5e47797cf20 100644 --- a/policyengine_us/tools/period_branch.py +++ b/policyengine_us/tools/period_branch.py @@ -22,7 +22,8 @@ parent's cache is then safe to share; this is the usual case, where a formula branches while its parent is still calculating the variable the branch overrides. Otherwise the branch drops every array it copied except - inputs, and calculates the rest itself. + inputs, each for the periods it was set for, and calculates the rest + itself. - A branch reused within its period with different inputs is created again. ``get_branch_for_period`` is the same with no inputs, for branches that @@ -30,10 +31,10 @@ and deletes the variables it recalculates. """ -from typing import Dict, Tuple, Union +from typing import Dict, Set, Tuple, Union import numpy as np -from policyengine_core.periods import Period +from policyengine_core.periods import ETERNITY, Period from policyengine_core.periods import period as to_period from policyengine_core.simulations import Simulation @@ -56,23 +57,41 @@ def _is_known(simulation: Simulation, variable: str, period: Period) -> bool: return holder.get_array(period, simulation.branch_name) is not None +def _input_keys(branch: Simulation) -> Set[Tuple[str, str, Period]]: + """The (variable, branch name, period) keys ``branch`` reads as inputs. + + policyengine-core records each key that ``set_input`` stores, whether + from the dataset, the situation or a branch, in one set shared by a + simulation and all its branches. ``branch`` reads its own keys and its + ancestors'. + """ + visible_branches = set(branch._get_visible_branch_names()) + return { + key + for key in getattr(branch, "_user_input_keys", set()) + if key[1] in visible_branches + } + + def drop_inherited_values(branch: Simulation) -> None: """Delete every array ``branch`` holds except inputs. - Inputs are the variables the simulation was built with and each value set - through ``set_input`` on a branch this one reads. + An array is kept only if ``set_input`` stored it, on this branch or one it + reads, for that variable and period. A value calculated for one period is + dropped even when the same variable is an input for another period. """ - input_variables = set(branch.input_variables) - user_input_keys = getattr(branch, "_user_input_keys", set()) - visible_branches = set(branch._get_visible_branch_names()) + input_keys = _input_keys(branch) + # An eternal variable stores every period under one key, while the input + # key records the period it was set for, so match it by branch only. + eternal_inputs = {(name, branch_name) for name, branch_name, _ in input_keys} for population in branch.populations.values(): for name, holder in population._holders.items(): - if name in input_variables: - continue + eternal = holder.variable.definition_period == ETERNITY for branch_name, known_period in holder.get_known_branch_periods(): if ( - branch_name in visible_branches - and (name, branch_name, known_period) in user_input_keys + (name, branch_name) in eternal_inputs + if eternal + else (name, branch_name, known_period) in input_keys ): continue # Exact key: ``Holder.delete_arrays`` would also delete any