(self, probs)
| 261 | |
| 262 | @torch.no_grad() |
| 263 | def distribution_alignment(self, probs): |
| 264 | # da |
| 265 | probs = probs * self.lb_prob_t / self.ulb_prob_t |
| 266 | probs = probs / probs.sum(dim=1, keepdim=True) |
| 267 | return probs.detach() |
| 268 | |
| 269 | |
| 270 | @torch.no_grad() |