monsoon-rl / build_continuous_historical.py
DHDRL's picture
Update build_continuous_historical.py
aacbbba verified
Raw
History Blame Contribute Delete
18.1 kB
"""
build_continuous_historical.py
==============================
Causal historical cache for Indonesian rice zones.
Pairs with:
- era5_data_pipeline.fetch_episode_context (valid_time-anchored, backward
30-day obs windows, archive URL for dates older than today-5)
- zone_observation.DataSource.OPENMETEO_ARCHIVE
- real_episode_sampler.DEFAULT_HOLDOUT_RANGES
This builder does NOT issue NWP. forecast_backend="baseline" is persistence
of causal precip_30d. Do not report forecast skill against this file.
Holdout seasons are stored in the pickle so evaluate_checkpoint_real.py can
see them. RealEpisodeIndex excludes DEFAULT_HOLDOUT_RANGES at train time.
"""
from __future__ import annotations
import argparse
import json
import logging
import pickle
import time
from collections import Counter
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
import zone_observation as _zo
assert _zo.SCHEMA_VERSION == 3, (
f"build_continuous_historical: schema mismatch "
f"(expected 3, got {_zo.SCHEMA_VERSION})"
)
from zone_observation import CropStage, DataSource, ForecastConfig
from indonesia_zones import (
INDONESIA_ZONES,
crop_stage_for_date,
register_indonesia_zones,
)
from era5_data_pipeline import fetch_episode_context
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s | %(levelname)s | %(message)s",
)
logger = logging.getLogger("continuous_cache")
CACHE_VERSION = "indonesia_continuous_v3_causal_obs"
PRIORITY_ZONES = [
"karawang_rice",
"indramayu_rice",
"central_java_rice",
"east_java_rice",
"lampung_rice",
"south_sumatra_rice",
"banten_rice",
"south_sulawesi_rice",
]
# Seasons stored in the pickle. Train vs eval is decided by
# real_episode_sampler.DEFAULT_HOLDOUT_RANGES, not by omitting rows here.
#
# DEFAULT_HOLDOUT_RANGES currently:
# 2023-05-01 -> 2024-04-30 el_nino_2023_24_strong
# 2017-05-01 -> 2018-04-30 neutral_2017_18
#
# neutral_2016_17 is the Neutral year that remains eligible for training.
PARADIGMATIC_SEASONS: List[Tuple[str, str, str, str]] = [
("elnino_2015_16_vstrong", "2015-05-01", "2016-04-30", "el_nino_very_strong"),
("neutral_2016_17", "2016-05-01", "2017-04-30", "neutral"),
("neutral_2017_18", "2017-05-01", "2018-04-30", "neutral"),
("elnino_2018_19", "2018-06-01", "2019-05-31", "el_nino_moderate"),
("lanina_2020_21", "2020-09-01", "2021-05-31", "la_nina_moderate"),
("lanina_2021_22", "2021-09-01", "2022-05-31", "la_nina_moderate"),
("lanina_2022_23", "2022-09-01", "2023-04-30", "la_nina_weak_moderate"),
("elnino_2023_24_strong", "2023-05-01", "2024-04-30", "el_nino_strong"),
]
# Mirrors real_episode_sampler.DEFAULT_HOLDOUT_RANGES for the QA sidecar.
# If you change the sampler holdout, change this list too (or pass --holdout).
DEFAULT_HOLDOUT_RANGES: Tuple[Tuple[str, str], ...] = (
("2023-05-01", "2024-04-30"),
("2017-05-01", "2018-04-30"),
)
_ARCHIVE_LAG_DAYS = 5
def _parse(s: str) -> datetime:
return datetime.strptime(s, "%Y-%m-%d").replace(tzinfo=timezone.utc)
def _daterange(start: datetime, end: datetime, step_days: int = 3):
if step_days < 1:
raise ValueError(f"step_days must be >= 1, got {step_days}")
cur = start
while cur <= end:
yield cur
cur += timedelta(days=step_days)
def _in_holdout(day: datetime, holdout: Tuple[Tuple[str, str], ...]) -> bool:
d = day.date().isoformat()
for start_s, end_s in holdout:
if start_s <= d <= end_s:
return True
return False
def _enrich_with_crop_stage(
obs_dict: Dict[str, Any], zone_id: str, valid_time: datetime
) -> Dict[str, Any]:
try:
stage, days_to_harvest, season_name = crop_stage_for_date(zone_id, valid_time)
obs_dict["crop_stage"] = stage.value if isinstance(stage, CropStage) else str(stage)
obs_dict["days_to_harvest"] = days_to_harvest
if obs_dict.get("extras") is None:
obs_dict["extras"] = {}
if season_name:
obs_dict["extras"]["season_name"] = season_name
except Exception as e:
logger.debug("crop_stage_for_date failed for %s @ %s: %s", zone_id, valid_time.date(), e)
return obs_dict
def _safe_to_dict(obj) -> Optional[Dict[str, Any]]:
if obj is None:
return None
if hasattr(obj, "to_dict"):
return obj.to_dict()
try:
return dict(obj.__dict__)
except Exception:
return None
def _source_value(obs) -> str:
src = getattr(obs, "source", None)
if src is None:
return "unknown"
return src.value if hasattr(src, "value") else str(src)
def _assert_causal_obs(ctx: Any, day: datetime) -> None:
"""Fail the point rather than write a leaked or synthetic row."""
obs = ctx.obs
if obs is None:
raise RuntimeError("fetch_episode_context returned ctx.obs=None")
obs_day = obs.valid_time
if getattr(obs_day, "tzinfo", None) is None:
obs_day = obs_day.replace(tzinfo=timezone.utc)
if obs_day.date() != day.date():
raise RuntimeError(
f"obs.valid_time {obs_day.date()} != anchor {day.date()}"
)
src = getattr(obs, "source", None)
if src == DataSource.SYNTHETIC:
raise RuntimeError(
f"synthetic obs refused at {day.date()} zone={obs.zone_id}"
)
forecast = getattr(ctx, "forecast", None)
if forecast is not None:
ft = getattr(forecast, "forecast_time", None)
if ft is not None:
if getattr(ft, "tzinfo", None) is None:
ft = ft.replace(tzinfo=timezone.utc)
if ft.date() > day.date():
raise RuntimeError(
f"forecast_time {ft.date()} is after anchor {day.date()}"
)
today = datetime.now(timezone.utc).date()
if day.date() < today - timedelta(days=_ARCHIVE_LAG_DAYS):
archive = getattr(DataSource, "OPENMETEO_ARCHIVE", None)
if archive is not None and src not in (
archive,
DataSource.ERA5_REANALYSIS,
DataSource.SATELLITE_PRECIP,
DataSource.SATELLITE_SOIL,
):
raise RuntimeError(
f"historical day {day.date()} stamped {src} "
f"(expected OPENMETEO_ARCHIVE / ERA5 / satellite). "
f"zone_observation.DataSource.OPENMETEO_ARCHIVE is missing "
f"or the pipeline on this machine is stale."
)
def _qa_payload(
trajectories: List[Dict[str, Any]],
failures: int,
zones: List[str],
cfg: ForecastConfig,
step_days: int,
holdout: Tuple[Tuple[str, str], ...],
) -> Dict[str, Any]:
source_counts: Counter = Counter()
model_counts: Counter = Counter()
n_train = 0
n_holdout = 0
n_missing_basin = 0
n_synthetic = 0
n_points = 0
for traj in trajectories:
for pt in traj.get("trajectory", []):
n_points += 1
obs = pt.get("obs") or {}
src = str(obs.get("source") or pt.get("data_source") or "unknown")
source_counts[src] += 1
if src.endswith("synthetic") or src == "DataSource.SYNTHETIC":
n_synthetic += 1
fcast = pt.get("forecast") or {}
model_counts[str(fcast.get("model_id") or "missing")] += 1
if pt.get("basin_context") is None:
n_missing_basin += 1
vt = pt.get("valid_time") or ""
day = _parse(vt[:10]) if vt else None
if day is None:
continue
if _in_holdout(day, holdout):
n_holdout += 1
else:
n_train += 1
return {
"version": CACHE_VERSION,
"n_trajectories": len(trajectories),
"total_points": n_points,
"failures_dropped": failures,
"zones": zones,
"seasons": [s[0] for s in PARADIGMATIC_SEASONS],
"step_days": step_days,
"holdout_ranges": [list(h) for h in holdout],
"n_train_eligible": n_train,
"n_holdout_stored": n_holdout,
"n_missing_basin": n_missing_basin,
"n_synthetic_obs": n_synthetic,
"source_histogram": dict(source_counts),
"forecast_model_histogram": dict(model_counts),
"config_snapshot": {
"forecast_backend": cfg.forecast_backend,
"use_climatology_anomalies": cfg.use_climatology_anomalies,
"include_basin_context": cfg.include_basin_context,
"require_real_basin_context": cfg.require_real_basin_context,
"force_data_source": (
cfg.force_data_source.value if cfg.force_data_source else None
),
"climatology_years": cfg.climatology_years,
},
"notes": [
"forecast_backend=baseline is persistence of causal precip_30d, not NWP",
"holdout rows are stored for eval; RealEpisodeIndex excludes them at train",
"n_synthetic_obs must be 0 before this pickle is used for training",
],
}
def build_continuous_cache(
output_path: str = "historical_continuous_indonesia_v3_causal_obs.pkl",
step_days: int = 3,
sleep_s: float = 0.7,
max_days_per_zone_season: int = 400,
resume: bool = True,
require_real_basin_context: bool = False,
holdout_ranges: Tuple[Tuple[str, str], ...] = DEFAULT_HOLDOUT_RANGES,
) -> None:
register_indonesia_zones()
available = {z.zone_id for z in INDONESIA_ZONES}
zones = [z for z in PRIORITY_ZONES if z in available]
missing = [z for z in PRIORITY_ZONES if z not in available]
if missing:
logger.warning("Priority zones not in registry (skipped): %s", missing)
if not zones:
raise RuntimeError("No priority zones registered — aborting")
logger.info("Priority zones (%d): %s", len(zones), zones)
if not hasattr(DataSource, "OPENMETEO_ARCHIVE"):
logger.warning(
"DataSource.OPENMETEO_ARCHIVE is missing. Historical points "
"will not pass the archive-stamp assertion. Upload the updated "
"zone_observation.py before rebuilding."
)
cfg = ForecastConfig(
forecast_backend="baseline",
use_climatology_anomalies=True,
include_basin_context=True,
require_real_basin_context=require_real_basin_context,
force_data_source=DataSource.OPENMETEO_LIVE,
real_data_ratio=1.0,
climatology_years=10,
)
trajectories: List[Dict[str, Any]] = []
failures = 0
t0 = time.time()
out = Path(output_path)
if resume and out.exists():
try:
with open(out, "rb") as f:
existing = pickle.load(f)
if existing.get("version") != CACHE_VERSION:
logger.warning(
"Resume file version=%s != %s — starting fresh so a "
"leaky v1/v2 cache cannot be extended.",
existing.get("version"), CACHE_VERSION,
)
else:
trajectories = existing.get("trajectories", [])
logger.info("Resuming from %d existing trajectories", len(trajectories))
except Exception as e:
logger.warning("Resume failed (%s) — starting fresh", e)
trajectories = []
already_done = {
(t["meta"]["label"], t["meta"]["zone_id"]) for t in trajectories
}
for label, start_s, end_s, regime in PARADIGMATIC_SEASONS:
start = _parse(start_s)
end = _parse(end_s)
logger.info("=== %s (%s → %s) [%s] ===", label, start_s, end_s, regime)
for zone_id in zones:
key = (label, zone_id)
if key in already_done:
logger.info(" %s already present — skipping", zone_id)
continue
traj_points: List[Dict[str, Any]] = []
days_fetched = 0
for day in _daterange(start, end, step_days=step_days):
if days_fetched >= max_days_per_zone_season:
logger.warning(
" %s hit max_days=%d before %s — raise --max-days",
zone_id, max_days_per_zone_season, end_s,
)
break
window_end = day + timedelta(days=30)
try:
ctx = fetch_episode_context(zone_id, (day, window_end), cfg)
_assert_causal_obs(ctx, day)
obs_dict = _safe_to_dict(ctx.obs) or {}
obs_dict = _enrich_with_crop_stage(obs_dict, zone_id, day)
point = {
"valid_time": day.isoformat(),
"zone_id": zone_id,
"obs": obs_dict,
"forecast": _safe_to_dict(ctx.forecast),
"basin_context": _safe_to_dict(
getattr(ctx, "basin_context", None)
),
"data_source": _source_value(ctx.obs),
"split": (
"holdout" if _in_holdout(day, holdout_ranges) else "train"
),
}
traj_points.append(point)
days_fetched += 1
time.sleep(sleep_s)
except Exception as e:
msg = str(e)
logger.warning(" Drop %s @ %s: %s", zone_id, day.date(), msg[:200])
failures += 1
time.sleep(sleep_s * 1.3)
continue
if traj_points:
trajectories.append({
"meta": {
"label": label,
"regime": regime,
"zone_id": zone_id,
"start": start_s,
"end": end_s,
"n_points": len(traj_points),
"step_days": step_days,
"n_holdout": sum(
1 for p in traj_points if p.get("split") == "holdout"
),
},
"trajectory": traj_points,
})
logger.info(" %s: %d causal points saved", zone_id, len(traj_points))
if len(trajectories) % 3 == 0:
_save(
trajectories, out, failures, zones, cfg,
step_days, holdout_ranges,
)
_save(trajectories, out, failures, zones, cfg, step_days, holdout_ranges)
elapsed = (time.time() - t0) / 60
qa = _qa_payload(trajectories, failures, zones, cfg, step_days, holdout_ranges)
logger.info("=" * 70)
logger.info("CAUSAL HISTORICAL CACHE COMPLETE")
logger.info(" Trajectories : %d", qa["n_trajectories"])
logger.info(" Total points : %d", qa["total_points"])
logger.info(" Train-eligible : %d", qa["n_train_eligible"])
logger.info(" Holdout stored : %d", qa["n_holdout_stored"])
logger.info(" Synthetic obs : %d", qa["n_synthetic_obs"])
logger.info(" Missing basin : %d", qa["n_missing_basin"])
logger.info(" Dropped : %d", failures)
logger.info(" Sources : %s", qa["source_histogram"])
logger.info(" Forecast models : %s", qa["forecast_model_histogram"])
logger.info(" Elapsed : %.1f min", elapsed)
logger.info(" Output : %s", out)
logger.info("=" * 70)
if qa["n_synthetic_obs"] != 0:
raise RuntimeError(
f"cache contains {qa['n_synthetic_obs']} synthetic obs rows — "
"do not train on this file"
)
if qa["n_train_eligible"] == 0:
raise RuntimeError("cache has 0 train-eligible points after holdout")
def _save(
trajectories,
out: Path,
failures: int,
zones,
cfg: ForecastConfig,
step_days: int,
holdout: Tuple[Tuple[str, str], ...],
) -> None:
qa = _qa_payload(trajectories, failures, zones, cfg, step_days, holdout)
payload = {
"version": CACHE_VERSION,
"created_utc": datetime.now(timezone.utc).isoformat(),
"design": "continuous_paradigmatic_seasons_causal_obs",
"n_trajectories": qa["n_trajectories"],
"total_points": qa["total_points"],
"priority_zones": zones,
"holdout_ranges": [list(h) for h in holdout],
"config_snapshot": qa["config_snapshot"],
"qa": qa,
"trajectories": trajectories,
}
tmp = out.with_suffix(out.suffix + ".tmp")
with open(tmp, "wb") as f:
pickle.dump(payload, f, protocol=pickle.HIGHEST_PROTOCOL)
tmp.replace(out)
with open(out.with_suffix(".summary.json"), "w") as f:
json.dump(qa, f, indent=2)
def main() -> None:
p = argparse.ArgumentParser(
description="Build a causal-obs Indonesia historical cache (v3)."
)
p.add_argument(
"--output",
default="historical_continuous_indonesia_v3_causal_obs.pkl",
)
p.add_argument("--step-days", type=int, default=3)
p.add_argument("--sleep", type=float, default=0.7)
p.add_argument("--max-days", type=int, default=400)
p.add_argument("--no-resume", action="store_true")
p.add_argument(
"--require-real-basin",
action="store_true",
help="Drop the point if RONI/DMI/SWPC cannot be fetched. "
"Off by default so a DMI 404 does not empty the cache; "
"ENSO/obs still have to be real.",
)
args = p.parse_args()
build_continuous_cache(
output_path=args.output,
step_days=args.step_days,
sleep_s=args.sleep,
max_days_per_zone_season=args.max_days,
resume=not args.no_resume,
require_real_basin_context=args.require_real_basin,
)
if __name__ == "__main__":
main()