PDENNEval / scripts /train.py
OneScience's picture
Upload folder using huggingface_hub
ede74c0 verified
Raw
History Blame Contribute Delete
5.93 kB
from __future__ import annotations
import argparse
from pathlib import Path
import numpy as np
import torch
import torch.nn as nn
from torch.optim import Adam, LBFGS
from onescience.utils.pinnsformer_util import get_data, get_n_params, make_time_sequence
from common import (
build_model,
ensure_runtime_dirs,
initial_condition,
load_config,
project_path,
seed_everything,
select_device,
)
def init_weights(module: nn.Module) -> None:
if isinstance(module, nn.Linear):
torch.nn.init.xavier_uniform_(module.weight)
module.bias.data.fill_(0.01)
def tensorize(array: np.ndarray, device: torch.device) -> torch.Tensor:
return torch.tensor(array, dtype=torch.float32, requires_grad=True, device=device)
def prepare_tensors(cfg: dict, device: torch.device, args: argparse.Namespace):
data_cfg = cfg["data"]
x_num = int(args.x_num or data_cfg["x_num"])
t_num = int(args.t_num or data_cfg["t_num"])
num_step = int(args.num_step or data_cfg["sequence"]["num_step"])
step = float(data_cfg["sequence"]["step"])
res, b_left, b_right, b_upper, b_lower = get_data(
data_cfg["x_range"],
data_cfg["t_range"],
x_num,
t_num,
)
tensors = []
for values in (res, b_left, b_right, b_upper, b_lower):
tensors.append(tensorize(make_time_sequence(values, num_step=num_step, step=step), device))
return tuple(tensors)
def loss_components(model: nn.Module, tensors: tuple[torch.Tensor, ...], cfg: dict):
res, b_left, b_right, b_upper, b_lower = tensors
x_res, t_res = res[:, :, 0:1], res[:, :, 1:2]
x_left, t_left = b_left[:, :, 0:1], b_left[:, :, 1:2]
x_upper, t_upper = b_upper[:, :, 0:1], b_upper[:, :, 1:2]
x_lower, t_lower = b_lower[:, :, 0:1], b_lower[:, :, 1:2]
pred_res = model(x_res, t_res)
pred_left = model(x_left, t_left)
pred_upper = model(x_upper, t_upper)
pred_lower = model(x_lower, t_lower)
u_t = torch.autograd.grad(
pred_res,
t_res,
grad_outputs=torch.ones_like(pred_res),
retain_graph=True,
create_graph=True,
)[0]
rate = float(cfg["equation"]["reaction_rate"])
target_ic = initial_condition(x_left[:, 0, :], cfg)
loss_res = torch.mean((u_t - rate * pred_res * (1 - pred_res)) ** 2)
loss_bc = torch.mean((pred_upper - pred_lower) ** 2)
loss_ic = torch.mean((pred_left[:, 0, :] - target_ic) ** 2)
loss = loss_res + loss_bc + loss_ic
return loss, (loss_res, loss_bc, loss_ic)
def build_optimizer(model: nn.Module, cfg: dict):
opt_cfg = cfg["training"]["optimizer"]
name = opt_cfg["name"].lower()
if name == "adam":
return Adam(model.parameters(), lr=float(opt_cfg.get("lr", 1e-3)))
if name == "lbfgs":
return LBFGS(
model.parameters(),
lr=float(opt_cfg.get("lr", 1.0)),
max_iter=int(opt_cfg.get("max_iter", 20)),
line_search_fn=opt_cfg.get("line_search_fn", "strong_wolfe"),
)
raise ValueError(f"Unsupported optimizer: {opt_cfg['name']}")
def save_checkpoint(path: Path, model: nn.Module, cfg: dict, loss_history: list[list[float]]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
torch.save(
{
"model_state_dict": model.state_dict(),
"config": cfg,
"loss_history": loss_history,
},
path,
)
def main() -> None:
parser = argparse.ArgumentParser(description="Train PINNsformer on the 1D reaction equation.")
parser.add_argument("--config", default=None, help="Path to config.yaml.")
parser.add_argument("--epochs", type=int, default=None, help="Override training epochs.")
parser.add_argument("--x-num", type=int, default=None, help="Override x grid count.")
parser.add_argument("--t-num", type=int, default=None, help="Override t grid count.")
parser.add_argument("--num-step", type=int, default=None, help="Override pseudo-sequence length.")
parser.add_argument("--device", default=None, help="Override runtime.device.")
args = parser.parse_args()
cfg = load_config(args.config)
ensure_runtime_dirs(cfg)
seed_everything(int(cfg["runtime"]["seed"]))
device = select_device(args.device or cfg["runtime"]["device"])
tensors = prepare_tensors(cfg, device, args)
model = build_model(cfg).to(device)
model.apply(init_weights)
optimizer = build_optimizer(model, cfg)
epochs = int(args.epochs or cfg["training"]["epochs"])
print(model)
print(f"parameters: {get_n_params(model)}")
print(f"device: {device}")
loss_history: list[list[float]] = []
for epoch in range(epochs):
latest: dict[str, float] = {}
def closure():
loss, parts = loss_components(model, tensors, cfg)
optimizer.zero_grad()
loss.backward()
latest["loss"] = float(loss.detach().cpu())
latest["loss_res"] = float(parts[0].detach().cpu())
latest["loss_bc"] = float(parts[1].detach().cpu())
latest["loss_ic"] = float(parts[2].detach().cpu())
return loss
if isinstance(optimizer, LBFGS):
optimizer.step(closure)
else:
closure()
optimizer.step()
loss_history.append([latest["loss_res"], latest["loss_bc"], latest["loss_ic"], latest["loss"]])
print(
f"epoch {epoch + 1}/{epochs} "
f"loss={latest['loss']:.6f} "
f"res={latest['loss_res']:.6f} "
f"bc={latest['loss_bc']:.6f} "
f"ic={latest['loss_ic']:.6f}"
)
checkpoint = project_path(cfg["training"]["checkpoint"])
save_checkpoint(checkpoint, model, cfg, loss_history)
np.save(project_path(cfg["paths"]["loss"]), np.asarray(loss_history, dtype=np.float32))
print(f"checkpoint saved to {checkpoint}")
if __name__ == "__main__":
main()