OneScience-Group/MeshGraphNet
022
1import logging2import os3import sys4import time5from pathlib import Path6 7import torch8import torch.nn as nn9from torch.amp import GradScaler, autocast10from torch.nn.parallel import DistributedDataParallel11 12PROJECT_ROOT = Path(__file__).resolve().parents[1]13sys.path.insert(0, str(PROJECT_ROOT))14 15from model.meshgraphnet import MeshGraphNet16from onescience.distributed import DistributedManager17from onescience.utils.YParams import YParams18from onescience.launch.utils import load_checkpoint, save_checkpoint19from fake_data import build_cylinder_flow_datapipe20 21 22def setup_logging(rank: int):23 level = logging.INFO if rank == 0 else logging.WARNING24 logging.basicConfig(level=level, format="%(asctime)s - %(levelname)s - %(message)s")25 return logging.getLogger("mesh_graph_net.train")26 27 28def build_model(model_params, device):29 mlp_act = "silu" if model_params.recompute_activation else "relu"30 return MeshGraphNet(31 input_dim_nodes=model_params.num_input_features,32 input_dim_edges=model_params.num_edge_features,33 output_dim=model_params.num_output_features,34 processor_size=model_params.processor_size,35 hidden_dim_processor=model_params.hidden_dim_processor,36 num_layers_node_processor=model_params.num_layers_node_processor,37 num_layers_edge_processor=model_params.num_layers_edge_processor,38 hidden_dim_node_encoder=model_params.hidden_dim_node_encoder,39 hidden_dim_edge_encoder=model_params.hidden_dim_edge_encoder,40 hidden_dim_node_decoder=model_params.hidden_dim_node_decoder,41 mlp_activation_fn=mlp_act,42 do_concat_trick=model_params.do_concat_trick,43 num_processor_checkpoint_segments=model_params.num_processor_checkpoint_segments,44 recompute_activation=model_params.recompute_activation,45 ).to(device)46 47 48def graph_from_batch(batch):49 return batch[0] if isinstance(batch, (tuple, list)) else batch50 51 52def resolve_device(device_name: str, manager: DistributedManager, gpuid: int):53 if manager.world_size > 1:54 return manager.device55 if device_name == "cpu":56 return torch.device("cpu")57 if device_name in ("cuda", "gpu"):58 if not torch.cuda.is_available():59 raise RuntimeError("Config requested cuda device, but torch.cuda.is_available() is false.")60 return torch.device(f"cuda:{gpuid}")61 return torch.device(f"cuda:{gpuid}" if torch.cuda.is_available() else "cpu")62 63 64def main():65 os.chdir(PROJECT_ROOT)66 DistributedManager.initialize()67 manager = DistributedManager()68 logger = setup_logging(manager.rank)69 70 config_path = PROJECT_ROOT / "config" / "config.yaml"71 cfg_model = YParams(config_path, "model")72 cfg_data = YParams(config_path, "datapipe")73 cfg_train = YParams(config_path, "training")74 model_params = cfg_model.specific_params[cfg_model.name]75 76 datapipe = build_cylinder_flow_datapipe(77 params=cfg_data,78 distributed=(manager.world_size > 1),79 project_root=PROJECT_ROOT,80 )81 train_loader, train_sampler = datapipe.train_dataloader()82 val_loader, val_sampler = datapipe.val_dataloader()83 84 device = resolve_device(getattr(cfg_train, "device", "auto"), manager, cfg_train.gpuid)85 logger.info("Using device: %s", device)86 model = build_model(model_params, device)87 if manager.world_size > 1:88 model = DistributedDataParallel(model, device_ids=[manager.local_rank], output_device=manager.local_rank)89 90 optimizer = torch.optim.Adam(model.parameters(), lr=cfg_train.lr)91 scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=lambda step: cfg_train.lr_decay_rate**step)92 loss_criterion = nn.MSELoss() if cfg_train.loss_criterion == "MSE" else nn.L1Loss()93 scaler = GradScaler(enabled=bool(cfg_train.amp))94 95 checkpoint_dir = PROJECT_ROOT / cfg_train.checkpoint_dir96 epoch_init = load_checkpoint(checkpoint_dir, models=model, optimizer=optimizer, scheduler=scheduler, scaler=scaler, device=device)97 best_valid_loss = float("inf")98 best_loss_epoch = epoch_init99 100 logger.info("Starting training")101 for epoch in range(epoch_init, cfg_train.max_epoch):102 if train_sampler is not None:103 train_sampler.set_epoch(epoch)104 start = time.time()105 model.train()106 train_loss = 0.0107 for idx, batch in enumerate(train_loader):108 graph = graph_from_batch(batch).to(device)109 optimizer.zero_grad(set_to_none=True)110 with autocast(device_type=device.type, enabled=bool(cfg_train.amp)):111 pred = model(graph.ndata["x"], graph.edata["x"], graph)112 loss = loss_criterion(pred, graph.ndata["y"])113 scaler.scale(loss).backward()114 scaler.step(optimizer)115 scaler.update()116 scheduler.step()117 train_loss += loss.item()118 if manager.rank == 0 and (idx + 1) % cfg_train.log_interval == 0:119 logger.info("Epoch %s/%s batch %s/%s loss %.6f", epoch + 1, cfg_train.max_epoch, idx + 1, len(train_loader), loss.item())120 121 train_loss /= max(len(train_loader), 1)122 model.eval()123 valid_loss = 0.0124 with torch.no_grad():125 for batch in val_loader:126 graph = graph_from_batch(batch).to(device)127 with autocast(device_type=device.type, enabled=bool(cfg_train.amp)):128 pred = model(graph.ndata["x"], graph.edata["x"], graph)129 loss = loss_criterion(pred, graph.ndata["y"])130 valid_loss += loss.item()131 valid_loss /= max(len(val_loader), 1)132 133 if manager.rank == 0:134 logger.info(135 "Epoch %s finished in %.2fs train_loss %.6f valid_loss %.6f",136 epoch + 1,137 time.time() - start,138 train_loss,139 valid_loss,140 )141 if valid_loss < best_valid_loss:142 best_valid_loss = valid_loss143 best_loss_epoch = epoch144 save_checkpoint(checkpoint_dir, models=model, optimizer=optimizer, scheduler=scheduler, scaler=scaler, epoch=epoch + 1)145 logger.info("Checkpoint saved to %s", checkpoint_dir)146 if (epoch - best_loss_epoch) > cfg_train.patience:147 logger.warning("Early stopping after %s stale epochs", cfg_train.patience)148 break149 150 logger.info("Training finished")151 152 153if __name__ == "__main__":154 main()155 