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()