monsoon-rl / hierarchical_search.py
DHDRL's picture
Update hierarchical_search.py
2ff3e89 verified
Raw
History Blame
10.6 kB
"""
hierarchical_search.py
=======================
Country-scale extension for the weather/crop-risk system. Does not modify
weather_forecast_env.py, zone_observation.py, crop_risk_scorer.py, or
gru_weather_policy.py.
Scaling problem: weather_forecast_env.py's action space is
Discrete(max_zones + 1), one zone per step -- unworkable at country scale
(20k+ zones), and gru_weather_policy.py's GRU input size is linear in
n_zones, baked into the weight matrices -- no single model works across
region sizes.
Fix: recursively bisect the map (split whichever axis -- lat or lon -- is
longer) and run the existing small env/model with exactly 2 "zones" (a
node's two children) at every level. The model never sees more than 2
zones regardless of country size. Binary branching also means this reuses
the exact model already trained/swept/fixed earlier in this project with
zero new training; a wider tree (e.g. province -> counties) would need a
fresh curriculum phase first. An administrative-hierarchy version (real
boundaries) is a natural upgrade once this scaffolding is validated, but
needs boundary data this repo doesn't have.
Tree depth maps to ForecastConfig.force_data_source (coarse/cheap ->
fine/expensive), so a clean branch prunes after one cheap fetch and
expensive fetches only happen for flagged nodes. Total fetches are capped
by `budget`, not by leaf-zone count.
Phase 1 (implemented, tested here): coarse-gate with
crop_risk_scorer.compute_risk_score() directly -- no RL/GPU in the loop.
Phase 2 (stub, NOT implemented -- see _decide_with_policy): swap in the
trained MaskablePPO model to make the recursion decision adaptively. This
is a real sim-to-real domain-transfer step, not validated here.
"""
from __future__ import annotations
import heapq
import itertools
import logging
from dataclasses import dataclass, field
from datetime import datetime, timedelta, timezone
from typing import List, Optional, Tuple
import zone_observation as _zo
assert _zo.SCHEMA_VERSION == 3, (
f"hierarchical_search: zone_observation schema mismatch "
f"(expected 3, got {_zo.SCHEMA_VERSION})"
)
from zone_observation import DataSource, ForecastConfig, GeoPolygon, AlertLevel
from era5_data_pipeline import register_zone, fetch_episode_context
from crop_risk_scorer import compute_risk_score
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Depth -> data resolution
# ---------------------------------------------------------------------------
RESOLUTION_BY_DEPTH: List[DataSource] = [
DataSource.ERA5_REANALYSIS, # depth 0-1: country / province scale
DataSource.ERA5_REANALYSIS,
DataSource.SATELLITE_PRECIP, # depth 2-3: county scale
DataSource.SATELLITE_SOIL,
DataSource.OPENMETEO_LIVE, # depth 4+: finest available
]
def _resolution_for_depth(depth: int) -> DataSource:
idx = min(depth, len(RESOLUTION_BY_DEPTH) - 1)
return RESOLUTION_BY_DEPTH[idx]
# ---------------------------------------------------------------------------
# Spatial bisection tree
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class BBox:
lat_min: float
lat_max: float
lon_min: float
lon_max: float
@property
def centroid(self) -> Tuple[float, float]:
return ((self.lat_min + self.lat_max) / 2.0, (self.lon_min + self.lon_max) / 2.0)
def split(self) -> Tuple["BBox", "BBox"]:
lat_span = self.lat_max - self.lat_min
lon_span = self.lon_max - self.lon_min
if lat_span >= lon_span:
mid = (self.lat_min + self.lat_max) / 2.0
return (
BBox(self.lat_min, mid, self.lon_min, self.lon_max),
BBox(mid, self.lat_max, self.lon_min, self.lon_max),
)
else:
mid = (self.lon_min + self.lon_max) / 2.0
return (
BBox(self.lat_min, self.lat_max, self.lon_min, mid),
BBox(self.lat_min, self.lat_max, mid, self.lon_max),
)
def to_polygon(self, zone_id: str) -> GeoPolygon:
return GeoPolygon(
zone_id=zone_id,
vertices=[
(self.lat_min, self.lon_min),
(self.lat_min, self.lon_max),
(self.lat_max, self.lon_max),
(self.lat_max, self.lon_min),
],
)
@dataclass(frozen=True)
class SpatialNode:
node_id: str
bbox: BBox
depth: int
def children(self) -> Tuple["SpatialNode", "SpatialNode"]:
left, right = self.bbox.split()
return (
SpatialNode(f"{self.node_id}_0", left, self.depth + 1),
SpatialNode(f"{self.node_id}_1", right, self.depth + 1),
)
@dataclass
class ResolvedLeaf:
node: SpatialNode
max_risk: float
alert_level: AlertLevel
depth: int
data_source: DataSource
# ---------------------------------------------------------------------------
# Node scoring (cheap, deterministic -- no RL, no GPU)
# ---------------------------------------------------------------------------
def _score_node(
node: SpatialNode,
date_range: Tuple[datetime, datetime],
config_template: ForecastConfig,
) -> Tuple[float, AlertLevel, DataSource]:
polygon = node.bbox.to_polygon(node.node_id)
register_zone(polygon)
cfg_dict = config_template.to_dict()
if config_template.force_data_source is None:
source = _resolution_for_depth(node.depth)
cfg_dict["force_data_source"] = source.value
else:
source = config_template.force_data_source
cfg = ForecastConfig.from_dict(cfg_dict)
ctx = fetch_episode_context(node.node_id, date_range, cfg)
risk = compute_risk_score(ctx.obs, ctx.forecast, cfg)
return risk.supply_shortfall_prob, risk.alert_level, source
# ---------------------------------------------------------------------------
# Phase 1: best-first bisection search, coarse-gated by compute_risk_score
# ---------------------------------------------------------------------------
def hierarchical_search(
root_bbox: BBox,
date_range: Tuple[datetime, datetime],
max_depth: int = 10,
budget: int = 500,
recurse_floor: AlertLevel = AlertLevel.WARNING,
config_template: Optional[ForecastConfig] = None,
) -> List[ResolvedLeaf]:
cfg_template = config_template or ForecastConfig()
root = SpatialNode("root", root_bbox, depth=0)
counter = itertools.count() # heapq tie-breaker; heapq is min-heap, so risk is negated for best-first
frontier: List[Tuple[float, int, SpatialNode]] = [(0.0, next(counter), root)]
results: List[ResolvedLeaf] = []
spent = 0
failed = 0
while frontier and spent < budget:
_, _, node = heapq.heappop(frontier)
spent += 1
try:
max_risk, alert, source = _score_node(node, date_range, cfg_template)
except Exception as e:
logger.warning("hierarchical_search: node %s failed to score: %s", node.node_id, e)
failed += 1
continue
should_stop = (
node.depth >= max_depth
or alert.severity() < recurse_floor.severity()
)
if should_stop:
results.append(ResolvedLeaf(node, max_risk, alert, node.depth, source))
continue
for child in node.children():
heapq.heappush(frontier, (-max_risk, next(counter), child))
if failed > 0 and failed == spent:
raise RuntimeError(
f"hierarchical_search: all {spent} scored nodes failed -- "
f"check config_template/credentials before trusting any result "
f"(including an empty one) from this run."
)
if failed > 0:
logger.warning(
"hierarchical_search: %d/%d node evaluations failed to score "
"(see warnings above) -- results below exclude those nodes, "
"they are neither 'resolved clean' nor 'pending'.", failed, spent,
)
if frontier:
logger.info(
"hierarchical_search: budget exhausted with %d nodes still pending "
"(unexplored -- not the same as 'clean')", len(frontier),
)
return results
def _decide_with_policy(node: SpatialNode, children_contexts, model, env_factory):
"""Phase 2 STUB -- not implemented, not tested.
Intended contract: node's two children become n_zones=2 "zones" for one
WeatherForecastEnv episode. Seed belief_map from the already-fetched
coarse compute_risk_score result per child (a real sim-to-real
assumption, not validated here). "Inspect zone i" = fetch child i at
the next-finer resolution and rescore; "terminate" = stop, uninspected
children keep their coarse read. Recursion only happens for children
the policy chose to inspect -- this is where RL adds value over Phase
1's fixed recurse_floor: adaptive, learned budget allocation.
Raises rather than returning a fake decision.
"""
raise NotImplementedError(
"Phase 2 (RL-driven recursion) is a design stub -- see docstring. "
"Phase 1 (hierarchical_search, using compute_risk_score directly) "
"is implemented and tested in this file."
)
# ---------------------------------------------------------------------------
# Self-test: runs fully offline (SYNTHETIC data source, no credentials needed)
# ---------------------------------------------------------------------------
if __name__ == "__main__":
logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(levelname)s | %(message)s")
indonesia_bbox = BBox(lat_min=-11.0, lat_max=6.0, lon_min=95.0, lon_max=141.0)
date_range = (
datetime(2026, 7, 1, tzinfo=timezone.utc),
datetime(2026, 7, 31, tzinfo=timezone.utc),
)
cfg = ForecastConfig(force_data_source=DataSource.SYNTHETIC)
results = hierarchical_search(
indonesia_bbox, date_range,
max_depth=8, budget=60,
recurse_floor=AlertLevel.WATCH,
config_template=cfg,
)
by_depth = {}
for r in results:
by_depth.setdefault(r.depth, []).append(r)
print(f"\n{len(results)} leaves resolved out of a budget of 60 node evaluations")
print(f"(worst case with no pruning at max_depth=8 would be up to {2**8} leaves)\n")
for depth in sorted(by_depth):
leaves = by_depth[depth]
alerts = [l.alert_level.value for l in leaves]
print(f" depth {depth}: {len(leaves)} branches stopped here -- alerts: {alerts}")