PINN-TC / scripts /train.py
zhangrenchao's picture
Upload folder using huggingface_hub
2f3c9e4 verified
Raw
History Blame Contribute Delete
7.43 kB
"""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()