Download scripts/train.py from OneScience-Group/PINNsformer: direct link, hf CLI and curl.
- Browser
- Download file 5.93 kB
-
https://huggingface.co/OneScience-Group/PINNsformer/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/PINNsformer/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/PINNsformer/resolve/main/scripts/train.py
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() | |