Team Ai
Modelpublic

OneScience-Group/MeshGraphNet

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
0likes22downloads
inference.py93 linesDownload Raw Back to scripts
1import logging2import os3import sys4from pathlib import Path5 6import numpy as np7import torch8 9PROJECT_ROOT = Path(__file__).resolve().parents[1]10sys.path.insert(0, str(PROJECT_ROOT))11 12 13from model.meshgraphnet import MeshGraphNet14from onescience.utils.YParams import YParams15from onescience.launch.utils import load_checkpoint 16from fake_data import build_cylinder_flow_datapipe17 18 19def build_model(model_params, device):20    mlp_act = "silu" if model_params.recompute_activation else "relu"21    return MeshGraphNet(22        input_dim_nodes=model_params.num_input_features,23        input_dim_edges=model_params.num_edge_features,24        output_dim=model_params.num_output_features,25        processor_size=model_params.processor_size,26        hidden_dim_processor=model_params.hidden_dim_processor,27        num_layers_node_processor=model_params.num_layers_node_processor,28        num_layers_edge_processor=model_params.num_layers_edge_processor,29        hidden_dim_node_encoder=model_params.hidden_dim_node_encoder,30        hidden_dim_edge_encoder=model_params.hidden_dim_edge_encoder,31        hidden_dim_node_decoder=model_params.hidden_dim_node_decoder,32        mlp_activation_fn=mlp_act,33        do_concat_trick=model_params.do_concat_trick,34        num_processor_checkpoint_segments=model_params.num_processor_checkpoint_segments,35        recompute_activation=model_params.recompute_activation,36    ).to(device)37 38 39def resolve_device(device_name: str):40    if device_name == "cpu":41        return torch.device("cpu")42    if device_name in ("cuda", "gpu"):43        if not torch.cuda.is_available():44            raise RuntimeError("Config requested cuda device, but torch.cuda.is_available() is false.")45        return torch.device("cuda:0")46    return torch.device("cuda:0" if torch.cuda.is_available() else "cpu")47 48 49def main():50    os.chdir(PROJECT_ROOT)51    logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")52    logger = logging.getLogger("mesh_graph_net.inference")53 54    config_path = PROJECT_ROOT / "config" / "config.yaml"55    cfg_model = YParams(config_path, "model")56    cfg_data = YParams(config_path, "datapipe")57    cfg_train = YParams(config_path, "training")58    cfg_inference = YParams(config_path, "inference")59    model_params = cfg_model.specific_params[cfg_model.name]60 61    device = resolve_device(getattr(cfg_inference, "device", "auto"))62    logger.info("Using device: %s", device)63    datapipe = build_cylinder_flow_datapipe(64        params=cfg_data,65        distributed=False,66        project_root=PROJECT_ROOT,67    )68    loader = datapipe.test_dataloader()69    model = build_model(model_params, device)70    checkpoint_dir = PROJECT_ROOT / getattr(cfg_inference, "checkpoint_dir", cfg_train.checkpoint_dir)71    epoch = load_checkpoint(checkpoint_dir, models=model, device=device)72    if epoch == 0:73        logger.warning("No checkpoint found in %s; running with randomly initialized weights", checkpoint_dir)74    model.eval()75 76    predictions, targets = [], []77    with torch.no_grad():78        for batch in loader:79            graph = batch[0] if isinstance(batch, (tuple, list)) else batch80            graph = graph.to(device)81            pred = model(graph.ndata["x"], graph.edata["x"], graph)82            predictions.append(pred.cpu().numpy())83            targets.append(graph.ndata["y"].cpu().numpy())84 85    output_path = PROJECT_ROOT / cfg_inference.output_path86    output_path.parent.mkdir(parents=True, exist_ok=True)87    np.savez(output_path, prediction=np.concatenate(predictions, axis=0), target=np.concatenate(targets, axis=0))88    logger.info("Saved inference results to %s", output_path)89 90 91if __name__ == "__main__":92    main()93