BatterySwapAI2026-MnesisLab / scripts /train_submission.py
CarlAlbertCode's picture
Submission 007 - temperature-gated causal identity ensemble
2bb61d5 verified
Raw
History Blame Contribute Delete
7.28 kB
from __future__ import annotations
import argparse
import dataclasses
import hashlib
import importlib.metadata
import json
from pathlib import Path
import joblib
import pandas as pd
from batteryswap_public.evaluate import evaluate_plan
from batteryswap_public.utils import iterate_scenarios, load_dataset
from batteryswapai.competition_features import (
build_trajectory_matrix,
scenario_history_snapshot,
attach_training_targets,
build_daily_features,
scenario_snapshot,
)
from batteryswapai.competition_model import fit_event_time_model
from batteryswapai.competition_planner import CompetitionPlanner, PlannerPolicy
DATASET_REVISION = "7f423ac4cb6ab146f7ea7a37872eb4dfc3c9705c"
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--dataset-path", type=Path, default=Path("data/raw/train"))
parser.add_argument(
"--artifact", type=Path, default=Path("submission_artifacts/planner.joblib")
)
parser.add_argument("--quantile", type=float, default=0.05)
parser.add_argument("--event-risk-threshold", type=float, default=0.50)
parser.add_argument("--prediction-offset-days", type=float, default=-5.0)
parser.add_argument("--risk-calibration-scale", type=float, default=1.5)
parser.add_argument("--expected-gain-margin", type=float, default=10.0)
parser.add_argument("--expected-service-cost-hours", type=float, default=2.0)
parser.add_argument("--emergency-operational-scale", type=float, default=0.0)
parser.add_argument("--scheduled-fraction", type=float, default=0.038)
parser.add_argument("--capacity-lookback-days", type=int, default=42)
parser.add_argument("--weekly-guard-fraction", type=float, default=0.95)
parser.add_argument("--hard-limit-penalty-multiplier", type=float, default=1.5)
parser.add_argument("--minimum-scheduled-batteries", type=int, default=8)
parser.add_argument("--maximum-scheduled-batteries", type=int, default=24)
parser.add_argument("--calibration-folds", type=int, default=5)
parser.add_argument("--skip-evaluation", action="store_true")
return parser.parse_args()
def main() -> None:
args = parse_args()
locations, timeseries, eol_times, scenarios = load_dataset(args.dataset_path)
daily = build_daily_features(timeseries)
snapshots = []
for scenario, locs, visible, _ in iterate_scenarios(locations, timeseries, eol_times, scenarios):
snapshot = scenario_history_snapshot(
visible, locs, scenario["name"], scenario["start_time"]
)
snapshot = attach_training_targets(
snapshot,
eol_times,
unobserved_eol_days=float(scenario["settings"].unobserved_eol_days),
)
snapshots.append(snapshot)
training = pd.concat(snapshots, ignore_index=True)
model = fit_event_time_model(
training,
quantile=args.quantile,
dataset_revision=DATASET_REVISION,
calibration_folds=args.calibration_folds,
trajectory=build_trajectory_matrix(daily),
eol_times=eol_times,
)
planner = CompetitionPlanner(
model,
PlannerPolicy(
event_risk_threshold=args.event_risk_threshold,
prediction_offset_days=args.prediction_offset_days,
use_expected_cost=True,
risk_calibration_scale=args.risk_calibration_scale,
expected_gain_margin=args.expected_gain_margin,
expected_service_cost_hours=args.expected_service_cost_hours,
emergency_operational_scale=args.emergency_operational_scale,
scheduled_fraction=args.scheduled_fraction,
capacity_weekly_limit_fraction=args.weekly_guard_fraction,
capacity_limit_penalty_multiplier=args.hard_limit_penalty_multiplier,
minimum_scheduled_batteries=args.minimum_scheduled_batteries,
maximum_scheduled_batteries=args.maximum_scheduled_batteries,
capacity_lookahead_days=21,
capacity_lookback_days=args.capacity_lookback_days,
),
)
args.artifact.parent.mkdir(parents=True, exist_ok=True)
joblib.dump(planner, args.artifact, compress=3)
report = {
"dataset_revision": DATASET_REVISION,
"artifact": args.artifact.as_posix(),
"quantile": args.quantile,
"event_risk_threshold": args.event_risk_threshold,
"prediction_offset_days": args.prediction_offset_days,
"risk_calibration_scale": args.risk_calibration_scale,
"expected_gain_margin": args.expected_gain_margin,
"expected_service_cost_hours": args.expected_service_cost_hours,
"emergency_operational_scale": args.emergency_operational_scale,
"scheduled_fraction": args.scheduled_fraction,
"capacity_lookback_days": args.capacity_lookback_days,
"weekly_guard_fraction": args.weekly_guard_fraction,
"hard_limit_penalty_multiplier": args.hard_limit_penalty_multiplier,
"minimum_scheduled_batteries": args.minimum_scheduled_batteries,
"maximum_scheduled_batteries": args.maximum_scheduled_batteries,
"calibration_folds": args.calibration_folds,
"planner_policy": dataclasses.asdict(planner.policy),
"survivor_rul_head": model.survivor_rul_estimator is not None,
"trajectory_weight": model.trajectory_weight,
"forecast_heads": sorted(model.forecasters or {}),
"training_rows": len(training),
"training_devices": int(training["battery"].nunique()),
"features": model.feature_columns,
"artifact_sha256": hashlib.sha256(args.artifact.read_bytes()).hexdigest(),
"runtime_versions": {
package: importlib.metadata.version(package)
for package in (
"batteryswap_public",
"fastparquet",
"joblib",
"numpy",
"pandas",
"scikit-learn",
)
},
"validation_report": "artifacts/submission_cv_final.json",
}
if not args.skip_evaluation:
scores = []
for scenario, locs, visible, not_dead in iterate_scenarios(
locations, timeseries, eol_times, scenarios
):
snapshot = scenario_history_snapshot(
visible, locs, scenario["name"], scenario["start_time"]
)
plan = planner.plan_snapshot(
snapshot,
locs,
scenario["travel_costs"],
scenario["settings"],
scenario["start_time"],
)
_, _, score = evaluate_plan(
plan,
locs,
scenario["travel_costs"],
scenario["settings"],
eol_times=not_dead,
start_time=pd.Timestamp(scenario["start_time"]),
verbose=0,
)
scores.append(score)
report["optimistic_train_score"] = pd.concat(scores, axis=1).mean(axis=1).to_dict()
report_path = args.artifact.with_suffix(".json")
report_path.write_text(json.dumps(report, indent=2, default=float), encoding="utf-8")
print(json.dumps(report, indent=2, default=float))
if __name__ == "__main__":
main()