MCPcopy Create free account
hub / github.com/NineAbyss/ZeroG / TDataset

Class TDataset

code/dataset_benchmark.py:204–257  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

202
203
204class TDataset(InMemoryDataset):
205 '''
206
207 '''
208
209 def __init__(self, root, name, transform=None, pre_transform=None):
210 self.name = name
211 assert self.name in ['tfinance', 'tsocial']
212
213 super().__init__(root, transform, pre_transform)
214 self.load(self.processed_paths[0])
215
216 @property
217 def raw_dir(self):
218 return osp.join(self.root, self.name, 'raw')
219
220 @property
221 def processed_dir(self):
222 return osp.join(self.root, self.name, 'processed')
223
224 @property
225 def raw_file_names(self):
226 names = [self.name]
227 return names
228
229 @property
230 def processed_file_names(self):
231 return 'data.pt'
232
233 def download(self):
234 pass
235
236 def process(self):
237 file_path = os.path.join(self.raw_dir, self.raw_file_names[0])
238 if not osp.exists(file_path):
239 try:
240 shutil.copy(f'data/{self.name}', file_path)
241 except:
242 raise ValueError('source file does not exist!')
243
244 data = load_graphs(file_path)[0][0]
245 features = data.ndata['feature']
246 labels = data.ndata['label']
247 train_mask = data.ndata['train_masks']
248 val_mask = data.ndata['val_masks']
249 test_mask = data.ndata['test_masks']
250
251 data = Data(x=features, edge_index=torch.vstack(
252 data.edges()), y=labels)
253 data.tran_mask = train_mask
254 data.val_mask = val_mask
255 data.test_mask = test_mask
256 data = data if self.pre_transform is None else self.pre_transform(data)
257 self.save([data], self.processed_paths[0])
258
259
260class TextDataset(InMemoryDataset):

Callers 2

load_node_datasetFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected