monsoon-rl / build_dynamics_pairs.py
DHDRL's picture
Upload 27 files
976eb45 verified
Raw
History Blame Contribute Delete
9.98 kB
#!/usr/bin/env python3
"""
build_dynamics_pairs.py
=======================
Last-mile path: historical cache → (current, next) ZoneStateTensor pairs
for physics_dynamics.DynamicsTrainer.
Previously the Dyna module was fully unit-tested but had no script that
turned real trajectories into training pairs. This closes that gap.
Contract
--------
Purpose: Emit list of (ZoneStateTensor, ZoneStateTensor) consecutive pairs
from historical_continuous_indonesia_v1.pkl (or compatible).
Allowed caller: offline data jobs, train_curriculum prep.
Forbidden: inventing precip; using look-ahead climatology labels as targets.
Writes: optional .pt file of pair list; JSON manifest.
Side effects: none on live systems.
Response: exit 0 + counts; structured manifest.
Notes on physics framing
------------------------
Pairs are consecutive valid_times along a trajectory (step_days from meta).
Precip channel = forecast precip_mm horizon (padded/truncated to horizon_days).
Belief / uncertainty are derived proxies from obs anomalies (not agent
beliefs) so DynamicsTrainer has a real-data fuel path without requiring a
trained policy. This is intentional for the first real-data dynamics fit.
Usage
-----
python build_dynamics_pairs.py \\
--pkl historical_continuous_indonesia_v1.pkl \\
--zones karawang_rice,indramayu_rice \\
--horizon-days 14 \\
--out dynamics_pairs.pt \\
--manifest dynamics_pairs_manifest.json
"""
from __future__ import annotations
import argparse
import json
import logging
import pickle
import sys
from datetime import date, datetime
from pathlib import Path
from typing import Any, Dict, List, Optional, Sequence, Tuple
import numpy as np
logger = logging.getLogger(__name__)
def _parse_day(s: str) -> date:
return date.fromisoformat(str(s)[:10])
def _forecast_precip_array(fc: Dict[str, Any], horizon_days: int) -> np.ndarray:
precip = fc.get("precip_mm") or ()
arr = np.array(list(precip)[:horizon_days], dtype=np.float32)
if arr.shape[0] < horizon_days:
pad = np.zeros(horizon_days - arr.shape[0], dtype=np.float32)
arr = np.concatenate([arr, pad])
return np.clip(arr, 0.0, 500.0)
def _belief_uncertainty_from_obs(obs: Dict[str, Any]) -> Tuple[float, float]:
"""
Proxy scalars so pairs need no trained agent.
belief ~ max of clipped anomaly magnitudes mapped to [0,1]
uncertainty ~ higher when quality_flag low / cloud high
"""
anom = abs(float(obs.get("precip_anomaly_idx") or 0.0))
soil = abs(float(obs.get("soil_moisture_anom") or 0.0))
belief = float(np.clip(max(anom, soil) / 3.0, 0.0, 1.0))
q = int(obs.get("quality_flag") or 3)
cloud = float(obs.get("cloud_cover_pct") or 0.0)
uncertainty = float(np.clip(0.2 + 0.1 * max(0, 3 - q) + cloud / 200.0, 0.05, 0.95))
return belief, uncertainty
def extract_pairs_from_pkl(
pkl_path: Path,
zone_ids: Sequence[str],
horizon_days: int = 14,
start: Optional[date] = None,
end: Optional[date] = None,
) -> Tuple[List[Tuple[Any, Any]], List[float], Dict[str, Any]]:
"""
Build consecutive pairs within each trajectory (same zone, ordered time).
Returns (pairs, dts, manifest_stats). dts[i] is the real-time gap in days
for pairs[i] (exact step_days only). Requires torch only when materializing
ZoneStateTensor — import deferred so --help works offline.
"""
import torch
from physics_dynamics import ZoneStateTensor
with open(pkl_path, "rb") as f:
cache = pickle.load(f)
zone_set = set(zone_ids)
pairs: List[Tuple[ZoneStateTensor, ZoneStateTensor]] = []
dts: List[float] = []
stats = {
"n_trajectories_seen": 0,
"n_trajectories_used": 0,
"n_points": 0,
"n_pairs": 0,
"zones": list(zone_ids),
"horizon_days": horizon_days,
"skipped_non_consecutive": 0,
"skipped_non_exact_step": 0,
"dt_values": {},
"default_dt_hint": None,
}
for traj in cache.get("trajectories") or []:
stats["n_trajectories_seen"] += 1
meta = traj.get("meta") or {}
zid = meta.get("zone_id")
if zid not in zone_set:
continue
points = list(traj.get("trajectory") or [])
if len(points) < 2:
continue
# filter by date if requested
filtered = []
for pt in points:
d = _parse_day(pt.get("valid_time", "1970-01-01"))
if start and d < start:
continue
if end and d > end:
continue
filtered.append(pt)
if len(filtered) < 2:
continue
stats["n_trajectories_used"] += 1
stats["n_points"] += len(filtered)
step_days = int(meta.get("step_days") or 5)
if stats["default_dt_hint"] is None:
stats["default_dt_hint"] = float(step_days)
for i in range(len(filtered) - 1):
a, b = filtered[i], filtered[i + 1]
da = _parse_day(a["valid_time"])
db = _parse_day(b["valid_time"])
gap = (db - da).days
if gap <= 0:
stats["skipped_non_consecutive"] += 1
continue
# Exact step only — avoids dt mismatch with PDE residual.
# (Previously allowed 2x step while physics_loss used dt=1.0.)
if gap != step_days:
stats["skipped_non_exact_step"] += 1
continue
pa = _forecast_precip_array(a["forecast"], horizon_days)
pb = _forecast_precip_array(b["forecast"], horizon_days)
ba, ua = _belief_uncertainty_from_obs(a["obs"])
bb, ub = _belief_uncertainty_from_obs(b["obs"])
# shapes: precip [1, 1, H], uncertainty [1, 1], belief [1, 1]
curr = ZoneStateTensor.from_numpy(
precip=pa.reshape(1, 1, -1),
uncertainty=np.array([ua], dtype=np.float32),
belief=np.array([ba], dtype=np.float32),
)
nxt = ZoneStateTensor.from_numpy(
precip=pb.reshape(1, 1, -1),
uncertainty=np.array([ub], dtype=np.float32),
belief=np.array([bb], dtype=np.float32),
)
pairs.append((curr, nxt))
dts.append(float(gap))
key = str(int(gap))
stats["dt_values"][key] = int(stats["dt_values"].get(key, 0)) + 1
stats["n_pairs"] += 1
return pairs, dts, stats
def main(argv: Optional[Sequence[str]] = None) -> int:
logging.basicConfig(level=logging.INFO, format="%(levelname)s %(name)s: %(message)s")
p = argparse.ArgumentParser(description="Build DynamicsTrainer pairs from historical pkl")
p.add_argument("--pkl", required=True)
p.add_argument("--zones", default="karawang_rice,indramayu_rice")
p.add_argument("--horizon-days", type=int, default=14)
p.add_argument("--start", default=None)
p.add_argument("--end", default=None)
p.add_argument("--out", default="dynamics_pairs.pt")
p.add_argument("--manifest", default="dynamics_pairs_manifest.json")
args = p.parse_args(list(argv) if argv is not None else None)
pkl_path = Path(args.pkl)
if not pkl_path.is_file():
print(f"FILE_NOT_FOUND: {pkl_path}", file=sys.stderr)
return 2
zones = [z.strip() for z in args.zones.split(",") if z.strip()]
start = _parse_day(args.start) if args.start else None
end = _parse_day(args.end) if args.end else None
try:
pairs, dts, stats = extract_pairs_from_pkl(
pkl_path, zones, args.horizon_days, start, end
)
except ImportError as e:
print(f"IMPORT_FAILED (need torch + physics_dynamics): {e}", file=sys.stderr)
return 3
print(
f"pairs={stats['n_pairs']} points={stats['n_points']} "
f"traj_used={stats['n_trajectories_used']}/{stats['n_trajectories_seen']} "
f"skipped_gap={stats['skipped_non_consecutive']} "
f"skipped_non_exact={stats['skipped_non_exact_step']} "
f"dt_values={stats['dt_values']}"
)
if not pairs:
print("NO_PAIRS", file=sys.stderr)
return 4
import torch
# Bundle pairs + dts so DynamicsTrainer.train(..., dts=...) gets the
# real gap. Legacy code that only expects a list of (curr, nxt) can still
# torch.load and take payload["pairs"].
payload = {
"pairs": pairs,
"dts": dts,
"default_dt": float(stats.get("default_dt_hint") or 5.0),
"note": "exact step_days pairs only; use dts with DynamicsTrainer",
}
torch.save(payload, args.out)
print(f"Wrote {args.out}")
manifest = {
**stats,
"out": str(args.out),
"start": start.isoformat() if start else None,
"end": end.isoformat() if end else None,
"note": (
"Belief/uncertainty are obs-derived proxies, not agent beliefs. "
"Temporal residual is a smoothness prior along forecast lead axis, "
"not spatial advection-diffusion. "
"Pairs are exact step_days only; payload includes dts for "
"DynamicsTrainer.train(..., dts=dts, default_dt=default_dt)."
),
}
with open(args.manifest, "w") as f:
json.dump(manifest, f, indent=2)
print(f"Wrote {args.manifest}")
return 0
def _self_test() -> None:
print("build_dynamics_pairs self-test (synthetic dicts)")
# Minimal offline check without pkl
pa = _forecast_precip_array({"precip_mm": tuple(range(20))}, 14)
assert pa.shape == (14,)
b, u = _belief_uncertainty_from_obs({"precip_anomaly_idx": 3.0, "quality_flag": 2})
assert 0.0 <= b <= 1.0 and 0.0 <= u <= 1.0
print(" precip pad/clip OK")
print(" belief/uncertainty proxy OK")
print("All build_dynamics_pairs self-tests passed.")
if __name__ == "__main__":
if len(sys.argv) == 1:
_self_test()
else:
raise SystemExit(main())