| """Train PINN-TC with Adam and an optional L-BFGS refinement under torchrun DDP.""" |
|
|
| import argparse |
| import importlib.util |
| import json |
| import os |
| import random |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| import yaml |
| from torch.nn.parallel import DistributedDataParallel |
|
|
|
|
| ROOT = Path(__file__).resolve().parents[1] |
|
|
|
|
| def load_model_module(): |
| spec = importlib.util.spec_from_file_location("pinn_tc_model", ROOT / "model/pinn-tc.py") |
| module = importlib.util.module_from_spec(spec) |
| spec.loader.exec_module(module) |
| return module |
|
|
|
|
| def choose_device(config, local_rank): |
| requested = config["runtime"]["device"] |
| if requested == "auto": |
| return torch.device("cuda", local_rank) if torch.cuda.is_available() else torch.device("cpu") |
| return torch.device(requested) |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--resume", action="store_true", help="resume optimizer and model state from the configured checkpoint") |
| args = parser.parse_args() |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) |
| module = load_model_module() |
| seed = int(config["seed"]) |
| random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) |
| distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1 |
| local_rank = int(os.environ.get("LOCAL_RANK", "0")) |
| if distributed: |
| torch.distributed.init_process_group("nccl" if torch.cuda.is_available() else "gloo") |
| rank = torch.distributed.get_rank() if distributed else 0 |
| device = choose_device(config, local_rank) |
| if device.type == "cuda": |
| torch.cuda.set_device(device); torch.cuda.manual_seed_all(seed) |
|
|
| source = np.load(ROOT / config["data"]["root"] / "training_points.npz") |
| observations = torch.from_numpy(source["observation_coordinates"]).to(device) |
| targets = torch.from_numpy(source["observation_targets"]).to(device) |
| collocation = torch.from_numpy(source["collocation_coordinates"]).to(device) |
| model = module.PINNTC(**config["model"]).to(device) |
| wrapped = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None) if distributed else model |
| bare = wrapped.module if distributed else wrapped |
| gamma = float(config["train"]["gamma"]) |
| batch_observations = int(config["train"]["observation_batch"]) |
| batch_collocation = int(config["train"]["collocation_batch"]) |
|
|
| def loss_for(obs_index, col_index): |
| prediction = wrapped(observations[obs_index])[:, :3] |
| speed = torch.sqrt(targets[obs_index, 0] ** 2 + targets[obs_index, 1] ** 2) |
| wind_weight = 1.0 + float(config["train"]["wind_speed_weight"]) * speed / 50.0 |
| normalized_error = (prediction - targets[obs_index]) / torch.tensor([45.0, 45.0, 6000.0], device=device) |
| data_loss = (normalized_error.square() * wind_weight[:, None]).mean() |
| points = collocation[col_index].detach().requires_grad_(True) |
| ru, rv, rc = module.pde_residuals(wrapped, points, **config["physics"]) |
| pde_loss = (ru.square().mean() + rv.square().mean()) / float(config["train"]["momentum_scale"]) ** 2 |
| pde_loss = pde_loss + rc.square().mean() / float(config["train"]["continuity_scale"]) ** 2 |
| return data_loss + gamma * pde_loss, data_loss, pde_loss |
|
|
| optimizer = torch.optim.Adam(wrapped.parameters(), lr=float(config["train"]["adam_lr"])) |
| history = [] |
| checkpoint_path = ROOT / config["paths"]["checkpoint"] |
| completed_adam_steps = 0 |
| checkpoint = None |
| if args.resume: |
| checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False) |
| if checkpoint["model_config"] != config["model"]: |
| raise ValueError("checkpoint model configuration mismatch") |
| bare.load_state_dict(checkpoint["model_state_dict"]) |
| optimizer.load_state_dict(checkpoint["adam_optimizer_state_dict"]) |
| history = list(checkpoint.get("history", [])) |
| completed_adam_steps = int(checkpoint.get("completed_adam_steps", 0)) |
| generator = torch.Generator(device="cpu").manual_seed(seed + rank) |
| for step in range(completed_adam_steps, completed_adam_steps + int(config["train"]["adam_steps"])): |
| obs_index = torch.randperm(len(observations), generator=generator)[:batch_observations].to(device) |
| col_index = torch.randperm(len(collocation), generator=generator)[:batch_collocation].to(device) |
| optimizer.zero_grad(set_to_none=True) |
| loss, data_loss, pde_loss = loss_for(obs_index, col_index) |
| loss.backward(); optimizer.step() |
| history.append({"stage": "Adam", "step": step + 1, "loss": float(loss.detach()), |
| "data_loss": float(data_loss.detach()), "pde_loss": float(pde_loss.detach())}) |
|
|
| lbfgs_steps = int(config["train"]["lbfgs_steps"]) |
| optimizer_lbfgs = None |
| if lbfgs_steps > 0: |
| |
| obs_index = torch.arange(min(batch_observations, len(observations)), device=device) |
| col_index = torch.arange(min(batch_collocation, len(collocation)), device=device) |
| optimizer_lbfgs = torch.optim.LBFGS(wrapped.parameters(), lr=float(config["train"]["lbfgs_lr"]), |
| max_iter=lbfgs_steps, history_size=int(config["train"]["lbfgs_history_size"]), |
| line_search_fn=None) |
| if checkpoint is not None and checkpoint.get("lbfgs_optimizer_state_dict") is not None: |
| optimizer_lbfgs.load_state_dict(checkpoint["lbfgs_optimizer_state_dict"]) |
| latest = {} |
|
|
| def closure(): |
| optimizer_lbfgs.zero_grad(set_to_none=True) |
| loss, data_loss, pde_loss = loss_for(obs_index, col_index) |
| loss.backward() |
| latest.update(loss=float(loss.detach()), data_loss=float(data_loss.detach()), pde_loss=float(pde_loss.detach())) |
| return loss |
|
|
| optimizer_lbfgs.step(closure) |
| history.append({"stage": "L-BFGS", "step": 1, **latest}) |
|
|
| if rank == 0: |
| metrics_path = ROOT / config["paths"]["training_metrics"] |
| checkpoint_path.parent.mkdir(parents=True, exist_ok=True); metrics_path.parent.mkdir(parents=True, exist_ok=True) |
| torch.save({"model": bare.state_dict(), "model_state_dict": bare.state_dict(), |
| "model_config": config["model"], "format_version": config["data"]["format_version"], |
| "adam_optimizer_state_dict": optimizer.state_dict(), |
| "lbfgs_optimizer_state_dict": optimizer_lbfgs.state_dict() if optimizer_lbfgs is not None else None, |
| "completed_adam_steps": completed_adam_steps + int(config["train"]["adam_steps"]), |
| "history": history, |
| "input_order": list(module.INPUT_ORDER), "output_order": list(module.OUTPUT_ORDER), |
| "domain": config["domain"], "physics": config["physics"], "seed": seed}, checkpoint_path) |
| metrics_path.write_text(json.dumps({"history": history, "distributed_world_size": int(os.environ.get("WORLD_SIZE", "1"))}, indent=2) + "\n") |
| print(f"checkpoint={checkpoint_path.relative_to(ROOT)} parameters={module.parameter_count(bare)} final_loss={history[-1]['loss']:.6g}") |
| if distributed: |
| torch.distributed.destroy_process_group() |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|