项目文件夹

文件
Chen Sirui afc83aa216 Graphsim (#2794)
* Add hgat example

* Add experiment

* Clean code

* clear the code

* Add index in README

* Add index in README

* Add index in README

* Add index in README

* Add index in README

* Add index in README

* Change the code title and folder name

* Ready to merge

* Prepare for rebase and change message passing function

* use git ignore to handle empty file

* change file permission to resolve empty file

* Change permission

* change file mode

* Finish Coding

* working code cpu

* pyg compare

* Accelerate with batching

* FastMode Enabled

* update readme

* Update README.md

* refractor code

* add graphsim code

* modified code

* few fix

* Modified graphsim

* Simple Model Added

* Clean up code

* Refractor the code for Merge

* Bugfix enable gradient when train

* update readme and format

Co-authored-by: Chen <chesirui@3c22fbe5458c.ant.amazon.com>
Co-authored-by: Tianjun Xiao <xiaotj1990327@gmail.com>
Co-authored-by: Ubuntu <ubuntu@ip-172-31-4-63.ap-northeast-1.compute.internal>
Co-authored-by: Ubuntu <ubuntu@ip-172-31-45-47.ap-northeast-1.compute.internal>
2021-04-12 19:02:36 +08:00

78 行
2.3 KiB
Python

import os
import copy
import numpy as np
import torch
import dgl
import networkx as nx
from torch.utils.data import Dataset, DataLoader
def build_dense_graph(n_particles):
g = nx.complete_graph(n_particles)
return dgl.from_networkx(g)
class MultiBodyDataset(Dataset):
def __init__(self, path):
self.path = path
self.zipfile = np.load(self.path)
self.node_state = self.zipfile['data']
self.node_label = self.zipfile['label']
self.n_particles = self.zipfile['n_particles']
def __len__(self):
return self.node_state.shape[0]
def __getitem__(self, idx):
if torch.is_tensor(idx):
idx = idx.tolist()
node_state = self.node_state[idx, :, :]
node_label = self.node_label[idx, :, :]
return (node_state, node_label)
class MultiBodyTrainDataset(MultiBodyDataset):
def __init__(self, data_path='./data/'):
super(MultiBodyTrainDataset, self).__init__(
data_path+'n_body_train.npz')
self.stat_median = self.zipfile['median']
self.stat_max = self.zipfile['max']
self.stat_min = self.zipfile['min']
class MultiBodyValidDataset(MultiBodyDataset):
def __init__(self, data_path='./data/'):
super(MultiBodyValidDataset, self).__init__(
data_path+'n_body_valid.npz')
class MultiBodyTestDataset(MultiBodyDataset):
def __init__(self, data_path='./data/'):
super(MultiBodyTestDataset, self).__init__(data_path+'n_body_test.npz')
self.test_traj = self.zipfile['test_traj']
self.first_frame = torch.from_numpy(self.zipfile['first_frame'])
# Construct fully connected graph
class MultiBodyGraphCollator:
def __init__(self, n_particles):
self.n_particles = n_particles
self.graph = dgl.from_networkx(nx.complete_graph(self.n_particles))
def __call__(self, batch):
graph_list = []
data_list = []
label_list = []
for frame in batch:
graph_list.append(copy.deepcopy(self.graph))
data_list.append(torch.from_numpy(frame[0]))
label_list.append(torch.from_numpy(frame[1]))
graph_batch = dgl.batch(graph_list)
data_batch = torch.vstack(data_list)
label_batch = torch.vstack(label_list)
return graph_batch, data_batch, label_batch