| 9 | |
| 10 | |
| 11 | class TorchDataset(Dataset): |
| 12 | |
| 13 | def __init__(self, ds: deeplake.Dataset, transform=None): |
| 14 | self.ds = ds |
| 15 | self.transform = transform |
| 16 | self.column_names = [col.name for col in ds.schema.columns] |
| 17 | |
| 18 | def __len__(self): |
| 19 | return len(self.ds) |
| 20 | |
| 21 | def __getitem__(self, idx): |
| 22 | sample = self.ds[idx] |
| 23 | if self.transform: |
| 24 | return self.transform(sample) |
| 25 | else: |
| 26 | out = {} |
| 27 | for col in self.column_names: |
| 28 | out[col] = sample[col] |
| 29 | return out |