| 101 | |
| 102 | @torch.no_grad() |
| 103 | def _dequeue_and_enqueue(self, keys1, keys2): |
| 104 | # gather keys before updating queue |
| 105 | keys1 = dist_collect(keys1) |
| 106 | keys2 = dist_collect(keys2) |
| 107 | |
| 108 | batch_size = keys1.shape[0] |
| 109 | |
| 110 | ptr = int(self.queue_ptr) |
| 111 | assert self.contrast_num_negative % batch_size == 0 # for simplicity |
| 112 | |
| 113 | # replace the keys at ptr (dequeue and enqueue) |
| 114 | self.queue1[:, ptr:ptr + batch_size] = keys1.T |
| 115 | self.queue2[:, ptr:ptr + batch_size] = keys2.T |
| 116 | ptr = (ptr + batch_size) % self.contrast_num_negative # move pointer |
| 117 | |
| 118 | self.queue_ptr[0] = ptr |
| 119 | |
| 120 | def contrastive_loss(self, q, k, queue): |
| 121 | |