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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@ sources = ["src"]

[project]
name = "jhechms"
version = "0.2.3"
version = "0.2.4"
description = "HEC-HMS hydrological model -- symfluence plugin"
readme = "README.md"
license = { file = "LICENSE" }
Expand Down
52 changes: 50 additions & 2 deletions src/jhechms/calibration/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,54 @@
import jax.numpy as jnp


def _calibration_slice(worker):
"""Calibration-period slice, tolerating symfluence releases without it.

``InMemoryModelWorker.get_calibration_slice()`` is the shared
implementation and is used whenever it is available. Releases at or
below symfluence 0.9.2 predate it, and this package must not quietly
fall back to scoring the whole post-warmup record there — losing the
calibration window is exactly the bug this guards against.

Args:
worker: The in-memory worker whose config and time index to read.

Returns:
``(start, end)`` within the post-warmup arrays, or None when no
calibration period is configured or it does not overlap the record.
"""
shared = getattr(worker, "get_calibration_slice", None)
if callable(shared):
return shared()

cal_period = worker._cfg(
"CALIBRATION_PERIOD", worker._cfg("EXPERIMENT_CALIBRATION_PERIOD", "")
)
if not cal_period or getattr(worker, "_time_index", None) is None:
return None
try:
dates = [d.strip() for d in str(cal_period).split(",")]
if len(dates) < 2:
return None
start_date = pd.Timestamp(dates[0])
end_date = pd.Timestamp(dates[1])

steps_fn = getattr(worker, "warmup_steps", None)
steps = steps_fn() if callable(steps_fn) else worker.warmup_days

after_warmup = worker._time_index[steps:]
if not isinstance(after_warmup, pd.DatetimeIndex):
after_warmup = pd.DatetimeIndex(after_warmup)

mask = (after_warmup >= start_date) & (after_warmup <= end_date)
hits = np.where(mask)[0]
if len(hits) == 0:
return None
return int(hits[0]), int(hits[-1] + 1)
except (ValueError, TypeError):
return None


class HecHmsWorker(InMemoryModelWorker):
"""Worker for HEC-HMS model calibration.

Expand Down Expand Up @@ -302,7 +350,7 @@ def compute_gradient(
if self._time_index is not None and len(self._time_index) > 0:
doy_start = self._time_index[0].timetuple().tm_yday

cal_slice = self.get_calibration_slice()
cal_slice = _calibration_slice(self)

def loss_fn(params_array, param_names):
params_dict = dict(zip(param_names, params_array))
Expand Down Expand Up @@ -366,7 +414,7 @@ def evaluate_with_gradient(
if self._time_index is not None and len(self._time_index) > 0:
doy_start = self._time_index[0].timetuple().tm_yday

cal_slice = self.get_calibration_slice()
cal_slice = _calibration_slice(self)

def loss_fn(params_array, param_names):
params_dict = dict(zip(param_names, params_array))
Expand Down
90 changes: 90 additions & 0 deletions tests/test_calibration_slice_fallback.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,90 @@
"""The calibration slice must work on symfluence releases without the shared helper.

``InMemoryModelWorker.get_calibration_slice()`` landed in symfluence after
the 0.9.2 release on PyPI. This package calls it from the differentiable
loss path, where an AttributeError would be swallowed by the surrounding
``except Exception`` in compute_gradient / evaluate_with_gradient — so an
older core would not raise, it would quietly stop restricting the loss to
the calibration period. That is the bug this whole change exists to fix.

These tests pin that the fallback produces the same window as the shared
implementation, and degrades to None rather than to "score everything".
"""

import pandas as pd
import pytest

import symfluence.optimization.workers.inmemory_worker as _iw

from jhechms.calibration.worker import _calibration_slice

WARMUP = 365
PERIOD = "2003-06-01, 2003-08-31"


class _Worker:
"""Only the attributes the slice logic reads."""

def __init__(self, idx, period=PERIOD, warmup=WARMUP):
self._time_index = idx
self._cfg_map = {"CALIBRATION_PERIOD": period}
self._warmup = warmup

def _cfg(self, key, default=None):
return self._cfg_map.get(key, default)

def warmup_steps(self):
return self._warmup

@property
def warmup_days(self):
return self._warmup


@pytest.fixture
def idx():
return pd.date_range("2002-01-01", periods=365 * 3, freq="D")


@pytest.fixture
def without_shared_helper(monkeypatch):
"""Simulate a symfluence predating get_calibration_slice."""
if hasattr(_iw.InMemoryModelWorker, "get_calibration_slice"):
monkeypatch.delattr(_iw.InMemoryModelWorker, "get_calibration_slice")


def _with_shared(worker):
"""Bind the real shared implementation, when this symfluence has one."""
shared = getattr(_iw.InMemoryModelWorker, "get_calibration_slice", None)
if shared is not None:
worker.get_calibration_slice = shared.__get__(worker)
return worker


def test_fallback_selects_the_calibration_window(idx, without_shared_helper):
start, end = _calibration_slice(_Worker(idx))
after_warmup = idx[WARMUP:]
assert after_warmup[start] == pd.Timestamp("2003-06-01")
assert after_warmup[end - 1] == pd.Timestamp("2003-08-31")
assert end - start == 92 # Jun 1 -> Aug 31 inclusive


def test_fallback_matches_the_shared_implementation(idx):
"""Both paths must agree, or upgrading symfluence would move results."""
shared = getattr(_iw.InMemoryModelWorker, "get_calibration_slice", None)
if shared is None:
pytest.skip("this symfluence has no shared implementation to compare against")
assert _calibration_slice(_with_shared(_Worker(idx))) == _calibration_slice(_Worker(idx))


@pytest.mark.parametrize("period", ["", "not-a-date", "2003-06-01"])
def test_unusable_period_returns_none(idx, period, without_shared_helper):
assert _calibration_slice(_Worker(idx, period=period)) is None


def test_non_overlapping_period_returns_none(idx, without_shared_helper):
assert _calibration_slice(_Worker(idx, period="1999-01-01, 1999-02-01")) is None


def test_missing_time_index_returns_none(without_shared_helper):
assert _calibration_slice(_Worker(None)) is None
Loading