File size: 5,228 Bytes
d46a58d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
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()