Download scripts/fake_data.py from OneScience-Group/MeshGraphNet: direct link, hf CLI and curl.
- Browser
- Download file 5.23 kB
-
https://huggingface.co/OneScience-Group/MeshGraphNet/resolve/5e387cabb1301867a633f76fa9ca4c39fdc9bbbd/scripts/fake_data.py
- Command line
-
hf download hf://OneScience-Group/MeshGraphNet@5e387cabb1301867a633f76fa9ca4c39fdc9bbbd/scripts/fake_data.py
-
curl -L -o fake_data.py https://huggingface.co/OneScience-Group/MeshGraphNet/resolve/5e387cabb1301867a633f76fa9ca4c39fdc9bbbd/scripts/fake_data.py
5.23 kB
| import os | |
| import sys | |
| from pathlib import Path | |
| import dgl | |
| import torch | |
| from dgl.dataloading import GraphDataLoader | |
| from torch.utils.data import Dataset | |
| PROJECT_ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(PROJECT_ROOT / "model")) | |
| from onescience.utils.YParams import YParams | |
| def make_graph(num_nodes: int = 12): | |
| src = torch.arange(num_nodes, dtype=torch.int32) | |
| dst = torch.roll(src, shifts=-1) | |
| graph = dgl.to_bidirected(dgl.graph((src, dst), num_nodes=num_nodes, idtype=torch.int32)) | |
| pos = torch.stack( | |
| ( | |
| torch.linspace(0.0, 1.0, num_nodes), | |
| torch.sin(torch.linspace(0.0, 3.14159, num_nodes)) * 0.2, | |
| ), | |
| dim=1, | |
| ) | |
| row, col = graph.edges() | |
| disp = pos[row.long()] - pos[col.long()] | |
| graph.edata["x"] = torch.cat( | |
| (disp, torch.linalg.norm(disp, dim=-1, keepdim=True)), | |
| dim=1, | |
| ) | |
| velocity = torch.randn(num_nodes, 2) * 0.1 | |
| node_type = torch.zeros(num_nodes, 4) | |
| node_type[:, 0] = 1.0 | |
| graph.ndata["x"] = torch.cat((velocity, node_type), dim=1) | |
| graph.ndata["y"] = torch.cat( | |
| (torch.randn(num_nodes, 2) * 0.01, torch.randn(num_nodes, 1) * 0.01), | |
| dim=1, | |
| ) | |
| graph.ndata["mesh_pos"] = pos | |
| cells = torch.tensor( | |
| [[i, i + 1, min(i + 2, num_nodes - 1)] for i in range(num_nodes - 2)], | |
| dtype=torch.int64, | |
| ) | |
| mask = torch.ones(num_nodes, 1, dtype=torch.bool) | |
| return {"graph": graph, "cells": cells, "mask": mask} | |
| class FakeGraphDataset(Dataset): | |
| def __init__(self, samples): | |
| self.samples = samples | |
| def __len__(self): | |
| return len(self.samples) | |
| def __getitem__(self, index): | |
| sample = self.samples[index] | |
| if isinstance(sample, dict) and "graph" in sample: | |
| return sample["graph"] | |
| return sample | |
| def _resolve_path(project_root: Path, path): | |
| path = Path(path) | |
| return path if path.is_absolute() else project_root / path | |
| def _torch_load(path: Path): | |
| try: | |
| return torch.load(path, map_location="cpu", weights_only=False) | |
| except TypeError: | |
| return torch.load(path, map_location="cpu") | |
| class FakeCylinderFlowDatapipe: | |
| def __init__(self, params, project_root: Path): | |
| self.params = params | |
| fake_data_path = _resolve_path(project_root, params.source.fake_data_path) | |
| if not fake_data_path.exists(): | |
| raise FileNotFoundError( | |
| f"Fake data file not found: {fake_data_path}. Run scripts/fake_data.py first." | |
| ) | |
| payload = _torch_load(fake_data_path) | |
| self.train_dataset = FakeGraphDataset(payload["train"]) | |
| self.val_dataset = FakeGraphDataset(payload["val"]) | |
| self.test_dataset = FakeGraphDataset(payload["test"]) | |
| self.stats = payload.get("stats", {}) | |
| def _loader(self, dataset, shuffle=False, drop_last=False): | |
| return GraphDataLoader( | |
| dataset, | |
| batch_size=self.params.dataloader.batch_size, | |
| drop_last=drop_last, | |
| num_workers=self.params.dataloader.num_workers, | |
| pin_memory=True, | |
| shuffle=shuffle, | |
| ) | |
| def train_dataloader(self): | |
| return self._loader(self.train_dataset, shuffle=True), None | |
| def val_dataloader(self): | |
| return self._loader(self.val_dataset), None | |
| def test_dataloader(self): | |
| return self._loader(self.test_dataset) | |
| def use_fake_data(params): | |
| return bool(getattr(params.source, "fake_data", False)) | |
| def build_cylinder_flow_datapipe(params, distributed: bool, project_root: Path): | |
| if use_fake_data(params): | |
| return FakeCylinderFlowDatapipe(params=params, project_root=project_root) | |
| from onescience.datapipes.cfd import DeepMind_CylinderFlowDatapipe | |
| return DeepMind_CylinderFlowDatapipe(params=params, distributed=distributed) | |
| def main(): | |
| os.chdir(PROJECT_ROOT) | |
| config_path = PROJECT_ROOT / "config" / "config.yaml" | |
| cfg_data = YParams(config_path, "datapipe") | |
| output_path = PROJECT_ROOT / cfg_data.source.fake_data_path | |
| output_path.parent.mkdir(parents=True, exist_ok=True) | |
| payload = { | |
| "train": [ | |
| make_graph() | |
| for _ in range(cfg_data.data.train_samples * (cfg_data.data.train_steps - 1)) | |
| ], | |
| "val": [ | |
| make_graph() | |
| for _ in range(cfg_data.data.val_samples * (cfg_data.data.val_steps - 1)) | |
| ], | |
| "test": [ | |
| make_graph() | |
| for _ in range(cfg_data.data.test_samples * (cfg_data.data.test_steps - 1)) | |
| ], | |
| "stats": { | |
| "edge_stats": { | |
| "edge_mean": torch.zeros(3), | |
| "edge_std": torch.ones(3), | |
| }, | |
| "node_stats": { | |
| "velocity_mean": torch.zeros(2), | |
| "velocity_std": torch.ones(2), | |
| "velocity_diff_mean": torch.zeros(2), | |
| "velocity_diff_std": torch.ones(2), | |
| "pressure_mean": torch.zeros(1), | |
| "pressure_std": torch.ones(1), | |
| }, | |
| }, | |
| } | |
| torch.save(payload, output_path) | |
| print(f"Fake data saved to {output_path.relative_to(PROJECT_ROOT)}") | |
| if __name__ == "__main__": | |
| main() | |