""" Real multi-stream Well continual-learning driver. Kept SEPARATE from continual_demo.py rather than adding a --real-streams flag to it: this keeps the synthetic demo stable as a fast, network-free smoke test, and makes the real-data claim of THIS script explicit. Prerequisites: - MultiScaleEncoder: channel-agnostic (model.py) - well_sample_to_fields / WellStreamAdapter (well_adapter.py) - fit_normalizer_for_domain + ReplayBuffer.sample(channels=, spatial=) - WellStreamDataset: hard-fails on real-stream failure (env.py) Live HF multi-stream + retention verified on Kaggle (gray_scott_reaction_diffusion C=2 → active_matter C=11 → shear_flow C=4). Run: python -m src.run_multistream \\ --datasets gray_scott_reaction_diffusion active_matter shear_flow \\ --max-samples 96 --epochs-per-domain 3 """ from __future__ import annotations import argparse import copy import os import torch from torch.utils.data import DataLoader import numpy as np from src.env import WellStreamDataset from src.well_adapter import WellStreamAdapter from src.model import MultiScaleEncoder, HierarchicalHyperbolicPredictor from src.physics_losses import combined_physics_loss from src.continual import ReplayBuffer, DiagonalEWC, fit_normalizer_for_domain from src.config import BEST_HPARAMS as BEST from src.provenance import DataLoadError WINDOW = 4 def collate(batch): return torch.stack([b["fields"] for b in batch]) def evaluate(model, norm, ds, device): model.eval() loader = DataLoader(ds, batch_size=8, collate_fn=collate) ps = model.pred_steps losses = [] with torch.no_grad(): for batch in loader: B, T, C, H, W = batch.shape batch = batch.to(device) flat = norm.transform(batch.view(B * T, C, H, W)).view(B, T, C, H, W) if T < WINDOW + ps: continue x = flat[:, :WINDOW] tgt = torch.stack([model.encode(flat[:, WINDOW + s]) for s in range(ps)], 1) pred = model(x) losses.append(model.hyperbolic_loss(pred, tgt).item()) model.train() return float(np.mean(losses)) if losses else float("nan") def load_real_domain(dataset_name: str, split: str, max_samples: int): stream = WellStreamDataset( dataset_name=dataset_name, split=split, n_steps_input=WINDOW, n_steps_output=BEST["pred_steps"], max_samples=max_samples, allow_synthetic_fallback=False, ) if stream.provenance != "REAL_STREAMED": raise DataLoadError( f"expected REAL_STREAMED provenance for {dataset_name}, got " f"{stream.provenance!r}", outcome_code="UNEXPECTED_PROVENANCE", ) return WellStreamAdapter(stream, include_output=True) def train_one_domain(model, norm, ds, opt, device, epochs, replay, ewc, teacher, mix_replay): loader = DataLoader(ds, batch_size=BEST["batch_size"], shuffle=True, collate_fn=collate) ps = model.pred_steps w_phys = BEST["w_phys"] probe = ds[0]["fields"] domain_c = int(probe.shape[1]) # fields are (T, C, H, W) domain_hw = (int(probe.shape[-2]), int(probe.shape[-1])) for ep in range(epochs): for batch in loader: B, T, C, H, W = batch.shape batch = batch.to(device) if replay is not None and len(replay) > 0 and np.random.rand() < mix_replay: old = replay.sample(max(1, B // 2), channels=domain_c, spatial=domain_hw) if old is not None: old = old.to(device) tmin = min(old.size(1), T) batch = torch.cat([batch[:, :tmin], old[:, :tmin]], dim=0) B = batch.size(0) T = tmin flat = norm.transform(batch.view(B * T, C, H, W)).view(B, T, C, H, W) if T < WINDOW + ps: raise RuntimeError( f"[TRAJECTORY_TOO_SHORT] domain trajectory length {T} < " f"required {WINDOW + ps} -- validate before training, " f"not mid-loop." ) x = flat[:, :WINDOW] with torch.no_grad(): tgt = torch.stack([model.encode(flat[:, WINDOW + s]) for s in range(ps)], 1) pred = model(x) loss = model.hyperbolic_loss(pred, tgt) loss = loss + combined_physics_loss(flat[:, :WINDOW + ps], w_smooth=w_phys, w_temp=w_phys) if ewc is not None: loss = loss + ewc.ewc_loss(model) if teacher is not None: with torch.no_grad(): t_lat = teacher.encode(x[:, -1]) s_lat = model.encode(x[:, -1]) from src.continual import hyperbolic_distillation_loss loss = loss + 0.1 * hyperbolic_distillation_loss(s_lat, t_lat, model.poincare) if not torch.isfinite(loss): raise RuntimeError(f"[NON_FINITE_LOSS] loss={loss.item()} during domain training") opt.zero_grad() loss.backward() if ewc is not None and np.random.rand() < 0.25: ewc.accumulate_fisher(model, loss) torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) opt.step() if replay is not None: for i in range(min(8, len(ds))): replay.add(ds[i]["fields"]) def main(): p = argparse.ArgumentParser() p.add_argument("--datasets", nargs="+", required=True, help="Real Well dataset names to stream in order, e.g. " "gray_scott_reaction_diffusion active_matter shear_flow") p.add_argument("--split", default="train") p.add_argument("--max-samples", type=int, default=128) p.add_argument("--epochs-per-domain", type=int, default=4) p.add_argument("--mix-replay", type=float, default=0.35) args = p.parse_args() device = "cpu" print("=" * 64) print(f"Real multi-stream continual run: {args.datasets}") print("Every domain hard-fails if it cannot stream real data -- no synthetic fallback.") print("=" * 64) enc = MultiScaleEncoder(hidden=BEST["hidden"], out_dim=8) model = HierarchicalHyperbolicPredictor( enc, c=BEST["curvature"], pred_steps=BEST["pred_steps"], levels=BEST["levels"] ).to(device) opt = torch.optim.Adam(model.parameters(), lr=BEST["lr"]) replay = ReplayBuffer(capacity=128) ewc = DiagonalEWC(model, lambda_ewc=500.0) teacher = None normalizers, datasets_loaded = [], [] for d_idx, name in enumerate(args.datasets): print(f"\n--- Domain {d_idx+1}/{len(args.datasets)}: {name} ---") ds = load_real_domain(name, args.split, args.max_samples) c = ds[0]["fields"].shape[1] print(f"[data] {name}: provenance={ds.provenance} C={c} n={len(ds)}") norm = fit_normalizer_for_domain(ds, max_fit=40) normalizers.append(norm) datasets_loaded.append(ds) train_one_domain(model, norm, ds, opt, device, args.epochs_per_domain, replay=replay, ewc=ewc if d_idx > 0 else None, teacher=teacher, mix_replay=0.0 if d_idx == 0 else args.mix_replay) print(f" replay buffer by channel count: {replay.counts_by_channels()}") ewc.finalize_domain(model) teacher = copy.deepcopy(model).eval() for param in teacher.parameters(): param.requires_grad = False # Retention: evaluate EVERY domain seen so far with ITS OWN normalizer print(" Retention after domain", d_idx + 1) for j in range(d_idx + 1): loss_j = evaluate(model, normalizers[j], datasets_loaded[j], device) print(f" Eval domain {j+1} ({args.datasets[j]}) loss: {loss_j:.4f}") os.makedirs("logs", exist_ok=True) torch.save({ "model": model.state_dict(), "params": BEST, "normalizers": [n.state_dict() for n in normalizers], "datasets": args.datasets, }, "logs/multistream_continual.pt") print("\nSaved logs/multistream_continual.pt") print("Done.") if __name__ == "__main__": main()