| 15 | |
| 16 | |
| 17 | def consistency_loss(logits_s, logits_w, name='ce', T=1.0, p_cutoff=0.0, use_hard_labels=True): |
| 18 | assert name in ['ce', 'L2'] |
| 19 | logits_w = logits_w.detach() |
| 20 | if name == 'L2': |
| 21 | assert logits_w.size() == logits_s.size() |
| 22 | return F.mse_loss(logits_s, logits_w, reduction='mean') |
| 23 | |
| 24 | elif name == 'L2_mask': |
| 25 | pass |
| 26 | |
| 27 | elif name == 'ce': |
| 28 | pseudo_label = torch.softmax(logits_w, dim=-1) |
| 29 | max_probs, max_idx = torch.max(pseudo_label, dim=-1) |
| 30 | mask = max_probs.ge(p_cutoff).float() |
| 31 | select = max_probs.ge(p_cutoff).long() |
| 32 | # strong_prob, strong_idx = torch.max(torch.softmax(logits_s, dim=-1), dim=-1) |
| 33 | # strong_select = strong_prob.ge(p_cutoff).long() |
| 34 | # select = select * strong_select * (strong_idx == max_idx) |
| 35 | if use_hard_labels: |
| 36 | masked_loss = ce_loss(logits_s, max_idx, use_hard_labels, reduction='none') * mask |
| 37 | else: |
| 38 | pseudo_label = torch.softmax(logits_w / T, dim=-1) |
| 39 | masked_loss = ce_loss(logits_s, pseudo_label, use_hard_labels) * mask |
| 40 | return masked_loss.mean(), mask.mean(), select, max_idx.long() |
| 41 | |
| 42 | else: |
| 43 | assert Exception('Not Implemented consistency_loss') |
| 44 | |