(mask, logits_s, logits_w, prob_model, label_hist)
| 22 | |
| 23 | |
| 24 | def entropy_loss(mask, logits_s, logits_w, prob_model, label_hist): |
| 25 | # select samples |
| 26 | logits_s = logits_s[mask] |
| 27 | |
| 28 | prob_s = logits_s.softmax(dim=-1) |
| 29 | _, pred_label_s = torch.max(prob_s, dim=-1) |
| 30 | |
| 31 | hist_s = torch.bincount(pred_label_s, minlength=logits_s.shape[1]).to(logits_w.dtype) |
| 32 | hist_s = hist_s / hist_s.sum() |
| 33 | |
| 34 | # modulate prob model |
| 35 | prob_model = prob_model.reshape(1, -1) |
| 36 | label_hist = label_hist.reshape(1, -1) |
| 37 | # prob_model_scaler = torch.nan_to_num(1 / label_hist, nan=0.0, posinf=0.0, neginf=0.0).detach() |
| 38 | prob_model_scaler = replace_inf_to_zero(1 / label_hist).detach() |
| 39 | mod_prob_model = prob_model * prob_model_scaler |
| 40 | mod_prob_model = mod_prob_model / mod_prob_model.sum(dim=-1, keepdim=True) |
| 41 | |
| 42 | # modulate mean prob |
| 43 | mean_prob_scaler_s = replace_inf_to_zero(1 / hist_s).detach() |
| 44 | # mean_prob_scaler_s = torch.nan_to_num(1 / hist_s, nan=0.0, posinf=0.0, neginf=0.0).detach() |
| 45 | mod_mean_prob_s = prob_s.mean(dim=0, keepdim=True) * mean_prob_scaler_s |
| 46 | mod_mean_prob_s = mod_mean_prob_s / mod_mean_prob_s.sum(dim=-1, keepdim=True) |
| 47 | |
| 48 | loss = mod_prob_model * torch.log(mod_mean_prob_s + 1e-12) |
| 49 | loss = loss.sum(dim=1) |
| 50 | return loss.mean(), hist_s.mean() |
| 51 | |
| 52 | |
| 53 |
no test coverage detected