File size: 1,546 Bytes
4a28d4d | 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 | import torch
import numpy as np
import torch_utils.distributed as dist
from itertools import islice
class InfiniteSampler(torch.utils.data.Sampler):
"""
sampler from official edm repo
"""
def __init__(self, dataset, rank=0, num_replicas=1, shuffle=False, seed=0, window_size=0.5):
assert len(dataset) > 0
assert num_replicas > 0
assert 0 <= rank < num_replicas
assert 0 <= window_size <= 1
super().__init__(dataset)
self.dataset = dataset
self.rank = rank
self.num_replicas = num_replicas
self.shuffle = shuffle
self.seed = seed
self.window_size = window_size
def __iter__(self):
order = np.arange(len(self.dataset))
rnd = None
window = 0
if self.shuffle:
rnd = np.random.RandomState(self.seed)
rnd.shuffle(order)
window = int(np.rint(order.size * self.window_size))
idx = 0
while True:
i = idx % order.size
if idx % self.num_replicas == self.rank:
yield order[i]
if window >= 2:
j = (i - rnd.randint(window)) % order.size
order[i], order[j] = order[j], order[i]
idx += 1
def split_by_node(src, group=None):
rank = dist.get_rank()
world_size = dist.get_world_size()
if world_size > 1:
yield from islice(src, rank, None, world_size)
else:
yield from src
def nosplit_by_node(src, group=None):
yield from src |