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

Class TextDataset

code/dataset_benchmark.py:260–306  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

258
259
260class TextDataset(InMemoryDataset):
261 '''
262 '''
263 url = 'https://github.com/pygod-team/data/raw/main/'
264 file_urls = {
265 "reddit": "reddit.pt.zip",
266 "weibo": "weibo.pt.zip",
267 }
268 file_names = {"reddit": "reddit.pt", "weibo": "weibo.pt"}
269
270 def __init__(self, root, name, transform=None, pre_transform=None):
271
272 self.name = name
273 assert self.name in ['reddit', 'weibo']
274
275 self.url = osp.join(self.url, self.file_urls[self.name])
276
277 super().__init__(root, transform, pre_transform)
278 self.load(self.processed_paths[0])
279
280 @property
281 def raw_dir(self):
282 return osp.join(self.root, self.name, 'raw')
283
284 @property
285 def processed_dir(self):
286 return osp.join(self.root, self.name, 'processed')
287
288 @property
289 def raw_file_names(self):
290 names = [self.file_names[self.name]]
291 return names
292
293 @property
294 def processed_file_names(self):
295 return 'data.pt'
296
297 def download(self):
298 path = download_url(self.url, self.raw_dir)
299 extract_zip(path, self.raw_dir)
300 os.unlink(path)
301
302 def process(self):
303 file_path = os.path.join(self.raw_dir, self.raw_file_names[0])
304 data = torch.load(file_path)
305 data = data if self.pre_transform is None else self.pre_transform(data)
306 self.save([data], self.processed_paths[0])
307
308
309# ======================================================================

Callers 1

load_node_datasetFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected