(self, momentum: float)
| 386 | |
| 387 | |
| 388 | def update_teacher(self, momentum: float): |
| 389 | if not self.enable_teacher: |
| 390 | return |
| 391 | |
| 392 | with torch.no_grad(): |
| 393 | for student_param, teacher_param in zip(self.trunk.parameters(), self.teacher_trunk.parameters()): |
| 394 | teacher_param.data = momentum * teacher_param.data + (1 - momentum) * student_param.data |
| 395 | |
| 396 | if self.vtp_config.training.train_clip: |
| 397 | for student_param, teacher_param in zip(self.proj.parameters(), self.teacher_proj.parameters()): |
| 398 | teacher_param.data = momentum * teacher_param.data + (1 - momentum) * student_param.data |
| 399 | |
| 400 | for student_param, teacher_param in zip(self.dino_head.parameters(), self.teacher_dino_head.parameters()): |
| 401 | teacher_param.data = momentum * teacher_param.data + (1 - momentum) * student_param.data |
| 402 | |
| 403 | def get_ssl_params(self): |
| 404 | ssl_params = [] |
nothing calls this directly
no outgoing calls
no test coverage detected