File size: 20,149 Bytes
976eb45
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
"""
train_kaggle.py
================
Standalone training script for WeatherForecastEnv, meant to run on Kaggle
(free GPU, no session-length hyperparameter-search restrictions like Colab).

WHAT THIS IS FOR
----------------
This trains a single fixed-n_zones snapshot of the environment (matching
what the Optuna sweep tuned against), with a configurable step count -- run
it short first (e.g. 100k-300k steps) to validate that the environment
fixes (reward scale, belief grounding, event injection, penalty
normalization, **budget pressure**) hold up, then re-run with a larger
--steps for a real training run.

BUDGET MODES (critical for multi-zone skill)
--------------------------------------------
The visit-once action mask makes the structural episode ceiling
``n_zones + 1`` (inspect each zone once, then terminate). Setting
``max_steps`` much larger than that (the old default of 250) never forces
the agent to *choose which zone to skip*. Full-tour is then both feasible
and reward-optimal (unvisited_zone_penalty → 0 after visiting everyone).

Under that regime policies learn a fixed inspect order (empirically: always
Indramayu under triage eval) and never learn risk-conditioned allocation.

Budget modes derive max_steps from n_zones unless --max-steps is explicit:

  full    max_steps = n_zones + 1   # can visit everyone (legacy behaviour)
  scarce  max_steps = n_zones      # can visit all only by skipping terminate
  triage  max_steps = max(1, n_zones - 1)  # MUST leave ≥1 zone unvisited

Default is **triage**. If you force max_steps >= n_zones+1, a hard WARNING
is logged: triage skill will not be trained.

This is intentionally NOT the 5-phase curriculum in train_curriculum.py --
it trains one fixed configuration end to end. Once you're confident in the
environment at this scale, train_curriculum.py's phase progression is the
next step (and now also carries budget_mode per phase).

POLICY (v3 equivariant)
-----------------------
Always use ZoneEquivariantMaskablePolicy from gru_weather_policy.py — NOT
the string "MultiInputPolicy". The latter builds logits from a pooled
feature vector and cannot express risk-conditioned zone choice (always
slot-0 degeneracy under triage). ZoneEquivariantMaskablePolicy scores each
zone before pooling and concatenates a terminate logit.

INPUT NORMALIZATION (--precip-scale)
-------------------------------------
forecast_precip runs roughly [0, 80] while zone_belief runs roughly
[0, 0.3], with no normalization layer between them and the GRU extractor.
Left unscaled, precip's larger raw magnitude can suppress the smaller but
more reliable belief signal during optimization, independent of which
feature actually carries more information. --precip-scale (default 40.0)
divides forecast_precip before it reaches the extractor. 40.0 is the value
that produced the validated single-dirty selection-accuracy results (see
the model card); it is confirmed working, not confirmed optimal --
recalibrate against your own separability probe if your event-injection
magnitudes differ.

KAGGLE SETUP
------------
1. Upload project files as a Kaggle Dataset (flat folder is fine).
2. Create a notebook, attach the dataset, enable GPU if available.
3. Example (triage, 2 zones, the recipe validated to produce single-dirty
   selection accuracy well above chance -- see model card)::

       !pip install -q gymnasium stable-baselines3 sb3-contrib tensorboard
       !PYTHONPATH=/kaggle/input/datasets/dhmmmreally/weather-modeller:$PYTHONPATH \\
       python train_kaggle.py \\
           --out /kaggle/working/run_nz2_triage \\
           --n-zones 2 \\
           --budget-mode triage \\
           --steps 150000 \\
           --clean-episode-ratio 0.80 \\
           --event-spatial-correlation 0.50 \\
           --precip-scale 40.0 \\
           --device auto

4. Resume mid-run with --resume-from pointing at a checkpoint under --out.

WHAT TO LOOK AT AFTERWARDS
---------------------------
  - ep_len_mean: under triage should sit near max_steps (not n_zones+1).
  - entropy_loss: should decline gradually, not collapse to ~0 immediately.
  - eval/mean_reward: should trend upward.
  - After training, run the zone-selection probe under max_steps=1: inspected
    pattern must NOT be 100% one action index; dirty-day shares should move
    off ~50/50 if allocation skill was learned.
"""

from __future__ import annotations

import argparse
import json
import logging
import sys
import time
from pathlib import Path
from typing import Optional, Tuple

logger = logging.getLogger("train_kaggle")
logging.basicConfig(
    level=logging.INFO,
    format="%(asctime)s | %(levelname)s | %(message)s",
)


def _add_file_logging(out_dir: Path) -> None:
    """Attach a FileHandler rooted at out_dir/training.log."""
    formatter = logging.Formatter("%(asctime)s | %(levelname)s | %(message)s")
    file_handler = logging.FileHandler(str(out_dir / "training.log"))
    file_handler.setFormatter(formatter)
    logging.getLogger().addHandler(file_handler)
    logger.info("File logging enabled: %s", out_dir / "training.log")


# ---------------------------------------------------------------------------
# Best hyperparameters found by the Optuna sweep (trial 6 of the 20-trial
# run against n_zones=3 / max_steps=250). Used as defaults below -- override
# any of them from the CLI if you want to explore further.
# ---------------------------------------------------------------------------
BEST_HYPERPARAMETERS = dict(
    learning_rate=6.916624987609979e-05,
    ent_coef=0.08779238696445962,
    hidden_size=128,
    spatial_size=12,
    n_steps=4096,
)


def resolve_max_steps(
    n_zones: int,
    budget_mode: str,
    max_steps_arg: Optional[int],
) -> Tuple[int, str]:
    """
    Derive episode max_steps from budget mode, or honour an explicit override.

    Returns (max_steps, resolved_mode_note).
    """
    n = max(1, int(n_zones))
    full_ceiling = n + 1

    mode = (budget_mode or "triage").strip().lower()
    if mode not in ("full", "scarce", "triage"):
        raise ValueError(
            f"budget_mode must be one of full|scarce|triage, got {budget_mode!r}"
        )

    if mode == "full":
        derived = full_ceiling
    elif mode == "scarce":
        derived = n
    else:  # triage
        derived = max(1, n - 1)

    if max_steps_arg is not None and int(max_steps_arg) > 0:
        ms = int(max_steps_arg)
        note = f"explicit --max-steps={ms} (budget-mode={mode} would have been {derived})"
        if ms >= full_ceiling:
            logger.warning(
                "BUDGET WARNING: max_steps=%d >= n_zones+1=%d. Full tour is "
                "feasible and reward-optimal (unvisited_zone_penalty → 0). "
                "The policy is NOT forced to differentiate which zone to "
                "inspect. For allocation skill use --budget-mode triage "
                "(or scarce) without overriding --max-steps, or set "
                "--max-steps < %d.",
                ms, full_ceiling, full_ceiling,
            )
        return ms, note

    return derived, f"budget-mode={mode} → max_steps={derived} (full ceiling={full_ceiling})"


def _add_dataset_to_path(dataset_dir: Optional[str]) -> None:
    """Make the uploaded project files importable."""
    if dataset_dir:
        p = str(Path(dataset_dir).resolve())
        if p not in sys.path:
            sys.path.insert(0, p)
            logger.info("Added to sys.path: %s", p)
    here = str(Path(__file__).resolve().parent)
    if here not in sys.path:
        sys.path.insert(0, here)


def _parse_args() -> argparse.Namespace:
    p = argparse.ArgumentParser(
        description="Single-config MaskablePPO trainer for WeatherForecastEnv (Kaggle)."
    )
    p.add_argument("--dataset-dir", default=None,
                    help="Path to the uploaded Kaggle dataset directory containing the .py files.")
    p.add_argument("--out", default="./run", help="Output directory for checkpoints/logs/final model.")
    p.add_argument("--resume-from", default=None, help="Path to a checkpoint .zip to resume from.")

    p.add_argument("--steps", type=int, default=50_000,
                    help="Total training timesteps. Start small (e.g. 100_000-300_000) to validate, "
                         "then scale up for a real run.")
    p.add_argument("--n-zones", type=int, default=2,
                    help="Number of zones. Default 2 (matches real-eval design window).")
    p.add_argument(
        "--budget-mode",
        choices=("full", "scarce", "triage"),
        default="triage",
        help="How tight the inspection budget is relative to n_zones. "
             "full=n_zones+1, scarce=n_zones, triage=max(1,n_zones-1). "
             "Default triage forces leaving ≥1 zone unvisited.",
    )
    p.add_argument(
        "--max-steps",
        type=int,
        default=None,
        help="Explicit episode length cap. If omitted, derived from --budget-mode. "
             "Setting this >= n_zones+1 disables triage pressure (WARNING logged).",
    )
    p.add_argument("--seed", type=int, default=42)

    p.add_argument("--lr", type=float, default=BEST_HYPERPARAMETERS["learning_rate"])
    p.add_argument("--ent-coef", type=float, default=BEST_HYPERPARAMETERS["ent_coef"])
    p.add_argument("--hidden-size", type=int, default=BEST_HYPERPARAMETERS["hidden_size"])
    p.add_argument("--spatial-size", type=int, default=BEST_HYPERPARAMETERS["spatial_size"])
    p.add_argument("--n-steps", type=int, default=BEST_HYPERPARAMETERS["n_steps"],
                    help="PPO rollout buffer size. If you change max_steps a lot, consider "
                         "resizing this to roughly 15-25x max_steps.")
    p.add_argument("--batch-size", type=int, default=None,
                    help="Defaults to max(32, n_steps // 8) if not given.")

    p.add_argument("--eval-freq", type=int, default=10_000)
    p.add_argument("--eval-episodes", type=int, default=20)
    p.add_argument("--checkpoint-freq", type=int, default=25_000)
    p.add_argument("--regression-check-freq", type=int, default=5_000)

    p.add_argument("--device", default="auto", help="'cpu', 'cuda', or 'auto'.")
    p.add_argument(
        "--clean-episode-ratio",
        type=float,
        default=0.90,
        help="Fraction of synthetic episodes with no regional hazard (match eval).",
    )
    p.add_argument(
        "--event-spatial-correlation",
        type=float,
        default=0.85,
        help="P(zone dirty | regional event). High → El Niño-style joint risk.",
    )
    p.add_argument(
        "--precip-scale",
        type=float,
        default=40.0,
        help="Fixed (non-learned) divisor applied to forecast_precip before "
             "it reaches the GRU extractor's forecast_proj/zone_encoder. "
             "forecast_precip runs roughly [0, 80] while zone_belief runs "
             "roughly [0, 0.3] with no normalization layer between them; "
             "left unscaled, precip's raw magnitude can suppress the "
             "smaller-but-more-reliable belief signal during optimization, "
             "independent of which feature is actually more informative. "
             "40.0 is the value that produced the validated single-dirty "
             "selection-accuracy results (see model card) -- confirmed "
             "working, not confirmed optimal. Recalibrate against your own "
             "separability probe's dirty-zone precip levels if your event "
             "injection magnitudes differ from the defaults.",
    )
    return p.parse_args()


def build_envs(args: argparse.Namespace):
    from weather_forecast_env import make_weather_env
    from zone_observation import ForecastConfig
    from stable_baselines3.common.monitor import Monitor

    train_config = ForecastConfig(
        n_zones=args.n_zones,
        max_steps=args.max_steps,
        seed=args.seed,
        clean_episode_ratio=args.clean_episode_ratio,
        event_spatial_correlation=args.event_spatial_correlation,
    )
    eval_config = ForecastConfig(
        n_zones=args.n_zones,
        max_steps=args.max_steps,
        seed=args.seed + 10_000,
        clean_episode_ratio=args.clean_episode_ratio,
        event_spatial_correlation=args.event_spatial_correlation,
    )
    train_env = Monitor(make_weather_env(train_config))
    eval_env = Monitor(make_weather_env(eval_config))
    return train_env, eval_env


def build_model(args: argparse.Namespace, train_env):
    from sb3_contrib import MaskablePPO
    from gru_weather_policy import (
        ZoneEquivariantMaskablePolicy,
        create_gru_weather_policy_kwargs,
    )

    if args.resume_from:
        logger.info("Resuming from checkpoint: %s", args.resume_from)
        return MaskablePPO.load(args.resume_from, env=train_env, device=args.device)

    policy_kwargs = create_gru_weather_policy_kwargs(
        hidden_size=args.hidden_size,
        spatial_output_size=args.spatial_size,
        features_dim=args.hidden_size * 2,
        basin_context_hidden=12,
        precip_scale=args.precip_scale,
    )
    batch_size = args.batch_size or max(32, args.n_steps // 8)

    try:
        import tensorboard  # noqa: F401
        tb_log = str(Path(args.out) / "tensorboard")
    except ImportError:
        logger.warning(
            "tensorboard not installed -- continuing without TensorBoard logs "
            "(install with `pip install tensorboard` if you want them)."
        )
        tb_log = None

    # ZoneEquivariantMaskablePolicy: inspect logits from per-zone scores
    # (pre-pool); terminate from pooled features. Do not use "MultiInputPolicy"
    # — that routes through the pooled action_net and erases slot identity.
    return MaskablePPO(
        ZoneEquivariantMaskablePolicy,
        train_env,
        learning_rate=args.lr,
        ent_coef=args.ent_coef,
        policy_kwargs=policy_kwargs,
        n_steps=args.n_steps,
        batch_size=batch_size,
        gamma=0.98,
        gae_lambda=0.95,
        clip_range=0.2,
        device=args.device,
        verbose=1,
        tensorboard_log=tb_log,
    )


class RegressionWatchCallback:
    """Flags signs of the old degenerate 'terminate immediately' collapse.

    Under triage, ep_len pinned near 1.0 with no belief movement is still a
    failure mode (never inspect). Under full budget, ep_len pinned at
    n_zones+1 is expected and not a regression by itself.
    """

    def __init__(
        self,
        total_timesteps: int,
        check_freq: int = 5_000,
        after_frac: float = 0.2,
        ep_len_threshold: float = 1.5,
    ):
        from stable_baselines3.common.callbacks import BaseCallback
        import numpy as np

        self._np = np
        self._BaseCallback = BaseCallback
        self.total_timesteps = total_timesteps
        self.check_freq = check_freq
        self.after_frac = after_frac
        self.ep_len_threshold = ep_len_threshold
        self._instance = self._build_instance()

    def _build_instance(self):
        np = self._np
        outer = self

        class _Impl(self._BaseCallback):
            def __init__(self):
                super().__init__()
                self._last_check = 0
                self.history = []  # (timestep, ep_len_mean, entropy_loss)

            def _on_step(self) -> bool:
                if self.num_timesteps - self._last_check < outer.check_freq:
                    return True
                self._last_check = self.num_timesteps

                ep_lens = (
                    [ep["l"] for ep in self.model.ep_info_buffer]
                    if self.model.ep_info_buffer
                    else []
                )
                ep_len_mean = float(np.mean(ep_lens)) if ep_lens else float("nan")

                entropy = None
                if self.model.logger is not None:
                    entropy = self.model.logger.name_to_value.get("train/entropy_loss")

                self.history.append((self.num_timesteps, ep_len_mean, entropy))

                frac = self.num_timesteps / max(1, outer.total_timesteps)
                if frac >= outer.after_frac and ep_lens and ep_len_mean < outer.ep_len_threshold:
                    logger.warning(
                        "REGRESSION WARNING at step %d: ep_len_mean=%.2f after %.0f%% "
                        "of training. This matches the original 'terminate immediately' "
                        "collapse pattern -- worth stopping to check config/reward before "
                        "trusting the rest of this run.",
                        self.num_timesteps, ep_len_mean, frac * 100,
                    )
                return True

        return _Impl()

    @property
    def instance(self):
        return self._instance


def run_training(args: argparse.Namespace) -> None:
    # Resolve budget before anything else so logs and envs agree.
    resolved_ms, budget_note = resolve_max_steps(
        args.n_zones, args.budget_mode, args.max_steps
    )
    args.max_steps = resolved_ms  # mutate so build_envs / summary see the real value

    out_dir = Path(args.out)
    out_dir.mkdir(parents=True, exist_ok=True)
    _add_file_logging(out_dir)

    logger.info(
        "Config: n_zones=%d max_steps=%d (%s) steps=%d clean=%.3f rho=%.3f "
        "lr=%.3g ent_coef=%.3g hidden_size=%d spatial_size=%d n_steps=%d "
        "precip_scale=%.3g device=%s",
        args.n_zones, args.max_steps, budget_note, args.steps,
        args.clean_episode_ratio, args.event_spatial_correlation,
        args.lr, args.ent_coef,
        args.hidden_size, args.spatial_size, args.n_steps,
        args.precip_scale,
        args.device,
    )
    logger.info(
        "Budget pressure: full_ceiling=%d  resolved_max_steps=%d  "
        "must_skip_zones=%s",
        args.n_zones + 1,
        args.max_steps,
        "yes" if args.max_steps < args.n_zones + 1 else "no (full tour allowed)",
    )

    train_env, eval_env = build_envs(args)
    model = build_model(args, train_env)

    from sb3_contrib.common.maskable.callbacks import MaskableEvalCallback
    try:
        from train_curriculum import CheckpointCallback as _ProjectCheckpointCallback
        checkpoint_cb = _ProjectCheckpointCallback(out_dir, save_freq=args.checkpoint_freq)
    except Exception as e:
        logger.warning(
            "Could not import CheckpointCallback from train_curriculum.py (%s); "
            "continuing without periodic checkpoints -- only the final model will "
            "be saved.", e
        )
        checkpoint_cb = None

    eval_cb = MaskableEvalCallback(
        eval_env,
        n_eval_episodes=args.eval_episodes,
        eval_freq=args.eval_freq,
        deterministic=True,
        best_model_save_path=str(out_dir / "best_model"),
        verbose=1,
    )

    regression_watch = RegressionWatchCallback(
        total_timesteps=args.steps, check_freq=args.regression_check_freq,
    )

    callbacks = [c for c in [checkpoint_cb, eval_cb, regression_watch.instance] if c is not None]

    t0 = time.time()
    model.learn(total_timesteps=args.steps, callback=callbacks, progress_bar=False)
    elapsed = time.time() - t0
    logger.info("Training finished in %.1f minutes.", elapsed / 60.0)

    final_path = out_dir / "final_model.zip"
    model.save(str(final_path))
    logger.info("Saved final model: %s", final_path)

    summary = {
        "args": vars(args),
        "budget_note": budget_note,
        "full_tour_ceiling": args.n_zones + 1,
        "must_skip_zones": args.max_steps < args.n_zones + 1,
        "elapsed_minutes": elapsed / 60.0,
        "best_mean_reward": eval_cb.best_mean_reward,
        "regression_check_history": regression_watch.instance.history,
    }
    summary_path = out_dir / "run_summary.json"
    with open(summary_path, "w") as f:
        json.dump(summary, f, indent=2, default=str)
    logger.info("Saved run summary: %s", summary_path)
    logger.info("Best mean eval reward this run: %s", eval_cb.best_mean_reward)


def main() -> None:
    args = _parse_args()
    _add_dataset_to_path(args.dataset_dir)
    run_training(args)


if __name__ == "__main__":
    main()