File size: 10,558 Bytes
976eb45
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2ff3e89
976eb45
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2ff3e89
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
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
"""
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}")