(q_logit, p_logit)
| 25 | return torch.from_numpy(d) |
| 26 | |
| 27 | def kl_div_with_logit(q_logit, p_logit): |
| 28 | |
| 29 | q = F.softmax(q_logit, dim=1) |
| 30 | logq = F.log_softmax(q_logit, dim=1) |
| 31 | logp = F.log_softmax(p_logit, dim=1) |
| 32 | |
| 33 | qlogq = ( q *logq).sum(dim=1).mean(dim=0) |
| 34 | qlogp = ( q *logp).sum(dim=1).mean(dim=0) |
| 35 | |
| 36 | return qlogq - qlogp |
| 37 | def vat_loss(model, ul_x, ul_y, xi=1e-6, eps=6, num_iters=1): |
| 38 | |
| 39 | # find r_adv |