(self)
| 78 | torch.nn.Dropout(dropout)) |
| 79 | |
| 80 | def get_target_encoder(self): |
| 81 | if self.target_encoder is None: |
| 82 | self.target_encoder = copy.deepcopy(self.online_encoder) |
| 83 | |
| 84 | for p in self.target_encoder.parameters(): |
| 85 | p.requires_grad = False |
| 86 | return self.target_encoder |
| 87 | |
| 88 | def update_target_encoder(self, momentum: float): |
| 89 | for p, new_p in zip(self.get_target_encoder().parameters(), self.online_encoder.parameters()): |
no outgoing calls
no test coverage detected