MeshGraphNet / scripts /fake_data.py
OneScience's picture
Upload folder using huggingface_hub
d46a58d verified
Raw
History Blame Contribute Delete
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()