r"""Make an iterator from the DataProto. This is built upon that TensorDict can be used as a normal Pytorch dataset. See https://pytorch.org/tensordict/stable/tutorials/data_fashion for more details. Args: mini_batch_size (int): mini-batch size when iterating the datase
(self, mini_batch_size, epochs, seed=None, dataloader_kwargs=None)
| 809 | return self |
| 810 | |
| 811 | def make_iterator(self, mini_batch_size, epochs, seed=None, dataloader_kwargs=None): |
| 812 | r"""Make an iterator from the DataProto. This is built upon that TensorDict can be used as a normal Pytorch |
| 813 | dataset. See https://pytorch.org/tensordict/stable/tutorials/data_fashion for more details. |
| 814 | |
| 815 | |
| 816 | Args: |
| 817 | mini_batch_size (int): mini-batch size when iterating the dataset. We require that |
| 818 | ``batch.batch_size[0] % mini_batch_size == 0``. |
| 819 | epochs (int): number of epochs when iterating the dataset. |
| 820 | dataloader_kwargs (Any): internally, it returns a DataLoader over the batch. The |
| 821 | dataloader_kwargs is the kwargs passed to the DataLoader. |
| 822 | |
| 823 | Returns: |
| 824 | Iterator: an iterator that yields a mini-batch data at a time. The total number of iteration |
| 825 | steps is ``self.batch.batch_size * epochs // mini_batch_size`` |
| 826 | """ |
| 827 | assert self.batch.batch_size[0] % mini_batch_size == 0, f"{self.batch.batch_size[0]} % {mini_batch_size} != 0" |
| 828 | # we can directly create a dataloader from TensorDict |
| 829 | if dataloader_kwargs is None: |
| 830 | dataloader_kwargs = {} |
| 831 | |
| 832 | if seed is not None: |
| 833 | generator = torch.Generator() |
| 834 | generator.manual_seed(seed) |
| 835 | else: |
| 836 | generator = None |
| 837 | |
| 838 | assert isinstance(dataloader_kwargs, dict) |
| 839 | train_dataloader = DataLoader( |
| 840 | dataset=self, batch_size=mini_batch_size, collate_fn=collate_fn, generator=generator, **dataloader_kwargs |
| 841 | ) |
| 842 | |
| 843 | def get_data(): |
| 844 | for _ in range(epochs): |
| 845 | for d in train_dataloader: |
| 846 | d.meta_info = self.meta_info |
| 847 | yield d |
| 848 | |
| 849 | return iter(get_data()) |
| 850 | |
| 851 | def is_padding_enabled(self): |
| 852 | """ |