MCPcopy Create free account
hub / github.com/InternLM/InternBootcamp / make_iterator

Method make_iterator

verl/verl/protocol.py:811–849  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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 """

Callers 5

train_mini_batchMethod · 0.80

Calls 1

get_dataFunction · 0.85

Tested by 2