File size: 2,139 Bytes
12fea4a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 | # Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the CC-by-NC license found in the
# LICENSE file in the root directory of this source tree.
# This implementation is adapted from https://github.com/pytorch/data/blob/main/torchdata/stateful_dataloader/sampler.py#L132
# which is released under BSD-3 license
import itertools
from typing import Any, Dict, Optional
import numpy as np
import torch
from torch import Tensor
from torch.utils.data import DataLoader, Dataset, Sampler
def cycle_loader(dataloader: DataLoader, sampler: Sampler = None) -> Tensor:
while 1:
if sampler is not None:
sampler.set_epoch(np.random.randint(0, 100000))
for data in dataloader:
yield data
class StatefulDistributedSampler(torch.utils.data.distributed.DistributedSampler):
"""
From: https://github.com/pytorch/data/blob/main/torchdata/stateful_dataloader/sampler.py#L132
"""
_YIELDED = "yielded"
def __init__(
self,
dataset: Dataset,
num_replicas: Optional[int] = None,
rank: Optional[int] = None,
shuffle: bool = True,
seed: int = 0,
drop_last: bool = False,
) -> None:
super().__init__(dataset, num_replicas, rank, shuffle, seed, drop_last)
self.yielded = 0
self.next_yielded = None
def __iter__(self):
self.yielded = 0
if self.next_yielded is not None:
self.yielded = self.next_yielded
self.next_yielded = None
it = super().__iter__()
for idx in itertools.islice(it, self.yielded, None):
self.yielded += 1
yield idx
def state_dict(self) -> Dict[str, Any]:
return {self._YIELDED: self.yielded}
def load_state_dict(self, state_dict: Dict[str, Any]) -> None:
if self._YIELDED not in state_dict:
raise ValueError("Invalid state_dict")
if state_dict[self._YIELDED] < 0:
raise ValueError("Cannot load state_dict with negative yielded value")
self.next_yielded = state_dict[self._YIELDED]
|