"""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: # Every DDP rank runs the same closure count and deterministic subset so gradient collectives remain aligned. 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()