""" 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.")