File size: 7,430 Bytes
2f3c9e4 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 | """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()
|