| 18 | |
| 19 | |
| 20 | def consistency_loss(logits_s, logits_w, class_acc, p_target, p_model, name='ce', |
| 21 | T=1.0, p_cutoff=0.0, use_hard_labels=True, use_DA=False): |
| 22 | assert name in ['ce', 'L2'] |
| 23 | logits_w = logits_w.detach() |
| 24 | if name == 'L2': |
| 25 | assert logits_w.size() == logits_s.size() |
| 26 | return F.mse_loss(logits_s, logits_w, reduction='mean') |
| 27 | |
| 28 | elif name == 'L2_mask': |
| 29 | pass |
| 30 | |
| 31 | elif name == 'ce': |
| 32 | pseudo_label = torch.softmax(logits_w, dim=-1) |
| 33 | if use_DA: |
| 34 | if p_model == None: |
| 35 | p_model = torch.mean(pseudo_label.detach(), dim=0) |
| 36 | else: |
| 37 | p_model = p_model * 0.999 + torch.mean(pseudo_label.detach(), dim=0) * 0.001 |
| 38 | pseudo_label = pseudo_label * p_target / p_model |
| 39 | pseudo_label = (pseudo_label / pseudo_label.sum(dim=-1, keepdim=True)) |
| 40 | |
| 41 | max_probs, max_idx = torch.max(pseudo_label, dim=-1) |
| 42 | # mask = max_probs.ge(p_cutoff * (class_acc[max_idx] + 1.) / 2).float() # linear |
| 43 | # mask = max_probs.ge(p_cutoff * (1 / (2. - class_acc[max_idx]))).float() # low_limit |
| 44 | mask = max_probs.ge(p_cutoff * (class_acc[max_idx] / (2. - class_acc[max_idx]))).float() # convex |
| 45 | # mask = max_probs.ge(p_cutoff * (torch.log(class_acc[max_idx] + 1.) + 0.5)/(math.log(2) + 0.5)).float() # concave |
| 46 | select = max_probs.ge(p_cutoff).long() |
| 47 | if use_hard_labels: |
| 48 | masked_loss = ce_loss(logits_s, max_idx, use_hard_labels, reduction='none') * mask |
| 49 | else: |
| 50 | pseudo_label = torch.softmax(logits_w / T, dim=-1) |
| 51 | masked_loss = ce_loss(logits_s, pseudo_label, use_hard_labels) * mask |
| 52 | return masked_loss.mean(), mask.mean(), select, max_idx.long(), p_model |
| 53 | |
| 54 | else: |
| 55 | assert Exception('Not Implemented consistency_loss') |