poincare-hyper / src /train.py
DHDRL's picture
Rename train.py to src/train.py
4c77ab8 verified
Raw
History Blame Contribute Delete
9.88 kB
"""
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.")