File size: 4,015 Bytes
0febf0f 5d26997 0febf0f 5d26997 0febf0f dfa8acb 0febf0f 2bb61d5 0febf0f 2bb61d5 ba502b0 2bb61d5 0febf0f 2bb61d5 0febf0f 5d26997 2bb61d5 0febf0f 2bb61d5 5d26997 dfa8acb 5d26997 2bb61d5 5d26997 0febf0f dfa8acb 5d26997 2bb61d5 5d26997 2bb61d5 0febf0f 2bb61d5 0febf0f 5d26997 0febf0f dfa8acb | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 | from __future__ import annotations
import os
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent / "src"))
import joblib
import pandas as pd
from batteryswap_public.utils import iterate_scenarios, load_dataset
from batteryswapai.competition_features import scenario_history_snapshot
def _assert_complete_plan(
plan: pd.DataFrame, locations: pd.DataFrame, scenario_name: str
) -> None:
required = {"day", "battery"}
missing = required - set(plan.columns)
if missing:
raise AssertionError(
f"plan {scenario_name} lacks columns: {sorted(missing)}"
)
expected = locations["battery"].astype(str)
actual = plan["battery"].astype(str)
if len(plan) != len(locations) or actual.duplicated().any():
raise AssertionError(
f"plan {scenario_name} must contain each live battery exactly once"
)
if set(actual) != set(expected):
missing_batteries = sorted(set(expected) - set(actual))
extra_batteries = sorted(set(actual) - set(expected))
raise AssertionError(
f"plan {scenario_name} battery mismatch: "
f"missing={missing_batteries[:3]}, extra={extra_batteries[:3]}"
)
parsed_days = pd.to_datetime(plan["day"], errors="coerce")
if parsed_days.isna().any():
raise AssertionError(f"plan {scenario_name} contains an invalid day")
def main() -> None:
dataset_path = Path(os.environ.get("BATTERYSWAP_DATASET_PATH", "/tmp/data"))
artifact_path = Path(
os.environ.get(
"BATTERYSWAP_PLANNER_PATH",
"submission_artifacts/weekly99_planner.joblib",
)
)
splits = [
split.strip()
for split in os.environ.get("BATTERYSWAP_SPLITS", "public,private").split(",")
if split.strip()
]
output_path = Path(os.environ.get("BATTERYSWAP_SUBMISSION_PATH", "submission.csv"))
if not artifact_path.is_file():
raise FileNotFoundError(f"Planner artifact does not exist: {artifact_path}")
planner = joblib.load(artifact_path)
plans = []
for split in splits:
locations, timeseries, eol_times_for_iterator, scenarios = load_dataset(
dataset_path / split
)
reset_split = getattr(planner, "reset_split", None)
if reset_split is not None:
reset_split(split)
for scenario, locs, visible_history, _ in iterate_scenarios(
locations,
timeseries,
eol_times_for_iterator,
scenarios,
):
snapshot = scenario_history_snapshot(
visible_history,
locs,
scenario["name"],
scenario["start_time"],
)
plan_scenario = getattr(planner, "plan_scenario", None)
if plan_scenario is None:
plan = planner.plan_snapshot(
snapshot,
locs,
scenario["travel_costs"],
scenario["settings"],
scenario["start_time"],
)
else:
plan = plan_scenario(
visible_history,
snapshot,
locs,
scenario["travel_costs"],
scenario["settings"],
scenario["start_time"],
)
_assert_complete_plan(plan, locs, str(scenario["name"]))
plan["split"] = split
plan["scenario"] = scenario["name"]
plans.append(plan)
submission = pd.concat(plans, ignore_index=True)
if submission.duplicated(["split", "scenario", "battery"]).any():
raise AssertionError("submission contains duplicate split/scenario/battery rows")
submission.to_csv(output_path, index=False)
if not output_path.exists():
raise RuntimeError(f"Submission was not created: {output_path}")
if __name__ == "__main__":
main()
|