svjack/Kolors-Controlnet_and_IPA
0
1import queue as Queue2import threading3import torch4from torch.utils.data import DataLoader5 6 7class PrefetchGenerator(threading.Thread):8 """A general prefetch generator.9 10 Reference: https://stackoverflow.com/questions/7323664/python-generator-pre-fetch11 12 Args:13 generator: Python generator.14 num_prefetch_queue (int): Number of prefetch queue.15 """16 17 def __init__(self, generator, num_prefetch_queue):18 threading.Thread.__init__(self)19 self.queue = Queue.Queue(num_prefetch_queue)20 self.generator = generator21 self.daemon = True22 self.start()23 24 def run(self):25 for item in self.generator:26 self.queue.put(item)27 self.queue.put(None)28 29 def __next__(self):30 next_item = self.queue.get()31 if next_item is None:32 raise StopIteration33 return next_item34 35 def __iter__(self):36 return self37 38 39class PrefetchDataLoader(DataLoader):40 """Prefetch version of dataloader.41 42 Reference: https://github.com/IgorSusmelj/pytorch-styleguide/issues/5#43 44 TODO:45 Need to test on single gpu and ddp (multi-gpu). There is a known issue in46 ddp.47 48 Args:49 num_prefetch_queue (int): Number of prefetch queue.50 kwargs (dict): Other arguments for dataloader.51 """52 53 def __init__(self, num_prefetch_queue, **kwargs):54 self.num_prefetch_queue = num_prefetch_queue55 super(PrefetchDataLoader, self).__init__(**kwargs)56 57 def __iter__(self):58 return PrefetchGenerator(super().__iter__(), self.num_prefetch_queue)59 60 61class CPUPrefetcher():62 """CPU prefetcher.63 64 Args:65 loader: Dataloader.66 """67 68 def __init__(self, loader):69 self.ori_loader = loader70 self.loader = iter(loader)71 72 def next(self):73 try:74 return next(self.loader)75 except StopIteration:76 return None77 78 def reset(self):79 self.loader = iter(self.ori_loader)80 81 82class CUDAPrefetcher():83 """CUDA prefetcher.84 85 Reference: https://github.com/NVIDIA/apex/issues/304#86 87 It may consume more GPU memory.88 89 Args:90 loader: Dataloader.91 opt (dict): Options.92 """93 94 def __init__(self, loader, opt):95 self.ori_loader = loader96 self.loader = iter(loader)97 self.opt = opt98 self.stream = torch.cuda.Stream()99 self.device = torch.device('cuda' if opt['num_gpu'] != 0 else 'cpu')100 self.preload()101 102 def preload(self):103 try:104 self.batch = next(self.loader) # self.batch is a dict105 except StopIteration:106 self.batch = None107 return None108 # put tensors to gpu109 with torch.cuda.stream(self.stream):110 for k, v in self.batch.items():111 if torch.is_tensor(v):112 self.batch[k] = self.batch[k].to(device=self.device, non_blocking=True)113 114 def next(self):115 torch.cuda.current_stream().wait_stream(self.stream)116 batch = self.batch117 self.preload()118 return batch119 120 def reset(self):121 self.loader = iter(self.ori_loader)122 self.preload()123 