| """ |
| Optuna hyperparameter search + supervised pretraining + light RL fine-tune. |
| |
| This script intentionally trains on SYNTHETIC data only (hyperparameter |
| search does not consume real datasets — that would be wasteful and is |
| exactly the kind of scarce resource the DatasetRegistry in provenance.py |
| is meant to protect). Data provenance is fetched via the explicit |
| get_synthetic_dataset() call, never implicit. |
| |
| For training against real Well data, use src/run_full.py instead, which |
| enforces the full contract (hard-fail on missing real data, dataset-reuse |
| registry, content-addressed atomic checkpoints). |
| """ |
| from __future__ import annotations |
| import os |
| import sys |
| sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) |
|
|
| import torch |
| import torch.nn.functional as F |
| from torch.utils.data import DataLoader |
| import optuna |
| from optuna.trial import Trial |
| import numpy as np |
| from tqdm import tqdm |
|
|
| from src.data_real import get_synthetic_dataset |
| from src.normalization import FieldNormalizer |
| from src.env import MultiStepPoincareEnv |
| from src.model import MultiScaleEncoder, HierarchicalHyperbolicPredictor, HyperbolicCritic |
| from src.physics_losses import combined_physics_loss |
| from src.provenance import CheckpointStore, hash_dataset, hash_code, validate_trajectory_lengths |
|
|
| WINDOW = 4 |
| SRC_DIR = os.path.dirname(os.path.abspath(__file__)) |
|
|
|
|
| def get_device(): |
| return "cuda" if torch.cuda.is_available() else "cpu" |
|
|
|
|
| def collate_fields(batch): |
| return torch.stack([b["fields"] for b in batch]) |
|
|
|
|
| def build_normalizer(dataset, max_fit: int = 64) -> FieldNormalizer: |
| norm = FieldNormalizer(mode="zscore") |
| samples = [dataset[i]["fields"] for i in range(min(len(dataset), max_fit))] |
| data = torch.stack(samples, dim=0) |
| norm.fit(data) |
| print(f"[Norm] fitted {data.shape[0]} trajs | mean={[round(x,4) for x in norm.mean.tolist()]}") |
| return norm |
|
|
|
|
| def objective(trial: Trial, max_epochs: int = 3, n_samples: int = 128) -> float: |
| device = get_device() |
| lr = trial.suggest_float("lr", 3e-5, 1.5e-3, log=True) |
| c = trial.suggest_float("curvature", 0.3, 1.8) |
| hidden = trial.suggest_categorical("hidden", [48, 64, 96]) |
| batch_size = trial.suggest_categorical("batch_size", [8, 16]) |
| pred_steps = trial.suggest_int("pred_steps", 2, 4) |
| w_phys = trial.suggest_float("w_phys", 1e-4, 5e-2, log=True) |
|
|
| ds, _provenance = get_synthetic_dataset(max_samples=n_samples, n_steps=14) |
| validate_trajectory_lengths(ds, required_length=WINDOW + pred_steps) |
| normalizer = build_normalizer(ds, max_fit=40) |
| loader = DataLoader(ds, batch_size=batch_size, shuffle=True, collate_fn=collate_fields) |
|
|
| encoder = MultiScaleEncoder(hidden=hidden, out_dim=8) |
| model = HierarchicalHyperbolicPredictor(encoder, c=c, pred_steps=pred_steps, levels=2).to(device) |
| opt = torch.optim.Adam(model.parameters(), lr=lr) |
|
|
| model.train() |
| losses = [] |
| for _ in range(max_epochs): |
| ep, n = 0.0, 0 |
| for batch in loader: |
| B, T, C, H, W = batch.shape |
| batch = batch.to(device) |
| flat = normalizer.transform(batch.view(B * T, C, H, W)).view(B, T, C, H, W) |
| x = flat[:, :WINDOW] |
| with torch.no_grad(): |
| tgt = torch.stack([model.encode(flat[:, WINDOW + s]) for s in range(pred_steps)], dim=1) |
| pred = model(x) |
| loss_h = model.hyperbolic_loss(pred, tgt) |
| loss_p = combined_physics_loss( |
| flat[:, : WINDOW + pred_steps], w_smooth=w_phys, w_temp=w_phys, w_cons=w_phys * 0.5 |
| ) |
| loss = loss_h + loss_p |
| if not torch.isfinite(loss): |
| raise optuna.TrialPruned(f"non-finite loss: {loss}") |
| opt.zero_grad() |
| loss.backward() |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) |
| opt.step() |
| ep += loss.item() |
| n += 1 |
| if n: |
| losses.append(ep / n) |
| if not losses: |
| raise optuna.TrialPruned("no batches were trained on this trial") |
| return float(np.mean(losses)) |
|
|
|
|
| def run_optuna(n_trials: int = 8, timeout: int = 120): |
| print("=" * 64) |
| print("Optuna - hierarchical Poincare 8D + physics priors (synthetic data)") |
| print("=" * 64) |
| study = optuna.create_study(direction="minimize") |
| study.optimize(lambda t: objective(t), n_trials=n_trials, timeout=timeout) |
| print("Best value:", round(study.best_value, 5)) |
| print("Best params:", study.best_params) |
| return study |
|
|
|
|
| def supervised_pretrain(study, epochs: int = 5): |
| device = get_device() |
| params = study.best_params |
| print("\n" + "=" * 64) |
| print("Supervised multi-step pre-training (best HPs, synthetic data)") |
| print("=" * 64) |
| ds, provenance = get_synthetic_dataset(max_samples=256, n_steps=14) |
| ps = params.get("pred_steps", 3) |
| validate_trajectory_lengths(ds, required_length=WINDOW + ps) |
| normalizer = build_normalizer(ds, max_fit=64) |
| loader = DataLoader(ds, batch_size=params.get("batch_size", 8), shuffle=True, collate_fn=collate_fields) |
|
|
| encoder = MultiScaleEncoder(hidden=params.get("hidden", 64), out_dim=8) |
| model = HierarchicalHyperbolicPredictor( |
| encoder, c=params.get("curvature", 1.0), pred_steps=ps, levels=2 |
| ).to(device) |
| opt = torch.optim.Adam(model.parameters(), lr=params.get("lr", 3e-4)) |
| w_phys = params.get("w_phys", 0.01) |
|
|
| for epoch in range(epochs): |
| total, n = 0.0, 0 |
| for batch in tqdm(loader, desc=f"Pretrain {epoch+1}/{epochs}"): |
| B, T, C, H, W = batch.shape |
| batch = batch.to(device) |
| flat = normalizer.transform(batch.view(B * T, C, H, W)).view(B, T, C, H, W) |
| x = flat[:, :WINDOW] |
| with torch.no_grad(): |
| tgt = torch.stack([model.encode(flat[:, WINDOW + s]) for s in range(ps)], dim=1) |
| pred = model(x) |
| loss = model.hyperbolic_loss(pred, tgt) + combined_physics_loss( |
| flat[:, : WINDOW + ps], w_smooth=w_phys, w_temp=w_phys |
| ) |
| if not torch.isfinite(loss): |
| raise RuntimeError(f"[NON_FINITE_LOSS] loss={loss.item()} at epoch {epoch+1}") |
| opt.zero_grad() |
| loss.backward() |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) |
| opt.step() |
| total += loss.item() |
| n += 1 |
| print(f" Epoch {epoch+1} loss: {total/max(n,1):.5f}") |
|
|
| return model, normalizer, params, ds, provenance |
|
|
|
|
| def simple_rl_finetune(model, normalizer, params, episodes: int = 40): |
| print("\n" + "=" * 64) |
| print("Simple on-policy RL fine-tune (hyperbolic critic, synthetic data)") |
| print("=" * 64) |
| device = get_device() |
| ds, _provenance = get_synthetic_dataset(max_samples=64, n_steps=14) |
| env = MultiStepPoincareEnv( |
| dataset=ds, normalizer=normalizer, encoder=model.encoder, |
| poincare_module=model.poincare, window=WINDOW, horizon=3, device=device |
| ) |
| critic = HyperbolicCritic(c=params.get("curvature", 1.0)).to(device) |
| policy_head = torch.nn.Sequential( |
| torch.nn.Linear(8, 32), torch.nn.GELU(), torch.nn.Linear(32, 8) |
| ).to(device) |
| log_std = torch.nn.Parameter(torch.zeros(8, device=device) - 1.2) |
| opt_c = torch.optim.Adam(critic.parameters(), lr=1e-3) |
| opt_p = torch.optim.Adam(list(policy_head.parameters()) + [log_std], lr=1e-3) |
|
|
| returns = [] |
| for ep in range(episodes): |
| obs, _ = env.reset() |
| logps, rewards, values = [], [], [] |
| for t in range(4): |
| obs_t = torch.as_tensor(obs, device=device).unsqueeze(0) |
| with torch.no_grad(): |
| z_ball = model.poincare.expmap0(obs_t) |
| mean = policy_head(obs_t).squeeze(0) |
| std = log_std.exp() |
| dist = torch.distributions.Normal(mean, std) |
| action_raw = dist.rsample() |
| action = action_raw.clamp(-1.5, 1.5) |
| logp = dist.log_prob(action_raw).sum() |
| val = critic(z_ball) |
| next_obs, reward, term, trunc, info = env.step( |
| action.detach().squeeze(0).cpu().numpy() |
| ) |
| logps.append(logp) |
| rewards.append(reward) |
| values.append(val) |
| obs = next_obs |
| if term or trunc: |
| break |
| R = 0.0 |
| returns_ep = [] |
| for r in reversed(rewards): |
| R = r + 0.95 * R |
| returns_ep.insert(0, R) |
| returns_t = torch.tensor(returns_ep, device=device) |
| values_t = torch.stack(values).reshape(-1) |
| adv = returns_t - values_t.detach() |
| loss_p = -(torch.stack(logps) * adv).mean() |
| loss_c = F.mse_loss(values_t, returns_t) |
| opt_p.zero_grad() |
| loss_p.backward() |
| opt_p.step() |
| opt_c.zero_grad() |
| loss_c.backward() |
| opt_c.step() |
| returns.append(sum(rewards)) |
| if (ep + 1) % 10 == 0: |
| print(f" RL episode {ep+1}: return={np.mean(returns[-10:]):.3f}") |
|
|
| print("RL fine-tune finished.") |
| return model |
|
|
|
|
| if __name__ == "__main__": |
| study = run_optuna(n_trials=6, timeout=100) |
| model, normalizer, params, ds, provenance = supervised_pretrain(study, epochs=4) |
| model = simple_rl_finetune(model, normalizer, params, episodes=30) |
| dataset_hash = hash_dataset(ds, sample_cap=64) |
| code_hash = hash_code(SRC_DIR) |
| store = CheckpointStore(checkpoints_dir="checkpoints") |
| result = store.save( |
| model_state={"model": model.state_dict()}, |
| config=params, |
| dataset_hash=dataset_hash, |
| code_hash=code_hash, |
| data_provenance=provenance, |
| extra={"normalizer": normalizer.state_dict(), "optuna_best_value": study.best_value}, |
| ) |
| print(f"\n[checkpoint] {result['outcome_code']} -> {result['path']}") |
| print("Hierarchical + physics + RL search/pretrain completed.") |
|
|