MCPcopy Create free account
hub / github.com/Tencent/digitalhuman / make_iterator

Method make_iterator

RLVER/code/verl/protocol.py:453–490  ·  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/tutorials/data_fashion for more details. Args: mini_batch_size (int): mini-batch size when iterating the dataset. We r

(self, mini_batch_size, epochs, seed=None, dataloader_kwargs=None)

Source from the content-addressed store, hash-verified

451 return self
452
453 def make_iterator(self, mini_batch_size, epochs, seed=None, dataloader_kwargs=None):
454 r"""Make an iterator from the DataProto. This is built upon that TensorDict can be used as a normal Pytorch
455 dataset. See https://pytorch.org/tensordict/tutorials/data_fashion for more details.
456
457
458 Args:
459 mini_batch_size (int): mini-batch size when iterating the dataset. We require that ``batch.batch_size[0] % mini_batch_size == 0``.
460 epochs (int): number of epochs when iterating the dataset.
461 dataloader_kwargs (Any): internally, it returns a DataLoader over the batch. The dataloader_kwargs is the kwargs passed to the DataLoader.
462
463 Returns:
464 Iterator: an iterator that yields a mini-batch data at a time. The total number of iteration steps is ``self.batch.batch_size * epochs // mini_batch_size``
465 """
466 assert self.batch.batch_size[0] % mini_batch_size == 0, f"{self.batch.batch_size[0]} % {mini_batch_size} != 0"
467 # we can directly create a dataloader from TensorDict
468 if dataloader_kwargs is None:
469 dataloader_kwargs = {}
470
471 if seed is not None:
472 generator = torch.Generator()
473 generator.manual_seed(seed)
474 else:
475 generator = None
476
477 assert isinstance(dataloader_kwargs, Dict)
478 train_dataloader = DataLoader(dataset=self,
479 batch_size=mini_batch_size,
480 collate_fn=collate_fn,
481 generator=generator,
482 **dataloader_kwargs)
483
484 def get_data():
485 for _ in range(epochs):
486 for d in train_dataloader:
487 d.meta_info = self.meta_info
488 yield d
489
490 return iter(get_data())
491
492 def chunk(self, chunks: int) -> List['DataProto']:
493 """Split the batch among dim=0 into chunks. The meta_info is passed to each DataProto after split.

Callers 3

Calls 1

get_dataFunction · 0.50

Tested by 1