NYCU-MLLab's picture
Upload folder using huggingface_hub
4a28d4d verified
Raw
History Blame Contribute Delete
1.55 kB
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