MCPcopy Create free account
hub / github.com/apple/ml-pointersect / MeshConcatDataset

Class MeshConcatDataset

pointersect/data/mesh_dataset_v2.py:376–419  ·  view source on GitHub ↗

r"""Dataset as a concatenation of multiple MeshDatasets. Arguments: datasets (sequence): List of datasets to be concatenated

Source from the content-addressed store, hash-verified

374
375
376class MeshConcatDataset(torch.utils.data.Dataset):
377 r"""Dataset as a concatenation of multiple MeshDatasets.
378
379 Arguments:
380 datasets (sequence): List of datasets to be concatenated
381 """
382
383 def __init__(
384 self,
385 datasets: T.Iterable[MeshDataset],
386 max_retry=30,
387 wait_sec=5,
388 ):
389
390 self.datasets = list(datasets)
391 assert len(self.datasets) > 0
392 self.concat_dataset = torch.utils.data.ConcatDataset(self.datasets)
393 self.max_retry = max_retry
394 self.wait_sec = wait_sec
395
396 def __len__(self):
397 return len(self.concat_dataset)
398
399 def __getitem__(self, idx):
400 d = None
401 for retry in range(self.max_retry):
402 try:
403 d = self.concat_dataset[idx]
404 if d is not None:
405 return d
406 except:
407 traceback.print_exc()
408 time.sleep(self.wait_sec)
409 return d
410
411 def get_all_num_pixels(self):
412 all_seq_lens = []
413 for dset in self.datasets:
414 seq_lens = dset.get_all_num_pixels()
415 if seq_lens is None:
416 return None
417 else:
418 all_seq_lens += seq_lens
419 return all_seq_lens

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected