CellTriage / src /utils /config.py
Sarvarbek13's picture
Sync with repository cleanup: drop build-tooling docs, new root markers
d92710f verified
Raw
History Blame Contribute Delete
42.7 kB
"""Typed, validated configuration loading for CellTriage.
WHAT: Pydantic schemas for all five YAML configs, plus loaders that parse and
validate them, and cross-config consistency checks.
WHY a schema rather than ``yaml.safe_load``: this project encodes its
methodological commitments as configuration. Its hard constraints are
expressible as invariants over these files -- escape cost must
dominate overkill cost, feature selection must be fitted inside CV folds,
neural networks are forbidden, and every excluded cell must carry a reason. A
plain dict load would let any of those be violated by a typo and the run would
proceed, producing numbers that look fine and are not. Validation turns each
commitment into something that must be *deliberately* defeated rather than
accidentally lost.
The validators here are therefore not defensive boilerplate; each one
corresponds to a stated project rule and is annotated with which.
"""
from __future__ import annotations
from pathlib import Path
from typing import Any, Literal
import yaml
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from src.utils.paths import CONFIG_DIR
# ---------------------------------------------------------------------------
# Base
# ---------------------------------------------------------------------------
class _StrictModel(BaseModel):
"""Base model that forbids unknown keys.
WHY ``extra='forbid'``: a mistyped key in a YAML file is otherwise silently
ignored, so the pipeline runs with the default rather than the intended
value. For a file that holds cost assumptions and leakage controls, a typo
that changes behaviour without complaining is the worst possible failure
mode.
"""
model_config = ConfigDict(extra="forbid", frozen=False)
# ---------------------------------------------------------------------------
# grading.yaml
# ---------------------------------------------------------------------------
class GradeSpec(_StrictModel):
"""One product tier."""
min_cycles: int = Field(ge=0, description="Inclusive lower bound on observed cycle life.")
tier: str = Field(min_length=1, description="Human-readable application tier.")
class LabellingSpec(_StrictModel):
boundary_inclusive: bool
class GradingConfig(_StrictModel):
"""Grade boundaries and product tiers (``configs/grading.yaml``)."""
grades: dict[str, GradeSpec]
warranty_target_cycles: int = Field(ge=0)
grade_order: list[str] = Field(min_length=2)
shipped_grades: list[str] = Field(min_length=1)
labelling: LabellingSpec
@model_validator(mode="after")
def _check_grade_structure(self) -> "GradingConfig":
if set(self.grade_order) != set(self.grades):
raise ValueError(
f"grade_order {self.grade_order} does not match the grades defined "
f"({sorted(self.grades)}). Every grade must appear exactly once."
)
if len(set(self.grade_order)) != len(self.grade_order):
raise ValueError(f"grade_order contains duplicates: {self.grade_order}")
# Ordering is load-bearing: it determines which off-diagonal cost-matrix
# entries count as escapes. If the order were not strictly decreasing in
# min_cycles, "more demanding tier" would be undefined.
thresholds = [self.grades[g].min_cycles for g in self.grade_order]
if thresholds != sorted(thresholds, reverse=True) or len(set(thresholds)) != len(thresholds):
raise ValueError(
f"grade_order must be strictly decreasing in min_cycles; got "
f"{dict(zip(self.grade_order, thresholds))}."
)
if thresholds[-1] != 0:
raise ValueError(
f"The least demanding grade ('{self.grade_order[-1]}') must have "
f"min_cycles == 0 so that every cell receives a grade; got {thresholds[-1]}."
)
unknown_shipped = set(self.shipped_grades) - set(self.grade_order)
if unknown_shipped:
raise ValueError(f"shipped_grades references undefined grades: {sorted(unknown_shipped)}")
if self.grade_order[-1] in self.shipped_grades:
raise ValueError(
f"The scrap grade '{self.grade_order[-1]}' cannot be a shipped grade; "
"escape rate is defined over shipped cells only."
)
highest, lowest_shipped = thresholds[0], self.grades[self.grade_order[-2]].min_cycles
if not lowest_shipped <= self.warranty_target_cycles <= highest:
raise ValueError(
f"warranty_target_cycles ({self.warranty_target_cycles}) must lie between the "
f"lowest shipped-grade boundary ({lowest_shipped}) and the highest "
f"({highest})."
)
return self
def grade_for(self, cycle_life: float) -> str:
"""Assign a grade to an observed cycle life.
Boundary handling is explicit (see ``labelling.boundary_inclusive``):
a value exactly on a boundary goes to the higher tier. An off-by-one
here silently moves cells between classes and makes the reported class
balance irreproducible.
"""
for grade in self.grade_order:
threshold = self.grades[grade].min_cycles
if (cycle_life >= threshold) if self.labelling.boundary_inclusive else (cycle_life > threshold):
return grade
return self.grade_order[-1]
# ---------------------------------------------------------------------------
# costs.yaml
# ---------------------------------------------------------------------------
class SensitivitySpec(_StrictModel):
escape_overkill_ratios: list[float] = Field(min_length=2)
@field_validator("escape_overkill_ratios")
@classmethod
def _ascending_and_positive(cls, v: list[float]) -> list[float]:
if any(r <= 0 for r in v):
raise ValueError(f"escape_overkill_ratios must all be positive; got {v}")
if v != sorted(v):
raise ValueError(f"escape_overkill_ratios must be ascending; got {v}")
return v
class UnitsSpec(_StrictModel):
basis: str
currency: str | None = None
class CostsConfig(_StrictModel):
"""Cost matrix and chamber economics (``configs/costs.yaml``)."""
cost_matrix: dict[str, list[float]]
cycle_test_cost: float = Field(gt=0)
chamber_slots: int = Field(gt=0)
batch_size: int = Field(gt=0)
sensitivity: SensitivitySpec
units: UnitsSpec
@model_validator(mode="after")
def _check_matrix(self) -> "CostsConfig":
n = len(self.cost_matrix)
for row_key, row in self.cost_matrix.items():
if not row_key.startswith("true_"):
raise ValueError(f"Cost-matrix row keys must be 'true_<GRADE>'; got '{row_key}'.")
if len(row) != n:
raise ValueError(
f"Cost matrix must be square: row '{row_key}' has {len(row)} entries, expected {n}."
)
if any(c < 0 for c in row):
raise ValueError(f"Cost-matrix entries must be non-negative; row '{row_key}' = {row}.")
rows = self.row_grades()
for i, grade in enumerate(rows):
diagonal = self.cost_matrix[f"true_{grade}"][i]
if diagonal != 0.0:
raise ValueError(
f"Diagonal entry for grade '{grade}' must be 0 (a correct assignment costs "
f"nothing); got {diagonal}."
)
escapes = self.escape_costs()
overkills = self.overkill_costs()
if not escapes or not overkills:
raise ValueError("Cost matrix must contain at least one escape and one overkill entry.")
# PROJECT RULE: escape cost must dominate overkill cost. A pack's life
# is governed by its weakest cell, so an escaped cell carries warranty
# and recall exposure while an overkilled cell costs only its
# manufacturing value. If this inversion were ever introduced by an
# edit, every decision-theoretic conclusion in the project would flip
# without any other symptom.
if min(escapes.values()) <= max(overkills.values()):
raise ValueError(
f"Every escape cost must exceed every overkill cost. "
f"Cheapest escape {min(escapes.values())} ({min(escapes, key=escapes.get)}) does not "
f"exceed costliest overkill {max(overkills.values())} "
f"({max(overkills, key=overkills.get)})."
)
# Capacity must actually bind, otherwise Phase 8D's allocation problem
# is vacuous: if every cell can be held, there is nothing to allocate.
if self.chamber_slots >= self.batch_size:
raise ValueError(
f"chamber_slots ({self.chamber_slots}) must be smaller than batch_size "
f"({self.batch_size}); otherwise the capacity constraint does not bind and the "
"allocation problem in Phase 8D is trivial."
)
return self
def row_grades(self) -> list[str]:
"""Grade labels in matrix row order, most demanding first."""
return [key.removeprefix("true_") for key in self.cost_matrix]
def _off_diagonal(self, kind: Literal["escape", "overkill"]) -> dict[tuple[str, str], float]:
"""Return off-diagonal entries of the requested kind, keyed (true, assigned).
Rows are true grades and columns assigned grades, both ordered most
demanding first. A cell assigned to a MORE demanding tier than it can
sustain is an escape, which means ``col_index < row_index`` -- the
lower-left triangle. See the orientation note in ``configs/costs.yaml``:
the build plan's prose says "above the diagonal" but its own numbers
place escapes below it, and the numbers govern.
"""
grades = self.row_grades()
out: dict[tuple[str, str], float] = {}
for i, true_grade in enumerate(grades):
row = self.cost_matrix[f"true_{true_grade}"]
for j, assigned_grade in enumerate(grades):
if i == j:
continue
is_escape = j < i
if (kind == "escape") == is_escape:
out[(true_grade, assigned_grade)] = row[j]
return out
def escape_costs(self) -> dict[tuple[str, str], float]:
"""Costs of assigning a cell to a tier it cannot sustain."""
return self._off_diagonal("escape")
def overkill_costs(self) -> dict[tuple[str, str], float]:
"""Costs of downgrading or scrapping a capable cell."""
return self._off_diagonal("overkill")
def escape_overkill_ratio(self) -> float:
"""Mean escape cost divided by mean overkill cost, for the sensitivity sweep."""
escapes = list(self.escape_costs().values())
overkills = list(self.overkill_costs().values())
return (sum(escapes) / len(escapes)) / (sum(overkills) / len(overkills))
# ---------------------------------------------------------------------------
# data.yaml
# ---------------------------------------------------------------------------
class DatasetSpec(_StrictModel):
name: str
chemistry: str
cell_model: str
nominal_capacity_ah: float = Field(gt=0)
nominal_voltage_v: float = Field(gt=0)
cycling_temperature_c: float
discharge_protocol: str
eol_capacity_fraction: float = Field(gt=0, lt=1)
class SourceFileSpec(_StrictModel):
"""One downloadable source file, with the provenance needed to re-fetch it."""
name: str = Field(min_length=1)
url: str = Field(pattern=r"^https://")
size_bytes: int = Field(gt=0)
batch: str = Field(min_length=1)
#: Populated by src/data/downloader.py from the downloaded bytes. Null until
#: the file has actually been fetched and hashed.
sha256: str | None = Field(default=None, pattern=r"^[0-9a-f]{64}$")
class SourceSpec(_StrictModel):
verified: bool
primary_citation: str
extension_citation: str
files: list[SourceFileSpec] = Field(default_factory=list)
format: str
verified_on: str | None = None
landing_page: str | None = None
total_size_bytes: int | None = None
extension_located: bool = False
@model_validator(mode="after")
def _verified_requires_evidence(self) -> "SourceSpec":
# Task 2.0 rule: do not hardcode a URL you have not verified. Flipping
# `verified` to true without a digest for every file would defeat the
# check this exists to enforce -- a URL that merely resolves is not a
# URL whose contents have been verified.
if self.verified:
if not self.files:
raise ValueError("source.verified is true but source.files is empty.")
undigested = [f.name for f in self.files if not f.sha256]
if undigested:
raise ValueError(
f"source.verified is true but these files have no sha256 digest: {undigested}. "
"A digest can only come from bytes that were actually downloaded."
)
if self.total_size_bytes is not None and self.files:
declared = sum(f.size_bytes for f in self.files)
if declared != self.total_size_bytes:
raise ValueError(
f"source.total_size_bytes ({self.total_size_bytes}) does not equal the sum of the "
f"per-file sizes ({declared})."
)
return self
class PathsSpec(_StrictModel):
raw: str
interim: str
processed: str
cells: str
class BatchSpec(_StrictModel):
date: str
role: Literal["in_distribution", "out_of_distribution"]
expected_cells_raw: int = Field(gt=0)
expected_cells_used: int = Field(ge=0)
verified: bool
note: str | None = None
@model_validator(mode="after")
def _used_not_more_than_raw(self) -> "BatchSpec":
if self.expected_cells_used > self.expected_cells_raw:
raise ValueError(
f"expected_cells_used ({self.expected_cells_used}) exceeds expected_cells_raw "
f"({self.expected_cells_raw})."
)
return self
class ContinuationSpec(_StrictModel):
primary: str
continued_as: str
known_extra_cycles: int = Field(gt=0)
verified: bool
@model_validator(mode="after")
def _distinct_records(self) -> "ContinuationSpec":
if self.primary == self.continued_as:
raise ValueError(f"A cell cannot continue as itself: '{self.primary}'.")
return self
class ExclusionSpec(_StrictModel):
"""One excluded cell.
``reason`` is required and must be substantive. PROJECT RULE: never
silently drop a cell -- always log the cell ID and the reason. Making the
reason a required, minimum-length field means an undocumented exclusion
cannot be added to this file at all.
"""
cell_id: str = Field(min_length=1)
batch: str = Field(min_length=1)
reason: str = Field(min_length=15)
verified: bool
@field_validator("reason")
@classmethod
def _reason_is_substantive(cls, v: str) -> str:
placeholders = {"tbd", "todo", "n/a", "na", "unknown", "excluded", "see above", "-"}
if v.strip().lower().rstrip(".") in placeholders:
raise ValueError(f"Exclusion reason '{v}' is a placeholder, not a reason.")
return v
class ReconciliationSpec(_StrictModel):
expected_cells_used: int = Field(gt=0)
expected_cells_used_with_extension: int = Field(gt=0)
verified: bool
@model_validator(mode="after")
def _extension_is_larger(self) -> "ReconciliationSpec":
if self.expected_cells_used_with_extension < self.expected_cells_used:
raise ValueError("The extended corpus cannot contain fewer cells than the primary one.")
return self
class ParsingSpec(_StrictModel):
"""Parsing and storage policy for Phase 2."""
within_cycle_max_cycle: int = Field(gt=0)
parquet_compression: Literal["zstd", "snappy", "gzip", "brotli", "none"]
class CycleLifeRecomputeSpec(_StrictModel):
method: Literal["first_below_threshold"]
smoothing_window: int = Field(ge=1)
min_cycle: int = Field(ge=1)
#: Fraction above the end-of-life threshold still counted as having reached
#: it. Derived empirically in Phase 2, not assumed -- see the extensive
#: rationale in configs/data.yaml. Capped well below 10% because a large
#: tolerance would start absorbing genuinely censored cells.
eol_tolerance_fraction: float = Field(ge=0.0, le=0.05)
class PlaceholderRowSpec(_StrictModel):
expect_at_most_one: bool
expect_at_first_cycle: bool
#: Batches empirically confirmed to carry the all-zero placeholder row.
batches_with_placeholder: list[str] = Field(default_factory=list)
class ValidationSpec(_StrictModel):
capacity_max_fraction_of_nominal: float = Field(gt=1.0)
voltage_range_v: list[float] = Field(min_length=2, max_length=2)
protocol_voltage_range_v: list[float] = Field(min_length=2, max_length=2)
temperature_range_c: list[float] = Field(min_length=2, max_length=2)
require_monotonic_cycle_index: bool
forbid_duplicate_cycle_index: bool
cycle_life_disagreement_tolerance: int = Field(ge=0)
cycle_life_recompute: CycleLifeRecomputeSpec
placeholder_row: PlaceholderRowSpec
systematic_violation_fraction: float = Field(gt=0.0, le=0.5)
systematic_cell_fraction: float = Field(gt=0.0, le=1.0)
max_cycle_duration_minutes: float = Field(gt=0.0)
@model_validator(mode="after")
def _validation_bound_contains_protocol(self) -> "ValidationSpec":
# The validation bound exists to catch mis-parsing, so it must be at
# least as wide as the protocol setpoints. A bound TIGHTER than the
# protocol would flag correctly parsed data as impossible.
if (self.voltage_range_v[0] > self.protocol_voltage_range_v[0]
or self.voltage_range_v[1] < self.protocol_voltage_range_v[1]):
raise ValueError(
f"voltage_range_v {self.voltage_range_v} must contain the protocol range "
f"{self.protocol_voltage_range_v}; a bound tighter than the protocol would reject "
"correctly parsed data."
)
return self
@field_validator("voltage_range_v", "temperature_range_c")
@classmethod
def _ordered_range(cls, v: list[float]) -> list[float]:
if v[0] >= v[1]:
raise ValueError(f"Range must be [low, high] with low < high; got {v}.")
return v
class DataConfig(_StrictModel):
"""Dataset structure, joins and exclusions (``configs/data.yaml``)."""
dataset: DatasetSpec
source: SourceSpec
paths: PathsSpec
batches: dict[str, BatchSpec]
continuations: list[ContinuationSpec]
exclusions: list[ExclusionSpec]
reconciliation: ReconciliationSpec
parsing: ParsingSpec
validation: ValidationSpec
@model_validator(mode="after")
def _check_cross_references(self) -> "DataConfig":
known_batches = set(self.batches)
unknown = {e.cell_id: e.batch for e in self.exclusions if e.batch not in known_batches}
if unknown:
raise ValueError(f"Exclusions reference undefined batches: {unknown}")
ids = [e.cell_id for e in self.exclusions]
duplicates = {c for c in ids if ids.count(c) > 1}
if duplicates:
raise ValueError(f"Duplicate cell IDs in the exclusion list: {sorted(duplicates)}")
# A cell cannot be both joined into another record and excluded: the
# two operations would disagree about whether it exists.
continuation_ids = {c.primary for c in self.continuations} | {
c.continued_as for c in self.continuations
}
conflict = continuation_ids & set(ids)
if conflict:
raise ValueError(
f"Cells appear in both the continuation joins and the exclusion list: {sorted(conflict)}. "
"A cell cannot be simultaneously joined and dropped."
)
if not any(b.role == "out_of_distribution" for b in self.batches.values()):
raise ValueError(
"At least one batch must be marked out_of_distribution; RQ4 has no test set otherwise."
)
return self
def excluded_cell_ids(self) -> set[str]:
return {e.cell_id for e in self.exclusions}
def unverified_assumptions(self) -> list[str]:
"""List every config entry still carrying ``verified: false``.
WHY: Phase 2 must close each of these out. Surfacing them as a list
means the outstanding assumptions are printable rather than buried in
the YAML.
"""
pending: list[str] = []
if not self.source.verified:
pending.append("source (URLs, checksums, file sizes)")
pending.extend(f"batches.{name} (cell counts)" for name, b in self.batches.items() if not b.verified)
pending.extend(
f"continuations.{c.primary}->{c.continued_as} (extra-cycle count)"
for c in self.continuations
if not c.verified
)
pending.extend(f"exclusions.{e.cell_id} (stated reason)" for e in self.exclusions if not e.verified)
if not self.reconciliation.verified:
pending.append("reconciliation (final cell count)")
return pending
# ---------------------------------------------------------------------------
# features.yaml
# ---------------------------------------------------------------------------
class CurveSpec(_StrictModel):
voltage_min_v: float
voltage_max_v: float
n_grid_points: int = Field(gt=1)
baseline_cycle: int = Field(gt=0)
min_baseline_cycle: int = Field(gt=0)
interpolation: str
@model_validator(mode="after")
def _check_grid(self) -> "CurveSpec":
if self.voltage_min_v >= self.voltage_max_v:
raise ValueError(
f"voltage_min_v ({self.voltage_min_v}) must be below voltage_max_v ({self.voltage_max_v})."
)
if self.min_baseline_cycle > self.baseline_cycle:
raise ValueError(
f"min_baseline_cycle ({self.min_baseline_cycle}) cannot exceed baseline_cycle "
f"({self.baseline_cycle})."
)
return self
class FeatureGroupSpec(_StrictModel):
enabled: bool
module: str
in_line_measurable: bool
is_process_recipe: bool
rationale: str = Field(min_length=20)
class MutualInformationSpec(_StrictModel):
enabled: bool
n_neighbors: int = Field(gt=0)
class StabilitySelectionSpec(_StrictModel):
enabled: bool
n_bootstrap: int = Field(gt=0)
sample_fraction: float = Field(gt=0, le=1)
selection_frequency_threshold: float = Field(gt=0, le=1)
class SelectionSpec(_StrictModel):
fit_inside_cv_only: bool
variance_threshold: float = Field(ge=0)
correlation_threshold: float = Field(gt=0, le=1)
mutual_information: MutualInformationSpec
stability_selection: StabilitySelectionSpec
@field_validator("fit_inside_cv_only")
@classmethod
def _must_fit_inside_cv(cls, v: bool) -> bool:
# HARD CONSTRAINT: with n ~ 124-169 cells, fitting selection on the full
# dataset is the fastest way to produce a fake result. This validator
# exists so that disabling the rule requires deliberately editing a
# schema, not just flipping a YAML flag.
if not v:
raise ValueError(
"selection.fit_inside_cv_only must be true. Feature selection fitted outside "
"cross-validation folds leaks the evaluation set into model construction."
)
return v
class PreprocessingSpec(_StrictModel):
scaler: str
impute_strategy: str
class FeatureOutputSpec(_StrictModel):
directory: str
filename_template: str
float_precision: str
class FeaturesConfig(_StrictModel):
"""Budgets, feature groups and selection rules (``configs/features.yaml``)."""
budgets: list[int] = Field(min_length=1)
reference_budget: int = Field(gt=0)
curve: CurveSpec
groups: dict[str, FeatureGroupSpec]
feature_sets: dict[str, list[str]]
selection: SelectionSpec
preprocessing: PreprocessingSpec
output: FeatureOutputSpec
@field_validator("budgets")
@classmethod
def _ascending_unique_positive(cls, v: list[int]) -> list[int]:
if any(b <= 0 for b in v):
raise ValueError(f"Budgets must be positive cycle counts; got {v}.")
if v != sorted(v) or len(set(v)) != len(v):
raise ValueError(f"Budgets must be strictly ascending and unique; got {v}.")
return v
@model_validator(mode="after")
def _check_consistency(self) -> "FeaturesConfig":
if self.reference_budget not in self.budgets:
raise ValueError(
f"reference_budget ({self.reference_budget}) must be one of the configured budgets "
f"{self.budgets}."
)
for set_name, members in self.feature_sets.items():
unknown = set(members) - set(self.groups)
if unknown:
raise ValueError(f"feature_sets.{set_name} references undefined groups: {sorted(unknown)}")
if len(set(members)) != len(members):
raise ValueError(f"feature_sets.{set_name} contains duplicate groups.")
# RQ4 control. Reporting results with and without process-recipe
# descriptors is required, so a feature set free of them must exist.
recipe_groups = {name for name, g in self.groups.items() if g.is_process_recipe}
if recipe_groups:
if not any(recipe_groups.isdisjoint(m) for m in self.feature_sets.values()):
raise ValueError(
f"No feature set excludes the process-recipe groups {sorted(recipe_groups)}. "
"RQ4 requires results reported both with and without them."
)
# The baseline cycle for DeltaQ(V) must be reachable at the smallest
# budget, otherwise the curve features are undefined there and the
# smallest budget silently yields an all-NaN column.
if self.curve.min_baseline_cycle > min(self.budgets):
raise ValueError(
f"curve.min_baseline_cycle ({self.curve.min_baseline_cycle}) exceeds the smallest "
f"budget ({min(self.budgets)}); DeltaQ(V) features would be undefined there."
)
return self
def enabled_groups(self) -> list[str]:
return [name for name, group in self.groups.items() if group.enabled]
# ---------------------------------------------------------------------------
# models.yaml
# ---------------------------------------------------------------------------
class ModelConstraintsSpec(_StrictModel):
allow_neural_networks: bool
@field_validator("allow_neural_networks")
@classmethod
def _no_neural_networks(cls, v: bool) -> bool:
# A methodological commitment of the project, made structural so it
# cannot be relaxed by editing a config file.
if v:
raise ValueError(
"allow_neural_networks must be false. This project is committed to classical "
"supervised ML: feature engineering rather than estimator capacity is the "
"bottleneck, and exact TreeSHAP attributions are what make a scrap decision "
"auditable by a process engineer."
)
return v
class TargetSpec(_StrictModel):
name: str
transform: Literal["log10", "log", "none"]
class CVFoldSpec(_StrictModel):
n_splits: int = Field(ge=2)
n_repeats: int = Field(ge=1)
purpose: str | None = None
class CVSpec(_StrictModel):
group_key: str
outer: CVFoldSpec
inner: CVFoldSpec
class SplitSchemeSpec(_StrictModel):
enabled: bool
role: str
train_batches: list[str] | None = None
test_batches: list[str] | None = None
n_folds: int | None = Field(default=None, ge=2)
@model_validator(mode="after")
def _no_batch_overlap(self) -> "SplitSchemeSpec":
# HARD CONSTRAINT: no cell in two splits. At batch granularity this is
# the strongest check available before the data is parsed.
if self.train_batches and self.test_batches:
overlap = set(self.train_batches) & set(self.test_batches)
if overlap:
raise ValueError(f"Batches appear in both train and test: {sorted(overlap)}")
return self
class SeedsSpec(_StrictModel):
global_: int = Field(alias="global")
numpy: int
outer_repeats: list[int]
bootstrap: int
model_config = ConfigDict(extra="forbid", populate_by_name=True)
@field_validator("outer_repeats")
@classmethod
def _unique_seeds(cls, v: list[int]) -> list[int]:
if len(set(v)) != len(v):
raise ValueError(f"Outer-repeat seeds must be unique, else repeats are not independent; got {v}.")
return v
class BootstrapSpec(_StrictModel):
n_resamples: int = Field(ge=1000)
confidence_level: float = Field(gt=0, lt=1)
method: str
class MetricsSpec(_StrictModel):
statistical: list[str] = Field(min_length=1)
manufacturing: list[str] = Field(min_length=1)
primary: str
@model_validator(mode="after")
def _primary_is_declared(self) -> "MetricsSpec":
if self.primary not in (*self.statistical, *self.manufacturing):
raise ValueError(f"metrics.primary '{self.primary}' is not among the declared metrics.")
return self
class SearchDimSpec(_StrictModel):
type: Literal["float", "int", "categorical"]
low: float | None = None
high: float | None = None
log: bool = False
choices: list[Any] | None = None
@model_validator(mode="after")
def _check_bounds(self) -> "SearchDimSpec":
if self.type == "categorical":
if not self.choices:
raise ValueError("A categorical search dimension requires 'choices'.")
return self
if self.low is None or self.high is None:
raise ValueError(f"A {self.type} search dimension requires 'low' and 'high'.")
if self.low >= self.high:
raise ValueError(f"Search bounds must satisfy low < high; got low={self.low}, high={self.high}.")
if self.log and self.low <= 0:
raise ValueError(f"Log-scaled search requires a positive lower bound; got {self.low}.")
return self
class ModelSpec(_StrictModel):
enabled: bool
estimator: str
role: str | None = None
params: dict[str, Any] = Field(default_factory=dict)
search_space: dict[str, SearchDimSpec] = Field(default_factory=dict)
@field_validator("estimator")
@classmethod
def _reject_neural_estimators(cls, v: str) -> str:
# Defence in depth behind the allow_neural_networks flag: catches an
# estimator that would violate the constraint even if the flag were
# somehow bypassed.
forbidden = ("MLP", "torch", "tensorflow", "keras", "neural_network")
if any(token.lower() in v.lower() for token in forbidden):
raise ValueError(f"Estimator '{v}' appears to be a neural network, which this project forbids.")
return v
class HyperparameterSearchSpec(_StrictModel):
engine: str
n_trials: int = Field(gt=0)
timeout_seconds: int = Field(gt=0)
scope: str
sampler: str
seed: int
@field_validator("scope")
@classmethod
def _inner_fold_only(cls, v: str) -> str:
# PROJECT RULE: never tune hyperparameters on the test set.
if v != "inner_fold_only":
raise ValueError(
f"hyperparameter_search.scope must be 'inner_fold_only'; got '{v}'. Tuning against "
"outer-fold performance is tuning on the test set."
)
return v
class QuantileSpec(_StrictModel):
enabled: bool
estimator: str
objective: str
quantiles: list[float] = Field(min_length=2)
@field_validator("quantiles")
@classmethod
def _valid_quantiles(cls, v: list[float]) -> list[float]:
if any(not 0 < q < 1 for q in v):
raise ValueError(f"Quantiles must lie strictly in (0, 1); got {v}.")
if v != sorted(v) or len(set(v)) != len(v):
raise ValueError(f"Quantiles must be strictly ascending and unique; got {v}.")
return v
@model_validator(mode="after")
def _symmetric_for_cqr(self) -> "QuantileSpec":
# Conformalized quantile regression needs symmetric lower/upper pairs
# to form intervals at a nominal level.
for q in self.quantiles:
if not any(abs((1 - q) - other) < 1e-9 for other in self.quantiles):
raise ValueError(
f"Quantile {q} has no complementary quantile {1 - q:.2f}; CQR intervals require "
"symmetric pairs."
)
return self
class SecondaryClassifierSpec(_StrictModel):
enabled: bool
role: str
estimators: list[str] = Field(min_length=1)
class_weight: str | None = None
min_class_count_warning: int = Field(ge=1)
class GradingRouteSpec(_StrictModel):
"""How A/B/C grades are produced (reviewed decision after Phase 3)."""
primary: Literal["ordinal_from_regression", "direct_classification"]
secondary: SecondaryClassifierSpec
@model_validator(mode="after")
def _secondary_is_not_the_headline(self) -> "GradingRouteSpec":
# The secondary route exists for interpretability only. Requiring it to
# declare that in its own role string keeps the distinction visible at
# the point of use rather than only in a design document.
if self.secondary.enabled and "secondary" not in self.secondary.role.lower():
raise ValueError(
"grading_route.secondary.role must state that it is a secondary view; it must "
"never be reported as the headline grading result."
)
return self
class EnsembleSpec(_StrictModel):
enabled: bool
meta_learner: str
use_out_of_fold_only: bool
@field_validator("use_out_of_fold_only")
@classmethod
def _oof_only(cls, v: bool) -> bool:
if not v:
raise ValueError(
"ensemble.use_out_of_fold_only must be true; stacking on in-fold predictions leaks "
"base-model training performance into the meta-learner."
)
return v
class RegistrySpec(_StrictModel):
directory: str
record_provenance: bool
class ModelsConfig(_StrictModel):
"""Model spaces, CV protocol and seeds (``configs/models.yaml``)."""
constraints: ModelConstraintsSpec
target: TargetSpec
cv: CVSpec
splits: dict[str, SplitSchemeSpec]
seeds: SeedsSpec
bootstrap: BootstrapSpec
metrics: MetricsSpec
hyperparameter_search: HyperparameterSearchSpec
models: dict[str, ModelSpec]
quantile: QuantileSpec
grading_route: GradingRouteSpec
ensemble: EnsembleSpec
registry: RegistrySpec
@model_validator(mode="after")
def _check_protocol(self) -> "ModelsConfig":
if len(self.seeds.outer_repeats) != self.cv.outer.n_repeats:
raise ValueError(
f"seeds.outer_repeats has {len(self.seeds.outer_repeats)} entries but "
f"cv.outer.n_repeats is {self.cv.outer.n_repeats}; every repeat must have a fixed seed "
"for the run to be reproducible."
)
if not any(m.role == "floor" for m in self.models.values()):
raise ValueError(
"At least one dummy model must be declared with role 'floor'. Without a floor, a "
"reported error is not interpretable."
)
if not any(m.role == "severson_reproduction_and_baseline" for m in self.models.values()):
raise ValueError(
"No model is designated for the Severson reproduction (Gate 2), which is the "
"credibility anchor of the project."
)
return self
def enabled_models(self) -> list[str]:
return [name for name, spec in self.models.items() if spec.enabled]
# ---------------------------------------------------------------------------
# Bundle and loaders
# ---------------------------------------------------------------------------
CONFIG_SCHEMAS: dict[str, type[BaseModel]] = {
"data": DataConfig,
"features": FeaturesConfig,
"models": ModelsConfig,
"grading": GradingConfig,
"costs": CostsConfig,
}
class ConfigBundle(_StrictModel):
"""All five validated configs, plus cross-file consistency checks.
WHY cross-file validation: the individually valid files can still disagree
with each other -- a cost matrix whose grades do not match the grading
scheme is the obvious case, and it would produce a decision layer that
indexes the wrong tier without ever raising.
"""
data: DataConfig
features: FeaturesConfig
models: ModelsConfig
grading: GradingConfig
costs: CostsConfig
@model_validator(mode="after")
def _check_cross_config(self) -> "ConfigBundle":
if self.costs.row_grades() != self.grading.grade_order:
raise ValueError(
f"Cost-matrix grades {self.costs.row_grades()} must match grading.grade_order "
f"{self.grading.grade_order} in both membership and order; otherwise the decision "
"layer indexes the wrong tier."
)
configured_batches = set(self.data.batches)
for name, scheme in self.models.splits.items():
for field in ("train_batches", "test_batches"):
declared = getattr(scheme, field) or []
unknown = set(declared) - configured_batches
if unknown:
raise ValueError(
f"splits.{name}.{field} references batches not defined in data.yaml: {sorted(unknown)}"
)
ood_batches = {n for n, b in self.data.batches.items() if b.role == "out_of_distribution"}
holdout = self.models.splits.get("batch_holdout")
if holdout and holdout.enabled and set(holdout.test_batches or []) != ood_batches:
raise ValueError(
f"batch_holdout tests on {holdout.test_batches} but data.yaml marks {sorted(ood_batches)} "
"as out_of_distribution. The OOD split must match the declared OOD batches."
)
return self
def unverified_assumptions(self) -> list[str]:
"""Outstanding ``verified: false`` entries that Phase 2 must close out."""
return self.data.unverified_assumptions()
def load_config(name: str, config_dir: Path | None = None) -> BaseModel:
"""Load and validate a single config by stem (e.g. ``"costs"``).
Raises:
KeyError: If ``name`` is not one of the five known configs.
FileNotFoundError: If the YAML file is absent.
pydantic.ValidationError: If the contents violate the schema.
"""
if name not in CONFIG_SCHEMAS:
raise KeyError(f"Unknown config '{name}'. Known configs: {sorted(CONFIG_SCHEMAS)}.")
path = (config_dir or CONFIG_DIR) / f"{name}.yaml"
if not path.exists():
raise FileNotFoundError(f"Config file not found: {path}")
with path.open("r", encoding="utf-8") as handle:
raw = yaml.safe_load(handle)
if raw is None:
raise ValueError(f"Config file is empty: {path}")
return CONFIG_SCHEMAS[name].model_validate(raw)
def load_configs(config_dir: Path | None = None) -> ConfigBundle:
"""Load and validate all five configs together, including cross-file checks."""
return ConfigBundle.model_validate({name: load_config(name, config_dir) for name in CONFIG_SCHEMAS})
def run_config_validation(config_dir: Path | None = None) -> ConfigBundle:
"""Validate every config and log a summary. Entry point for Phase 1 acceptance."""
from src.utils.logger import get_logger, log_section
logger = get_logger("config")
log_section(logger, "Configuration validation")
bundle = load_configs(config_dir)
logger.info("All %d configs loaded and validated against their schemas.", len(CONFIG_SCHEMAS))
logger.info(
"Grades: %s | warranty target: %d cycles",
", ".join(f"{g}>={bundle.grading.grades[g].min_cycles}" for g in bundle.grading.grade_order),
bundle.grading.warranty_target_cycles,
)
logger.info(
"Cost matrix: mean escape/overkill ratio = %.1f (sweep: %s)",
bundle.costs.escape_overkill_ratio(),
bundle.costs.sensitivity.escape_overkill_ratios,
)
logger.info("Budgets: %s (reference %d)", bundle.features.budgets, bundle.features.reference_budget)
logger.info("Enabled feature groups: %s", ", ".join(bundle.features.enabled_groups()))
logger.info("Enabled models: %s", ", ".join(bundle.models.enabled_models()))
logger.info("Exclusions declared: %d cells, each with a reason.", len(bundle.data.exclusions))
pending = bundle.unverified_assumptions()
logger.warning("Unverified assumptions outstanding for Phase 2: %d", len(pending))
for item in pending:
logger.warning(" [ ] %s", item)
return bundle
if __name__ == "__main__":
run_config_validation()