(x, y, alpha=3)
| 422 | return CL_loss, CL_acc |
| 423 | |
| 424 | def sce_loss(x, y, alpha=3): |
| 425 | x = F.normalize(x, p=2, dim=-1) |
| 426 | y = F.normalize(y, p=2, dim=-1) |
| 427 | |
| 428 | # loss = - (x * y).sum(dim=-1) |
| 429 | # loss = (x_h - y_h).norm(dim=1).pow(alpha) |
| 430 | |
| 431 | loss = (1 - (x * y).sum(dim=-1)).pow_(alpha) |
| 432 | |
| 433 | loss = loss.mean() |
| 434 | return loss |
| 435 | |
| 436 | def create_k_shot_mask(labels, args): |
| 437 | k = args.shot |