OneScience's picture
Upload folder using huggingface_hub
e9b4f6f verified
Raw
History Blame Contribute Delete
8.51 kB
from __future__ import annotations
import logging
import os
import sys
import time
import importlib.util
from pathlib import Path
import torch
import torch.distributed as dist
import torch.nn as nn
from torch.nn.parallel import DistributedDataParallel
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
os.chdir(ROOT)
from model import Transolver3D, Transolver3D_plus
import onescience
from onescience.distributed.manager import DistributedManager
from onescience.utils.YParams import YParams
def load_shapenet_car_datapipe():
module_path = Path(onescience.__file__).resolve().parent / "datapipes/cfd/ShapeNetCar.py"
spec = importlib.util.spec_from_file_location("_onescience_shapenetcar", module_path)
if spec is None or spec.loader is None:
raise ImportError(f"Unable to load ShapeNetCarDatapipe from {module_path}")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module.ShapeNetCarDatapipe
def setup_logging(rank: int) -> logging.Logger:
level = logging.INFO if rank == 0 else logging.WARNING
logging.basicConfig(
level=level,
format="%(asctime)s - %(levelname)s - %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
logging.getLogger().setLevel(level)
return logging.getLogger(__name__)
def build_model(model_name: str, model_params, device: torch.device) -> torch.nn.Module:
model_cls = {
"Transolver": Transolver3D,
"Transolver_plus": Transolver3D_plus,
}.get(model_name)
if model_cls is None:
raise NotImplementedError(f"Model {model_name} initialization not implemented.")
return model_cls(
n_hidden=model_params.n_hidden,
n_layers=model_params.n_layers,
space_dim=model_params.space_dim,
fun_dim=model_params.fun_dim,
n_head=model_params.n_head,
mlp_ratio=model_params.mlp_ratio,
out_dim=model_params.out_dim,
slice_num=model_params.slice_num,
unified_pos=model_params.unified_pos,
).to(device)
def resolve_device(gpuid: int) -> torch.device:
if torch.cuda.is_available() and int(gpuid) >= 0:
return torch.device(f"cuda:{gpuid}")
return torch.device("cpu")
def save_checkpoint(model, optimizer, scheduler, epoch: int, loss: float, ckp_dir: str, model_name: str) -> None:
Path(ckp_dir).mkdir(parents=True, exist_ok=True)
model_to_save = model.module if hasattr(model, "module") else model
torch.save(
{
"model_state_dict": model_to_save.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"scheduler_state_dict": scheduler.state_dict(),
"epoch": epoch,
"loss": loss,
},
Path(ckp_dir) / f"{model_name}.pth",
)
def main() -> None:
DistributedManager.initialize()
manager = DistributedManager()
logger = setup_logging(manager.rank)
config_file_path = str(ROOT / "conf/config.yaml")
cfg = YParams(config_file_path, "model")
cfg_data = YParams(config_file_path, "datapipe")
cfg_train = YParams(config_file_path, "training")
model_name = cfg.name
if model_name not in cfg.specific_params:
raise ValueError(f"Model '{model_name}' not found in model.specific_params.")
model_params = cfg.specific_params[model_name]
cfg_data.model_hparams = model_params
logger.info("Initializing ShapeNetCar datapipe...")
ShapeNetCarDatapipe = load_shapenet_car_datapipe()
datapipe = ShapeNetCarDatapipe(params=cfg_data, distributed=(manager.world_size > 1))
train_dataloader, train_sampler = datapipe.train_dataloader()
val_dataloader, val_sampler = datapipe.val_dataloader()
if manager.world_size > 1:
device = torch.device(f"cuda:{manager.local_rank}" if torch.cuda.is_available() else "cpu")
else:
device = resolve_device(cfg_train.gpuid)
model = build_model(model_name, model_params, device)
if manager.rank == 0:
total_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
logger.info("Model: %s, trainable params: %.2fM", model_name, total_params / 1e6)
if manager.world_size > 1:
model = DistributedDataParallel(
model,
device_ids=[manager.local_rank],
output_device=manager.local_rank,
find_unused_parameters=True,
)
optimizer = torch.optim.Adam(model.parameters(), lr=cfg_train.lr)
scheduler = torch.optim.lr_scheduler.OneCycleLR(
optimizer,
max_lr=cfg_train.lr,
total_steps=max(1, len(train_dataloader) * cfg_train.max_epoch),
)
if cfg_train.loss_criterion == "MSE":
loss_criterion = nn.MSELoss(reduction="none")
elif cfg_train.loss_criterion == "MAE":
loss_criterion = nn.L1Loss(reduction="none")
else:
raise ValueError(f"Unknown loss criterion: {cfg_train.loss_criterion}")
best_valid_loss = 1.0e6
best_loss_epoch = 0
logger.info("Starting training...")
for epoch in range(cfg_train.max_epoch):
epoch_start_time = time.time()
if manager.world_size > 1:
train_sampler.set_epoch(epoch)
if val_sampler is not None:
val_sampler.set_epoch(epoch)
model.train()
train_loss = train_loss_press = train_loss_velo = 0.0
for data in train_dataloader:
data = data.to(device)
optimizer.zero_grad()
out = model(data)
targets = data.y
loss_press = loss_criterion(out[data.surf, -1], targets[data.surf, -1]).mean()
loss_velo = loss_criterion(out[:, :-1], targets[:, :-1]).mean()
loss = loss_velo + cfg_train.loss_weight * loss_press
loss.backward()
optimizer.step()
scheduler.step()
train_loss += loss.item()
train_loss_press += loss_press.item()
train_loss_velo += loss_velo.item()
train_loss /= max(1, len(train_dataloader))
train_loss_press /= max(1, len(train_dataloader))
train_loss_velo /= max(1, len(train_dataloader))
valid_loss = valid_loss_press = valid_loss_velo = 0.0
if (epoch + 1) % cfg_train.val_iter == 0 or epoch == cfg_train.max_epoch - 1:
model.eval()
with torch.no_grad():
for data in val_dataloader:
data = data.to(device)
out = model(data)
targets = data.y
loss_press = loss_criterion(out[data.surf, -1], targets[data.surf, -1]).mean()
loss_velo = loss_criterion(out[:, :-1], targets[:, :-1]).mean()
loss = loss_velo + cfg_train.loss_weight * loss_press
if manager.world_size > 1:
dist.all_reduce(loss, op=dist.ReduceOp.AVG)
dist.all_reduce(loss_press, op=dist.ReduceOp.AVG)
dist.all_reduce(loss_velo, op=dist.ReduceOp.AVG)
valid_loss += loss.item()
valid_loss_press += loss_press.item()
valid_loss_velo += loss_velo.item()
valid_loss /= max(1, len(val_dataloader))
valid_loss_press /= max(1, len(val_dataloader))
valid_loss_velo /= max(1, len(val_dataloader))
if manager.rank == 0:
logger.info(
"Epoch [%d/%d] | Time: %.2fs | Train: %.6f (velo %.6f, press %.6f) | "
"Valid: %.6f (velo %.6f, press %.6f)",
epoch + 1,
cfg_train.max_epoch,
time.time() - epoch_start_time,
train_loss,
train_loss_velo,
train_loss_press,
valid_loss,
valid_loss_velo,
valid_loss_press,
)
if valid_loss > 0 and valid_loss < best_valid_loss:
best_valid_loss = valid_loss
best_loss_epoch = epoch
save_checkpoint(model, optimizer, scheduler, epoch, valid_loss, cfg_train.checkpoint_dir, model_name)
logger.info("New best checkpoint saved to %s/%s.pth", cfg_train.checkpoint_dir, model_name)
if epoch - best_loss_epoch > cfg_train.patience:
logger.warning("Validation loss has not improved for %d epochs. Stopping.", cfg_train.patience)
break
if __name__ == "__main__":
main()