OneScience-Group/MeshGraphNet
022
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 