| """ |
| 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]) |
| 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 |
|
|
| |
| 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() |
|
|