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