262 lines
9.0 KiB
Python
262 lines
9.0 KiB
Python
"""Tests for the leakage-safe nonlinear candidate."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import importlib.util
|
|
import math
|
|
import sys
|
|
import unittest
|
|
from datetime import datetime
|
|
from pathlib import Path
|
|
from types import SimpleNamespace
|
|
from typing import Dict, List
|
|
|
|
|
|
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
|
MODULE_PATH = PROJECT_ROOT / "scripts/train_tree_model.py"
|
|
SPEC = importlib.util.spec_from_file_location("train_tree_model", MODULE_PATH)
|
|
if SPEC is None or SPEC.loader is None: # pragma: no cover
|
|
raise RuntimeError("Could not load scripts/train_tree_model.py")
|
|
tree = importlib.util.module_from_spec(SPEC)
|
|
sys.modules[SPEC.name] = tree
|
|
SPEC.loader.exec_module(tree)
|
|
baselines = tree.baselines
|
|
|
|
SKLEARN_AVAILABLE = (
|
|
importlib.util.find_spec("numpy") is not None
|
|
and importlib.util.find_spec("sklearn") is not None
|
|
)
|
|
|
|
|
|
def episode(
|
|
number: int,
|
|
partition: str,
|
|
year: int,
|
|
target: int,
|
|
age: float,
|
|
make: str,
|
|
county: str = "salt_lake",
|
|
audit: bool = False,
|
|
) -> baselines.EpisodeRow:
|
|
bucket = 5 if audit else 25 + (number % 70)
|
|
numeric: Dict[str, float] = {
|
|
"vehicle_age": age,
|
|
"prior_episode_count": 1.0 + number % 8,
|
|
"prior_total_attempt_count": 1.0 + number % 12,
|
|
"prior_attempt_count": 1.0 + number % 3,
|
|
"days_since_prior_episode": 300.0 + number % 150,
|
|
"days_since_prior_adverse": 200.0 + number % 500,
|
|
"prior_nonpass_rate": (number % 10) / 10.0,
|
|
}
|
|
categorical = {
|
|
"public_county": county,
|
|
"target_season": ("summer" if number % 2 else "winter"),
|
|
"prior_first_outcome": ("fail" if number % 4 == 0 else "pass"),
|
|
"prior_final_outcome": ("fail" if number % 5 == 0 else "pass"),
|
|
"last_observed_make": make,
|
|
"last_observed_model": ("model_a" if number % 3 else "model_b"),
|
|
}
|
|
return baselines.EpisodeRow(
|
|
vehicle_token="{:064x}".format(number + year * 100_000),
|
|
vehicle_bucket=bucket,
|
|
is_vin_audit=audit,
|
|
episode_number=2 + number % 10,
|
|
episode_start=datetime(year, 6, 1),
|
|
partition=partition,
|
|
target=target,
|
|
source_era="slc",
|
|
target_outcome_label_source="overall_result",
|
|
prior_first_outcome=categorical["prior_first_outcome"],
|
|
numeric=numeric,
|
|
categorical=categorical,
|
|
)
|
|
|
|
|
|
def partition_rows(partition: str, year: int, count: int, offset: int) -> List[object]:
|
|
rows: List[object] = []
|
|
for index in range(count):
|
|
number = offset + index
|
|
audit = index % 10 == 0
|
|
age = float((number * 7) % 30)
|
|
county = "utah" if number % 2 else "salt_lake"
|
|
signal = (age > 14.0) != (county == "utah")
|
|
target = int(signal)
|
|
if number % 17 == 0:
|
|
target = 1 - target
|
|
if audit:
|
|
make = "audit_only_make"
|
|
elif partition == "train":
|
|
make = ("toyota" if number % 3 else "ford")
|
|
else:
|
|
make = ("future_make" if number % 7 == 0 else "toyota")
|
|
rows.append(
|
|
episode(
|
|
number=number,
|
|
partition=partition,
|
|
year=year,
|
|
target=target,
|
|
age=age,
|
|
make=make,
|
|
county=county,
|
|
audit=audit,
|
|
)
|
|
)
|
|
return rows
|
|
|
|
|
|
def synthetic_mart() -> baselines.MartData:
|
|
rows_by_partition = {
|
|
"train": partition_rows("train", 2022, 360, 0),
|
|
"tune": partition_rows("tune", 2023, 140, 10_000),
|
|
"calibrate": partition_rows("calibrate", 2024, 140, 20_000),
|
|
"locked_test": partition_rows("locked_test", 2025, 140, 30_000),
|
|
}
|
|
total = sum(len(rows) for rows in rows_by_partition.values())
|
|
return baselines.MartData(
|
|
rows_by_partition=rows_by_partition,
|
|
input_rows=total,
|
|
ineligible_rows=0,
|
|
eligible_rows=total,
|
|
)
|
|
|
|
|
|
class TreeEncoderTests(unittest.TestCase):
|
|
def test_future_categories_and_values_do_not_change_training_preprocessing(self) -> None:
|
|
train_rows = [
|
|
episode(1, "train", 2022, 0, 10.0, "toyota"),
|
|
episode(2, "train", 2022, 1, 20.0, "toyota"),
|
|
episode(3, "train", 2022, 0, 30.0, "ford"),
|
|
]
|
|
future = episode(4, "tune", 2023, 1, 9_999.0, "future_make")
|
|
encoder = tree.TreeFeatureEncoder.fit(
|
|
train_rows, min_category_count=2, max_categories=16
|
|
)
|
|
|
|
self.assertEqual(encoder.numeric_medians["vehicle_age"], 20.0)
|
|
self.assertEqual(
|
|
encoder.category_code("last_observed_make", "toyota"),
|
|
tree.FIRST_KNOWN_CATEGORY_CODE,
|
|
)
|
|
self.assertEqual(
|
|
encoder.category_code("last_observed_make", "ford"),
|
|
tree.RARE_CATEGORY_CODE,
|
|
)
|
|
self.assertEqual(
|
|
encoder.category_code("last_observed_make", "future_make"),
|
|
tree.UNKNOWN_CATEGORY_CODE,
|
|
)
|
|
self.assertNotIn(
|
|
"future_make", encoder.seen_categories["last_observed_make"]
|
|
)
|
|
self.assertEqual(encoder.numeric_medians["vehicle_age"], 20.0)
|
|
self.assertTrue(set(tree.EXCLUDED_FROM_PREDICTORS).isdisjoint(
|
|
encoder.feature_names
|
|
))
|
|
self.assertNotIn("source_era", encoder.feature_names)
|
|
self.assertNotIn("target_outcome_label_source", encoder.feature_names)
|
|
|
|
if SKLEARN_AVAILABLE:
|
|
import numpy as np
|
|
|
|
matrix = encoder.transform([future], np)
|
|
category_index = 2 * len(tree.NUMERIC_FEATURES) + list(
|
|
tree.CATEGORICAL_FEATURES
|
|
).index("last_observed_make")
|
|
self.assertEqual(
|
|
matrix[0, category_index], float(tree.UNKNOWN_CATEGORY_CODE)
|
|
)
|
|
self.assertTrue(np.isfinite(matrix).all())
|
|
|
|
@unittest.skipUnless(SKLEARN_AVAILABLE, "scikit-learn is not installed")
|
|
def test_convergence_gate_rejects_iteration_limit_and_nonfinite_scores(self) -> None:
|
|
import numpy as np
|
|
|
|
at_limit = SimpleNamespace(
|
|
n_iter_=100,
|
|
train_score_=np.asarray([-1.0]),
|
|
validation_score_=np.asarray([-1.0]),
|
|
)
|
|
converged, n_iter = tree._hist_converged(at_limit, [], 100, np)
|
|
self.assertFalse(converged)
|
|
self.assertEqual(n_iter, 100)
|
|
|
|
nonfinite = SimpleNamespace(
|
|
n_iter_=20,
|
|
train_score_=np.asarray([-1.0, math.nan]),
|
|
validation_score_=np.asarray([-1.0]),
|
|
)
|
|
converged, _ = tree._hist_converged(nonfinite, [], 100, np)
|
|
self.assertFalse(converged)
|
|
|
|
|
|
@unittest.skipUnless(SKLEARN_AVAILABLE, "scikit-learn is not installed")
|
|
class TreeTrainingTests(unittest.TestCase):
|
|
@classmethod
|
|
def setUpClass(cls) -> None:
|
|
cls.mart = synthetic_mart()
|
|
cls.trained = tree.train_tree_model(
|
|
mart=cls.mart,
|
|
candidates=[tree.TreeCandidate(0.1, 15, 1.0)],
|
|
min_category_count=3,
|
|
max_categories=32,
|
|
min_samples_leaf=10,
|
|
max_iter=150,
|
|
n_iter_no_change=8,
|
|
calibration_max_iter=1000,
|
|
seed=20260715,
|
|
)
|
|
|
|
def test_fit_excludes_audit_and_future_categories(self) -> None:
|
|
seen = self.trained.encoder.seen_categories["last_observed_make"]
|
|
self.assertNotIn("audit_only_make", seen)
|
|
self.assertNotIn("future_make", seen)
|
|
self.assertLess(self.trained.model_n_iter, self.trained.max_iter)
|
|
self.assertTrue(
|
|
all(item["eligible_for_selection"] for item in self.trained.tuning_results)
|
|
)
|
|
|
|
def test_locked_metrics_require_explicit_gate(self) -> None:
|
|
closed_metrics, _closed_bins = tree.evaluate_tree_model(
|
|
self.mart,
|
|
self.trained,
|
|
evaluate_locked=False,
|
|
bin_count=5,
|
|
)
|
|
self.assertNotIn(
|
|
"locked_test", {row["partition"] for row in closed_metrics}
|
|
)
|
|
|
|
open_metrics, _open_bins = tree.evaluate_tree_model(
|
|
self.mart,
|
|
self.trained,
|
|
evaluate_locked=True,
|
|
bin_count=5,
|
|
)
|
|
locked = [
|
|
row for row in open_metrics if row["partition"] == "locked_test"
|
|
]
|
|
self.assertTrue(locked)
|
|
self.assertIn("never_fit_audit", {row["cohort"] for row in locked})
|
|
|
|
def test_probabilities_are_finite(self) -> None:
|
|
probabilities = tree.tree_probabilities(
|
|
self.trained,
|
|
self.mart.rows_by_partition["calibrate"],
|
|
calibrated=True,
|
|
)
|
|
self.assertTrue(probabilities)
|
|
self.assertTrue(all(math.isfinite(value) for value in probabilities))
|
|
self.assertTrue(all(0.0 <= value <= 1.0 for value in probabilities))
|
|
|
|
def test_cli_locked_flag_defaults_closed(self) -> None:
|
|
closed = tree.parse_args(["--mart", "data/private/mart.parquet"])
|
|
opened = tree.parse_args(
|
|
["--mart", "data/private/mart.parquet", "--evaluate-locked"]
|
|
)
|
|
self.assertFalse(closed.evaluate_locked)
|
|
self.assertTrue(opened.evaluate_locked)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|