| 237 | |
| 238 | @torch.no_grad() |
| 239 | def update_prob_t(self, lb_probs, ulb_probs): |
| 240 | ulb_prob_t = ulb_probs.mean(0) |
| 241 | self.ulb_prob_t = self.ema_p * self.ulb_prob_t + (1 - self.ema_p) * ulb_prob_t |
| 242 | |
| 243 | lb_prob_t = lb_probs.mean(0) |
| 244 | self.lb_prob_t = self.ema_p * self.lb_prob_t + (1 - self.ema_p) * lb_prob_t |
| 245 | |
| 246 | max_probs, max_idx = ulb_probs.max(dim=-1) |
| 247 | prob_max_mu_t = torch.mean(max_probs) |
| 248 | prob_max_var_t = torch.var(max_probs, unbiased=True) |
| 249 | self.prob_max_mu_t = self.ema_p * self.prob_max_mu_t + (1 - self.ema_p) * prob_max_mu_t.item() |
| 250 | self.prob_max_var_t = self.ema_p * self.prob_max_var_t + (1 - self.ema_p) * prob_max_var_t.item() |
| 251 | |
| 252 | @torch.no_grad() |
| 253 | def calculate_mask(self, probs): |