CarlAlbertCode's picture
Submission 016: weekly fraction 0.99
ba502b0
Raw
History Blame Contribute Delete
4.02 kB
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()