MCPcopy Create free account
hub / github.com/NVlabs/DiffusionNFT / DistributedKRepeatSampler

Class DistributedKRepeatSampler

scripts/train_nft_sd3.py:118–152  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

116
117
118class DistributedKRepeatSampler(Sampler):
119 def __init__(self, dataset, batch_size, k, num_replicas, rank, seed=0):
120 self.dataset = dataset
121 self.batch_size = batch_size
122 self.k = k
123 self.num_replicas = num_replicas
124 self.rank = rank
125 self.seed = seed
126
127 self.total_samples = self.num_replicas * self.batch_size
128 assert (
129 self.total_samples % self.k == 0
130 ), f"k can not div n*b, k{k}-num_replicas{num_replicas}-batch_size{batch_size}"
131 self.m = self.total_samples // self.k
132 self.epoch = 0
133
134 def __iter__(self):
135 while True:
136 g = torch.Generator()
137 g.manual_seed(self.seed + self.epoch)
138 indices = torch.randperm(len(self.dataset), generator=g)[: self.m].tolist()
139 repeated_indices = [idx for idx in indices for _ in range(self.k)]
140
141 shuffled_indices = torch.randperm(len(repeated_indices), generator=g).tolist()
142 shuffled_samples = [repeated_indices[i] for i in shuffled_indices]
143
144 per_card_samples = []
145 for i in range(self.num_replicas):
146 start = i * self.batch_size
147 end = start + self.batch_size
148 per_card_samples.append(shuffled_samples[start:end])
149 yield per_card_samples[self.rank]
150
151 def set_epoch(self, epoch):
152 self.epoch = epoch
153
154
155def gather_tensor_to_all(tensor, world_size):

Callers 1

mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected