| 30 | |
| 31 | |
| 32 | class MoBY(nn.Module): |
| 33 | def __init__(self, |
| 34 | cfg, |
| 35 | encoder, |
| 36 | encoder_k, |
| 37 | contrast_momentum=0.99, |
| 38 | contrast_temperature=0.2, |
| 39 | contrast_num_negative=4096, |
| 40 | proj_num_layers=2, |
| 41 | pred_num_layers=2, |
| 42 | **kwargs): |
| 43 | super().__init__() |
| 44 | |
| 45 | self.cfg = cfg |
| 46 | |
| 47 | self.encoder = encoder |
| 48 | self.encoder_k = encoder_k |
| 49 | |
| 50 | self.contrast_momentum = contrast_momentum |
| 51 | self.contrast_temperature = contrast_temperature |
| 52 | self.contrast_num_negative = contrast_num_negative |
| 53 | |
| 54 | self.proj_num_layers = proj_num_layers |
| 55 | self.pred_num_layers = pred_num_layers |
| 56 | |
| 57 | self.projector = MoBYMLP(in_dim=self.encoder.num_features, num_layers=proj_num_layers) |
| 58 | self.projector_k = MoBYMLP(in_dim=self.encoder.num_features, num_layers=proj_num_layers) |
| 59 | self.predictor = MoBYMLP(num_layers=pred_num_layers) |
| 60 | |
| 61 | for param_q, param_k in zip(self.encoder.parameters(), self.encoder_k.parameters()): |
| 62 | param_k.data.copy_(param_q.data) # initialize |
| 63 | param_k.requires_grad = False # not update by gradient |
| 64 | |
| 65 | for param_q, param_k in zip(self.projector.parameters(), self.projector_k.parameters()): |
| 66 | param_k.data.copy_(param_q.data) |
| 67 | param_k.requires_grad = False |
| 68 | |
| 69 | if self.cfg.MODEL.SWIN.NORM_BEFORE_MLP == 'bn': |
| 70 | nn.SyncBatchNorm.convert_sync_batchnorm(self.encoder) |
| 71 | nn.SyncBatchNorm.convert_sync_batchnorm(self.encoder_k) |
| 72 | |
| 73 | nn.SyncBatchNorm.convert_sync_batchnorm(self.projector) |
| 74 | nn.SyncBatchNorm.convert_sync_batchnorm(self.projector_k) |
| 75 | nn.SyncBatchNorm.convert_sync_batchnorm(self.predictor) |
| 76 | |
| 77 | self.K = int(self.cfg.DATA.TRAINING_IMAGES * 1. / dist.get_world_size() / self.cfg.DATA.BATCH_SIZE) * self.cfg.TRAIN.EPOCHS |
| 78 | self.k = int(self.cfg.DATA.TRAINING_IMAGES * 1. / dist.get_world_size() / self.cfg.DATA.BATCH_SIZE) * self.cfg.TRAIN.START_EPOCH |
| 79 | |
| 80 | # create the queue |
| 81 | self.register_buffer("queue1", torch.randn(256, self.contrast_num_negative)) |
| 82 | self.register_buffer("queue2", torch.randn(256, self.contrast_num_negative)) |
| 83 | self.queue1 = F.normalize(self.queue1, dim=0) |
| 84 | self.queue2 = F.normalize(self.queue2, dim=0) |
| 85 | |
| 86 | self.register_buffer("queue_ptr", torch.zeros(1, dtype=torch.long)) |
| 87 | |
| 88 | @torch.no_grad() |
| 89 | def _momentum_update_key_encoder(self): |