Team Ai
Modelpublic

OneScience-Group/MeshGraphNet

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
0likes22downloads
train.py155 linesDownload Raw Back to scripts
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