Skip to content
1 change: 1 addition & 0 deletions changelog.d/fix-carry-over-order.fixed.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Auto-carry-over now carries only inputs, taking the latest one stored for a period that starts no later than the requested period, so a carried value no longer depends on which periods were calculated first or on a later input.
1 change: 1 addition & 0 deletions changelog.d/fix-uprating-order.fixed.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Uprating now starts from the latest earlier input in the variable's own unit and skips periods the simulation calculated, so calculating intermediate periods first no longer compounds rounding or truncation, or carries an eligibility mask or a default into later uprated values.
41 changes: 40 additions & 1 deletion policyengine_core/data_storage/in_memory_storage.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,8 +47,19 @@ def __init__(self, is_eternal: bool):
# with a copy the first time it is read. A key left here after code
# outside this class empties ``_arrays`` costs one extra copy at most.
self._shared = set()
# Keys whose value was stored with ``put(..., derived=True)``: calculated
# by the simulation rather than taken as input. A key counts only
# while it is stored, and every ``put`` sets or clears its mark.
self._derived = set()
self.is_eternal = is_eternal

def __setstate__(self, state: dict) -> None:
# A storage pickled before derived marks or shared arrays existed has
# neither: its values count as inputs and as its own.
state.setdefault("_derived", set())
state.setdefault("_shared", set())
self.__dict__.update(state)

def clone(self, share_arrays: bool = False) -> "InMemoryStorage":
"""Copy this storage.

Expand Down Expand Up @@ -83,6 +94,7 @@ def clone(self, share_arrays: bool = False) -> "InMemoryStorage":
clone._arrays[key] = array.copy()
else:
clone._arrays = {key: array.copy() for key, array in self._arrays.items()}
clone._derived = set(self._derived)
return clone

def get(self, period: Period, branch_name: str = "default") -> ArrayLike:
Expand All @@ -101,8 +113,29 @@ def get(self, period: Period, branch_name: str = "default") -> ArrayLike:
self._shared.discard(key)
return values

def has(self, period: Period, branch_name: str = "default") -> bool:
"""Whether a value is stored for ``period`` under ``branch_name``.

Unlike ``get``, this never copies an array shared by ``clone``.
"""
if self.is_eternal:
period = periods.period(periods.ETERNITY)
return f"{branch_name}:{periods.period(period)}" in self._arrays

def is_derived(self, period: Period, branch_name: str = "default") -> bool:
"""Whether the value stored for ``period`` under ``branch_name`` was
stored with ``derived=True``; ``False`` if none is stored."""
if self.is_eternal:
period = periods.period(periods.ETERNITY)
key = f"{branch_name}:{periods.period(period)}"
return key in self._derived and key in self._arrays

def put(
self, value: ArrayLike, period: Period, branch_name: str = "default"
self,
value: ArrayLike,
period: Period,
branch_name: str = "default",
derived: bool = False,
) -> None:
if self.is_eternal:
period = periods.period(periods.ETERNITY)
Expand All @@ -126,6 +159,10 @@ def put(
key = f"{branch_name}:{period}"
self._arrays[key] = value
self._shared.discard(key)
if derived:
self._derived.add(key)
else:
self._derived.discard(key)

def delete(self, period: Period = None, branch_name: str = "default") -> None:
if period is None:
Expand All @@ -138,6 +175,7 @@ def delete(self, period: Period = None, branch_name: str = "default") -> None:
if not period_item.startswith(branch_prefix)
}
self._shared.intersection_update(self._arrays)
self._derived.intersection_update(self._arrays)
return

if self.is_eternal:
Expand All @@ -156,6 +194,7 @@ def delete(self, period: Period = None, branch_name: str = "default") -> None:
)
}
self._shared.intersection_update(self._arrays)
self._derived.intersection_update(self._arrays)

def get_known_periods(self) -> list:
# Split on the first colon only: an anchored period's string form
Expand Down
41 changes: 40 additions & 1 deletion policyengine_core/data_storage/on_disk_storage.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,10 +22,19 @@ def __init__(
):
self._files = {}
self._enums = {}
# File keys stored with ``put(..., derived=True)``; see
# ``InMemoryStorage``.
self._derived = set()
self.is_eternal = is_eternal
self.preserve_storage_dir = preserve_storage_dir
self.storage_dir = storage_dir

def __setstate__(self, state: dict) -> None:
# A storage pickled before derived marks existed has none: its values
# count as inputs.
state.setdefault("_derived", set())
self.__dict__.update(state)

def clone(self) -> "OnDiskStorage":
"""Create a private metadata view over this storage directory.

Expand All @@ -43,6 +52,7 @@ def clone(self) -> "OnDiskStorage":
)
clone._files = self._files.copy()
clone._enums = self._enums.copy()
clone._derived = set(self._derived)
clone._storage_dir_owner = getattr(self, "_storage_dir_owner", self)
return clone

Expand All @@ -63,8 +73,29 @@ def get(self, period: Period, branch_name: str = "default") -> ArrayLike:
return None
return self._decode_file(values)

def has(self, period: Period, branch_name: str = "default") -> bool:
"""Whether a value is stored for ``period`` under ``branch_name``.

Unlike ``get``, this reads no file.
"""
if self.is_eternal:
period = periods.period(periods.ETERNITY)
return f"{branch_name}_{periods.period(period)}" in self._files

def is_derived(self, period: Period, branch_name: str = "default") -> bool:
"""Whether the value stored for ``period`` under ``branch_name`` was
stored with ``derived=True``; ``False`` if none is stored."""
if self.is_eternal:
period = periods.period(periods.ETERNITY)
key = f"{branch_name}_{periods.period(period)}"
return key in self._derived and key in self._files

def put(
self, value: ArrayLike, period: Period, branch_name: str = "default"
self,
value: ArrayLike,
period: Period,
branch_name: str = "default",
derived: bool = False,
) -> None:
if self.is_eternal:
period = periods.period(periods.ETERNITY)
Expand All @@ -77,6 +108,10 @@ def put(
value = value.view(numpy.ndarray)
numpy.save(path, value)
self._files[filename] = path
if derived:
self._derived.add(filename)
else:
self._derived.discard(filename)

def delete(self, period: Period = None, branch_name: str = "default") -> None:
if period is None:
Expand All @@ -89,6 +124,7 @@ def delete(self, period: Period = None, branch_name: str = "default") -> None:
for period_item, value in self._files.items()
if not period_item.startswith(branch_prefix)
}
self._derived.intersection_update(self._files)
return

if self.is_eternal:
Expand All @@ -101,6 +137,7 @@ def delete(self, period: Period = None, branch_name: str = "default") -> None:
for period_item, value in self._files.items()
if not period_item == f"{branch_name}_{period}"
}
self._derived.intersection_update(self._files)

def get_known_periods(self) -> list:
return list([periods.period(x.split("_")[1]) for x in self._files.keys()])
Expand All @@ -113,6 +150,8 @@ def get_known_branch_periods(self) -> list:

def restore(self) -> None:
self._files = files = {}
# Files read back from a directory carry no derived marks.
self._derived = set()
# Restore self._files from content of storage_dir.
for filename in os.listdir(self.storage_dir):
if not filename.endswith(".npy"):
Expand Down
117 changes: 112 additions & 5 deletions policyengine_core/holders/holder.py
Original file line number Diff line number Diff line change
Expand Up @@ -347,9 +347,15 @@ def _set(
value: ArrayLike,
branch_name: str = "default",
validate_nan: bool = False,
derived: bool = False,
) -> None:
simulation = getattr(self, "simulation", None)
user_input_contexts = getattr(simulation, "_user_input_contexts", None)
# A value calculated while an input is being set (say, by a
# ``set_input`` helper that calculates) is not part of that input: it
# belongs to the branch it was calculated on.
user_input_contexts = (
None if derived else getattr(simulation, "_user_input_contexts", None)
)
if user_input_contexts and branch_name == "default":
branch_name = user_input_contexts[-1]
value = self._to_array(value, validate_nan=validate_nan)
Expand All @@ -367,17 +373,34 @@ def _set(
)

if should_store_on_disk:
self._disk_storage.put(value, period, branch_name)
self._disk_storage.put(value, period, branch_name, derived=derived)
else:
self._memory_storage.put(value, period, branch_name)
self._memory_storage.put(value, period, branch_name, derived=derived)
if user_input_contexts:
if not hasattr(simulation, "_user_input_keys"):
simulation._user_input_keys = set()
simulation._user_input_keys.add((self.variable.name, branch_name, period))

def put_in_cache(
self, value: ArrayLike, period: Period, branch_name: str = "default"
self,
value: ArrayLike,
period: Period,
branch_name: str = "default",
derived: bool = False,
) -> None:
"""Cache ``value`` for ``period``.

``derived`` marks a value the simulation calculated rather than took
as input: a formula result, a carried, uprated or default value, a
twelfth of a yearly flow cached at a month by ``calculate_divide``,
or a sum over several sub-periods cached by ``calculate_add``.
Auto-carry-over never carries such a value into another period (see
``is_derived``). The mark is stored with the value, for the (branch,
period) key written, and any later write to that key replaces it.

A derived value never replaces an input that ``get_array(period,
branch_name)`` reads: the input is kept and nothing is stored.
"""
if self._do_not_store:
return

Expand All @@ -388,11 +411,95 @@ def put_in_cache(
):
return

self._set(period, value, branch_name)
if (
derived
and self._branch_storing(period, branch_name) is not None
and not self.is_derived(period, branch_name)
):
return

self._set(period, value, branch_name, derived=derived)

def default_array(self) -> ArrayLike:
"""
Return a new array of the appropriate length for the entity, filled with the variable default values.
"""

return self.variable.default_array(self.population.count)

def _stores(self, period: Period, branch_name: str) -> bool:
"""Whether a value is stored for ``period`` under ``branch_name``,
without reading or copying it."""
return self._memory_storage.has(period, branch_name) or (
self._disk_storage is not None
and self._disk_storage.has(period, branch_name)
)

def _readable_branches(self, branch_name: str = "default") -> List[str]:
"""``get_array``'s lookup order: the branch, its ``parent_branch``
ancestors, then ``default``."""
names = [branch_name]
if branch_name != "default":
parent = (
getattr(self.simulation, "parent_branch", None)
if self.simulation
else None
)
while parent is not None:
names.append(parent.branch_name)
parent = getattr(parent, "parent_branch", None)
names.append("default")
return list(dict.fromkeys(names))

def _branch_storing(self, period: Period, branch_name: str = "default") -> str:
"""The branch whose stored value ``get_array(period, branch_name)``
reads, or ``None`` if none stores one."""
for name in self._readable_branches(branch_name):
if self._stores(period, name):
return name
return None

def is_derived(self, period: Period, branch_name: str = "default") -> bool:
"""Whether the value ``get_array(period, branch_name)`` reads was
calculated by the simulation rather than set as an input.

The answer comes from the branch that stores the value read: the
branch itself, else its ``parent_branch`` ancestors, else
``default``. ``False`` if no value is stored for ``period``.
"""
storing = self._branch_storing(period, branch_name)
if storing is None:
return False
if self._memory_storage.has(period, storing):
return self._memory_storage.is_derived(period, storing)
return self._disk_storage.is_derived(period, storing)

def get_input_periods(self, branch_name: str = "default") -> List[Period]:
"""The periods for which the value ``get_array(period, branch_name)``
reads is an input rather than a value the simulation calculated (see
``put_in_cache``). Periods stored only under branches this one cannot
read are left out.

One pass over the stored keys: for each period, the key ``get_array``
reads first (the branch before its ancestors, memory before disk).
"""
rank = {
name: index
for index, name in enumerate(self._readable_branches(branch_name))
}
storages = [self._memory_storage]
if self._disk_storage is not None:
storages.append(self._disk_storage)
read = {}
for order, storage in enumerate(storages):
for stored_branch, period in storage.get_known_branch_periods():
if stored_branch not in rank:
continue
key = (rank[stored_branch], order)
if period not in read or key < read[period][0]:
read[period] = (key, storage, stored_branch)
return [
period
for period, (_, storage, stored_branch) in read.items()
if not storage.is_derived(period, stored_branch)
]
Loading
Loading