(self)
| 70 | torch.nn.Dropout(dropout)) |
| 71 | |
| 72 | def get_target_encoder(self): |
| 73 | if self.target_encoder is None: |
| 74 | self.target_encoder = copy.deepcopy(self.online_encoder) |
| 75 | |
| 76 | for p in self.target_encoder.parameters(): |
| 77 | p.requires_grad = False |
| 78 | return self.target_encoder |
| 79 | |
| 80 | def update_target_encoder(self, momentum: float): |
| 81 | for p, new_p in zip(self.get_target_encoder().parameters(), self.online_encoder.parameters()): |
no outgoing calls
no test coverage detected