| """ |
| Demonstration of continual learning on successive scientific domains |
| using the hierarchical Poincaré model + Replay + EWC. |
| |
| Uses the exact best hyperparameters from the long Optuna study: |
| lr=3.82e-4, curvature=0.455, hidden=96, batch_size=8, |
| pred_steps=4, w_phys=9.6e-4, levels=2 |
| """ |
| from __future__ import annotations |
| import os, sys, copy |
| sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) |
|
|
| import torch |
| from torch.utils.data import DataLoader |
| import numpy as np |
| from tqdm import tqdm |
|
|
| from src.normalization import FieldNormalizer |
| from src.synthetic_fields import SyntheticWellLike |
| from src.model import MultiScaleEncoder, HierarchicalHyperbolicPredictor |
| from src.physics_losses import combined_physics_loss |
| from src.continual import ReplayBuffer, DiagonalEWC, hyperbolic_distillation_loss, fit_normalizer_for_domain |
| from src.config import BEST_HPARAMS as BEST |
|
|
| def collate(batch): |
| return torch.stack([b["fields"] for b in batch]) |
|
|
| def make_domain(seed: int, n_channels: int = 2, n_samples=96, n_steps=14): |
| """Slightly different synthetic regimes act as successive scientific domains.""" |
| torch.manual_seed(seed) |
| return SyntheticWellLike(n_samples=n_samples, n_steps=n_steps, height=32, width=32, |
| n_channels=n_channels, noise=0.12 + 0.04*(seed%3)) |
|
|
| def evaluate(model, norm, ds, device, pred_steps=None): |
| if pred_steps is None: |
| pred_steps = model.pred_steps |
| model.eval() |
| loader = DataLoader(ds, batch_size=8, collate_fn=collate) |
| 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) |
| win = 4 |
| if T < win + pred_steps: |
| continue |
| x = flat[:, :win] |
| tgt = torch.stack([model.encode(flat[:, win+s]) for s in range(pred_steps)], 1) |
| pred = model(x) |
| losses.append(model.hyperbolic_loss(pred, tgt).item()) |
| model.train() |
| return float(np.mean(losses)) if losses else 1e6 |
|
|
| def train_domain(model, norm, ds, opt, device, epochs=5, replay: ReplayBuffer=None, |
| ewc: DiagonalEWC=None, teacher=None, mix_replay=0.4): |
| loader = DataLoader(ds, batch_size=BEST["batch_size"], shuffle=True, collate_fn=collate) |
| w_phys = BEST["w_phys"] |
| ps = model.pred_steps |
| probe = ds[0]["fields"] |
| domain_c = int(probe.shape[1]) |
| domain_hw = (int(probe.shape[2]), int(probe.shape[3])) |
| 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 if batch.dim()==5 else B*batch.size(1), C, H, W)) |
| if batch.dim() == 5: |
| flat = flat.view(B, T, C, H, W) |
| else: |
| flat = flat.view(B, -1, C, H, W) |
| T = flat.size(1) |
| win = 4 |
| if T < win + ps: |
| continue |
| x = flat[:, :win] |
| with torch.no_grad(): |
| tgt = torch.stack([model.encode(flat[:, win+s]) for s in range(ps)], 1) |
| pred = model(x) |
| loss = model.hyperbolic_loss(pred, tgt) |
| loss = loss + combined_physics_loss(flat[:, :win+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] if x.dim()==5 else x) |
| s_lat = model.encode(x[:, -1] if x.dim()==5 else x) |
| loss = loss + 0.1 * hyperbolic_distillation_loss(s_lat, t_lat, model.poincare) |
| 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(): |
| device = "cpu" |
| print("=" * 64) |
| print("Continual learning demo – Replay + EWC on Poincaré hierarchical model") |
| print("Using Optuna best HPs:", BEST) |
| print("=" * 64) |
|
|
| domain_specs = [ |
| {"seed": 11, "n_channels": 2}, |
| {"seed": 22, "n_channels": 5}, |
| {"seed": 33, "n_channels": 2}, |
| ] |
| domains = [make_domain(seed=s["seed"], n_channels=s["n_channels"]) for s in domain_specs] |
|
|
| 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 = [] |
|
|
| history = {f"domain_{i}": [] for i in range(3)} |
|
|
| for d_idx, ds in enumerate(domains): |
| c = domain_specs[d_idx]["n_channels"] |
| print(f"\n--- Training domain {d_idx+1}/3 (C={c}) ---") |
| norm = fit_normalizer_for_domain(ds, max_fit=40) |
| normalizers.append(norm) |
|
|
| train_domain(model, norm, ds, opt, device, epochs=4, |
| replay=replay, |
| ewc=ewc if d_idx > 0 else None, |
| teacher=teacher, mix_replay=0.0 if d_idx == 0 else 0.35) |
| print(f" replay buffer by channel count: {replay.counts_by_channels()}") |
| ewc.finalize_domain(model) |
| teacher = copy.deepcopy(model).eval() |
| for p in teacher.parameters(): |
| p.requires_grad = False |
|
|
| for j in range(d_idx + 1): |
| loss_j = evaluate(model, normalizers[j], domains[j], device) |
| history[f"domain_{j}"].append(loss_j) |
| print(f" Eval domain {j+1} (C={domain_specs[j]['n_channels']}) loss: {loss_j:.4f}") |
|
|
| print("\n" + "=" * 64) |
| print("Retention summary (loss after each successive domain)") |
| for k, v in history.items(): |
| print(f" {k}: {[round(x,4) for x in v]}") |
|
|
| os.makedirs("logs", exist_ok=True) |
| path = "logs/poincare8d_continual.pt" |
| torch.save({ |
| "model": model.state_dict(), |
| "params": BEST, |
| "normalizers": [n.state_dict() for n in normalizers], |
| "domain_specs": domain_specs, |
| "history": history, |
| }, path) |
| print(f"\nSaved {path}") |
| print("Continual learning module integrated and demonstrated (cross-channel-count).") |
|
|
| if __name__ == "__main__": |
| main() |
|
|