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
4 changes: 3 additions & 1 deletion forecast_interface/__init__.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
__version__ = "0.1.17"
__version__ = "0.1.18"

from .common import AggregationMethod
from .input import (
Expand Down Expand Up @@ -26,6 +26,7 @@
ModelResult,
ModelSuccess,
RetrainableModel,
RunConfig,
TrainedArtifact,
)
from .output import (
Expand Down Expand Up @@ -65,6 +66,7 @@
"PastKnownVariable",
"QuantileData",
"RetrainableModel",
"RunConfig",
"SpatialInputs",
"SpatialInputSpec",
"SpatialRepresentation",
Expand Down
2 changes: 2 additions & 0 deletions forecast_interface/interface/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__ = [
Expand All @@ -13,5 +14,6 @@
"ModelResult",
"ModelSuccess",
"RetrainableModel",
"RunConfig",
"TrainedArtifact",
]
13 changes: 7 additions & 6 deletions forecast_interface/interface/protocol.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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: ...
26 changes: 26 additions & 0 deletions forecast_interface/interface/run_config.py
Original file line number Diff line number Diff line change
@@ -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
4 changes: 2 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
@@ -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"
Expand All @@ -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
Expand Down
61 changes: 61 additions & 0 deletions tests/test_run_config.py
Original file line number Diff line number Diff line change
@@ -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)
2 changes: 1 addition & 1 deletion uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading