diff --git a/forecast_interface/__init__.py b/forecast_interface/__init__.py index f303f40..ba93093 100644 --- a/forecast_interface/__init__.py +++ b/forecast_interface/__init__.py @@ -1,4 +1,4 @@ -__version__ = "0.1.17" +__version__ = "0.1.18" from .common import AggregationMethod from .input import ( @@ -26,6 +26,7 @@ ModelResult, ModelSuccess, RetrainableModel, + RunConfig, TrainedArtifact, ) from .output import ( @@ -65,6 +66,7 @@ "PastKnownVariable", "QuantileData", "RetrainableModel", + "RunConfig", "SpatialInputs", "SpatialInputSpec", "SpatialRepresentation", diff --git a/forecast_interface/interface/__init__.py b/forecast_interface/interface/__init__.py index 4b33dd2..a72d70d 100644 --- a/forecast_interface/interface/__init__.py +++ b/forecast_interface/interface/__init__.py @@ -2,6 +2,7 @@ from .failure import FailureCause from .protocol import BatchHindcastModel, ForecastModel, RetrainableModel from .result import ModelFailure, ModelResult, ModelSuccess +from .run_config import RunConfig from .scope import ArtifactScope __all__ = [ @@ -13,5 +14,6 @@ "ModelResult", "ModelSuccess", "RetrainableModel", + "RunConfig", "TrainedArtifact", ] diff --git a/forecast_interface/interface/protocol.py b/forecast_interface/interface/protocol.py index 03e9d2d..74d34cf 100644 --- a/forecast_interface/interface/protocol.py +++ b/forecast_interface/interface/protocol.py @@ -1,4 +1,4 @@ -from collections.abc import Sequence +from collections.abc import Mapping, Sequence from datetime import datetime from random import Random from typing import Any, Protocol, runtime_checkable @@ -18,13 +18,12 @@ def input_requirement(self) -> InputRequirement: ... artifact_scope: ArtifactScope # REQUIRED training contract — cold full rebuild is the baseline. + # `config` is an opaque mapping each model self-validates; `RunConfig` + # is FI's generic cross-model subset a model may parse out of it. def train( - self, inputs: ModelInputs, *, config: Any, rng: Random + self, inputs: ModelInputs, *, config: Mapping[str, Any], rng: Random ) -> TrainedArtifact: ... - # ^ PROVISIONAL: `config` model params are co-designed with SAP3 (Q8). - # Typed Any until that contract lands. - def predict( self, artifact: TrainedArtifact, @@ -57,11 +56,13 @@ def hindcast( class RetrainableModel(ForecastModel, Protocol): # Warm-start retrain — OPTIONAL. SAP3 checks isinstance(model, RetrainableModel) # to know whether warm-start is supported; otherwise it falls back to `train`. + # `config` is an opaque mapping each model self-validates; `RunConfig` + # is FI's generic cross-model subset a model may parse out of it. def retrain( self, base_artifact: TrainedArtifact, inputs: ModelInputs, *, - config: Any, # PROVISIONAL: model params are co-designed with SAP3 (Q8). + config: Mapping[str, Any], rng: Random, ) -> TrainedArtifact: ... diff --git a/forecast_interface/interface/run_config.py b/forecast_interface/interface/run_config.py new file mode 100644 index 0000000..086816c --- /dev/null +++ b/forecast_interface/interface/run_config.py @@ -0,0 +1,26 @@ +from pydantic import BaseModel, Field, field_validator + + +class RunConfig(BaseModel): + """Runtime sampling counts for forecast generation and emission. + + num_samples is aleatoric draws per weight. num_weight_samples is epistemic + weight draws. The pooled num_weight_samples * num_samples draws are the full + predictive distribution. num_trajectories is how many raw paths to emit + (<= pool, 0 = none): a retention count, not a generation count. + """ + + quantile_levels: list[float] | None = None + num_weight_samples: int | None = Field(default=None, gt=0) + num_trajectories: int | None = Field(default=None, ge=0) + num_samples: int | None = Field(default=None, gt=0) + + @field_validator("quantile_levels") + @classmethod + def _validate_levels(cls, v: list[float] | None) -> list[float] | None: + if v is None: + return None + for level in v: + if not (0 < level < 1): + raise ValueError(f"quantile levels must be in (0, 1), got {level}") + return v diff --git a/pyproject.toml b/pyproject.toml index e2e9cb5..51edc08 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "forecastinterface" -version = "0.1.17" +version = "0.1.18" description = "Add your description here" readme = "README.md" requires-python = ">=3.11" @@ -25,7 +25,7 @@ init_typed = true warn_required_dynamic_aliases = true [tool.bumpversion] -current_version = "0.1.17" +current_version = "0.1.18" commit = false tag = false allow_dirty = true diff --git a/tests/test_run_config.py b/tests/test_run_config.py new file mode 100644 index 0000000..31c0eb6 --- /dev/null +++ b/tests/test_run_config.py @@ -0,0 +1,61 @@ +from pydantic import ValidationError +import pytest + +from forecast_interface import RunConfig + + +class TestRunConfig: + def test_empty_config_valid(self) -> None: + config = RunConfig() + + assert config.quantile_levels is None + assert config.num_weight_samples is None + assert config.num_trajectories is None + assert config.num_samples is None + + def test_four_fields_validate(self) -> None: + config = RunConfig( + quantile_levels=[0.1, 0.5, 0.9], + num_trajectories=50, + num_samples=100, + num_weight_samples=5, + ) + + assert config.quantile_levels == [0.1, 0.5, 0.9] + assert config.num_weight_samples == 5 + assert config.num_trajectories == 50 + assert config.num_samples == 100 + + @pytest.mark.parametrize( + "quantile_levels", + ([0.1, 0.5, 1.0], [0.0, 0.5, 0.9], [-0.1, 0.5, 0.9]), + ) + def test_quantile_level_out_of_range_raises( + self, + quantile_levels: list[float], + ) -> None: + with pytest.raises(ValidationError, match="quantile levels must be in"): + RunConfig(quantile_levels=quantile_levels) + + def test_negative_counts_raise(self) -> None: + with pytest.raises(ValidationError): + RunConfig(num_trajectories=-1) + + with pytest.raises(ValidationError): + RunConfig(num_samples=0) + + def test_num_weight_samples_validates(self) -> None: + config = RunConfig(num_weight_samples=3) + + assert config.num_weight_samples == 3 + + with pytest.raises(ValidationError): + RunConfig(num_weight_samples=0) + + def test_num_trajectories_zero_allowed(self) -> None: + config = RunConfig(num_trajectories=0) + + assert config.num_trajectories == 0 + + with pytest.raises(ValidationError): + RunConfig(num_trajectories=-1) diff --git a/uv.lock b/uv.lock index c53dd3d..03ddebc 100644 --- a/uv.lock +++ b/uv.lock @@ -129,7 +129,7 @@ wheels = [ [[package]] name = "forecastinterface" -version = "0.1.17" +version = "0.1.18" source = { virtual = "." } dependencies = [ { name = "polars" },