Wraps an arbitrary object with __len__ and __getitem__ into a pytorch dataset
| 292 | return parser |
| 293 | |
| 294 | class WrappedDataset(Dataset): |
| 295 | """Wraps an arbitrary object with __len__ and __getitem__ into a pytorch dataset""" |
| 296 | |
| 297 | def __init__(self, dataset): |
| 298 | self.data = dataset |
| 299 | |
| 300 | def __len__(self): |
| 301 | return len(self.data) |
| 302 | |
| 303 | def __getitem__(self, idx): |
| 304 | return self.data[idx] |
| 305 | class ConcatDataset(Dataset): |
| 306 | def __init__(self, *datasets): |
| 307 | self.datasets = datasets |