1def softmax(x, dim=None):
2 e_x = torch.exp(x - torch.max(x, dim=dim, keepdim=True)[0])
3 return e_x / e_x.sum(dim=dim, keepdim=True)
4
5def ghostmax(x, dim=None):
6 e_x = torch.exp(x - torch.max(x, dim=dim, keepdim=True)[0])
7 return e_x / (1+e_x.sum(dim=dim, keepdim=True) )
8