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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions changelog.d/549.changed.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Raise the first simulation cache memory warning from 8 GiB to 16 GiB while retaining the 32 GiB warning.
7 changes: 4 additions & 3 deletions src/policyengine/core/cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@

logger = logging.getLogger(__name__)

_MEMORY_THRESHOLDS_GB = [8, 16, 32]
_MEMORY_THRESHOLDS_GIB = [16, 32]
_warned_thresholds: set[int] = set()

T = TypeVar("T")
Expand Down Expand Up @@ -50,10 +50,11 @@ def _check_memory_usage(self) -> None:
process = psutil.Process()
memory_gb = process.memory_info().rss / (1024**3)

for threshold in _MEMORY_THRESHOLDS_GB:
for threshold in _MEMORY_THRESHOLDS_GIB:
if memory_gb >= threshold and threshold not in _warned_thresholds:
logger.warning(
f"Memory usage has reached {memory_gb:.2f}GB (threshold: {threshold}GB). "
f"Memory usage has reached {memory_gb:.2f} GiB "
f"(threshold: {threshold} GiB). "
f"Cache contains {len(self._cache)} items."
)
_warned_thresholds.add(threshold)
74 changes: 74 additions & 0 deletions tests/test_cache.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,13 @@
import os
import tempfile
from types import SimpleNamespace

import pandas as pd
import pytest
from microdf import MicroDataFrame

from policyengine.core import Simulation
from policyengine.core import cache as cache_module
from policyengine.core.cache import LRUCache
from policyengine.tax_benefit_models.uk import (
PolicyEngineUKDataset,
Expand Down Expand Up @@ -143,3 +146,74 @@ def test_lru_cache_clear():
assert cache.get("a") is None
assert cache.get("b") is None
assert cache.get("c") is None


@pytest.mark.parametrize("memory_gib", [0, 8, 15.99])
def test_lru_cache_does_not_warn_below_16_gib(
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
memory_gib: float,
) -> None:
cache_module._warned_thresholds.clear()
monkeypatch.setattr(
cache_module.psutil,
"Process",
lambda: SimpleNamespace(
memory_info=lambda: SimpleNamespace(rss=memory_gib * 1024**3)
),
)

with caplog.at_level("WARNING", logger=cache_module.__name__):
LRUCache[str]().add("key", "value")

assert caplog.records == []


def test_lru_cache_warns_once_at_16_gib_and_again_at_32_gib(
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
) -> None:
cache_module._warned_thresholds.clear()
memory = SimpleNamespace(gib=16.0)
monkeypatch.setattr(
cache_module.psutil,
"Process",
lambda: SimpleNamespace(
memory_info=lambda: SimpleNamespace(rss=memory.gib * 1024**3)
),
)
cache = LRUCache[str]()

with caplog.at_level("WARNING", logger=cache_module.__name__):
cache.add("first", "value")
cache.add("second", "value")
memory.gib = 32.0
cache.add("third", "value")

messages = [record.getMessage() for record in caplog.records]
assert len(messages) == 2
assert "threshold: 16 GiB" in messages[0]
assert "threshold: 32 GiB" in messages[1]


def test_lru_cache_clear_resets_memory_warning_state(
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
) -> None:
cache_module._warned_thresholds.clear()
monkeypatch.setattr(
cache_module.psutil,
"Process",
lambda: SimpleNamespace(memory_info=lambda: SimpleNamespace(rss=16 * 1024**3)),
)
cache = LRUCache[str]()

with caplog.at_level("WARNING", logger=cache_module.__name__):
cache.add("first", "value")
cache.clear()
cache.add("second", "value")

assert (
sum("threshold: 16 GiB" in record.getMessage() for record in caplog.records)
== 2
)
Loading