Team Ai
Modelpublic

OneScience-Group/MeshGraphNet

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
0likes22downloads
fake_data.py168 linesDownload Raw Back to scripts
1import os2import sys3from pathlib import Path4 5import dgl6import torch7from dgl.dataloading import GraphDataLoader8from torch.utils.data import Dataset9 10PROJECT_ROOT = Path(__file__).resolve().parents[1]11sys.path.insert(0, str(PROJECT_ROOT / "model"))12 13from onescience.utils.YParams import YParams14 15 16def make_graph(num_nodes: int = 12):17    src = torch.arange(num_nodes, dtype=torch.int32)18    dst = torch.roll(src, shifts=-1)19    graph = dgl.to_bidirected(dgl.graph((src, dst), num_nodes=num_nodes, idtype=torch.int32))20 21    pos = torch.stack(22        (23            torch.linspace(0.0, 1.0, num_nodes),24            torch.sin(torch.linspace(0.0, 3.14159, num_nodes)) * 0.2,25        ),26        dim=1,27    )28    row, col = graph.edges()29    disp = pos[row.long()] - pos[col.long()]30    graph.edata["x"] = torch.cat(31        (disp, torch.linalg.norm(disp, dim=-1, keepdim=True)),32        dim=1,33    )34 35    velocity = torch.randn(num_nodes, 2) * 0.136    node_type = torch.zeros(num_nodes, 4)37    node_type[:, 0] = 1.038    graph.ndata["x"] = torch.cat((velocity, node_type), dim=1)39    graph.ndata["y"] = torch.cat(40        (torch.randn(num_nodes, 2) * 0.01, torch.randn(num_nodes, 1) * 0.01),41        dim=1,42    )43    graph.ndata["mesh_pos"] = pos44 45    cells = torch.tensor(46        [[i, i + 1, min(i + 2, num_nodes - 1)] for i in range(num_nodes - 2)],47        dtype=torch.int64,48    )49    mask = torch.ones(num_nodes, 1, dtype=torch.bool)50    return {"graph": graph, "cells": cells, "mask": mask}51 52 53class FakeGraphDataset(Dataset):54    def __init__(self, samples):55        self.samples = samples56 57    def __len__(self):58        return len(self.samples)59 60    def __getitem__(self, index):61        sample = self.samples[index]62        if isinstance(sample, dict) and "graph" in sample:63            return sample["graph"]64        return sample65 66 67def _resolve_path(project_root: Path, path):68    path = Path(path)69    return path if path.is_absolute() else project_root / path70 71 72def _torch_load(path: Path):73    try:74        return torch.load(path, map_location="cpu", weights_only=False)75    except TypeError:76        return torch.load(path, map_location="cpu")77 78 79class FakeCylinderFlowDatapipe:80    def __init__(self, params, project_root: Path):81        self.params = params82        fake_data_path = _resolve_path(project_root, params.source.fake_data_path)83        if not fake_data_path.exists():84            raise FileNotFoundError(85                f"Fake data file not found: {fake_data_path}. Run scripts/fake_data.py first."86            )87 88        payload = _torch_load(fake_data_path)89        self.train_dataset = FakeGraphDataset(payload["train"])90        self.val_dataset = FakeGraphDataset(payload["val"])91        self.test_dataset = FakeGraphDataset(payload["test"])92        self.stats = payload.get("stats", {})93 94    def _loader(self, dataset, shuffle=False, drop_last=False):95        return GraphDataLoader(96            dataset,97            batch_size=self.params.dataloader.batch_size,98            drop_last=drop_last,99            num_workers=self.params.dataloader.num_workers,100            pin_memory=True,101            shuffle=shuffle,102        )103 104    def train_dataloader(self):105        return self._loader(self.train_dataset, shuffle=True), None106 107    def val_dataloader(self):108        return self._loader(self.val_dataset), None109 110    def test_dataloader(self):111        return self._loader(self.test_dataset)112 113 114def use_fake_data(params):115    return bool(getattr(params.source, "fake_data", False))116 117 118def build_cylinder_flow_datapipe(params, distributed: bool, project_root: Path):119    if use_fake_data(params):120        return FakeCylinderFlowDatapipe(params=params, project_root=project_root)121 122    from onescience.datapipes.cfd import DeepMind_CylinderFlowDatapipe123 124    return DeepMind_CylinderFlowDatapipe(params=params, distributed=distributed)125 126 127def main():128    os.chdir(PROJECT_ROOT)129    config_path = PROJECT_ROOT / "config" / "config.yaml"130    cfg_data = YParams(config_path, "datapipe")131    output_path = PROJECT_ROOT / cfg_data.source.fake_data_path132    output_path.parent.mkdir(parents=True, exist_ok=True)133 134    payload = {135        "train": [136            make_graph()137            for _ in range(cfg_data.data.train_samples * (cfg_data.data.train_steps - 1))138        ],139        "val": [140            make_graph()141            for _ in range(cfg_data.data.val_samples * (cfg_data.data.val_steps - 1))142        ],143        "test": [144            make_graph()145            for _ in range(cfg_data.data.test_samples * (cfg_data.data.test_steps - 1))146        ],147        "stats": {148            "edge_stats": {149                "edge_mean": torch.zeros(3),150                "edge_std": torch.ones(3),151            },152            "node_stats": {153                "velocity_mean": torch.zeros(2),154                "velocity_std": torch.ones(2),155                "velocity_diff_mean": torch.zeros(2),156                "velocity_diff_std": torch.ones(2),157                "pressure_mean": torch.zeros(1),158                "pressure_std": torch.ones(1),159            },160        },161    }162    torch.save(payload, output_path)163    print(f"Fake data saved to {output_path.relative_to(PROJECT_ROOT)}")164 165 166if __name__ == "__main__":167    main()168