RabbitRUI/ruispace
0
1import torch2from torch import nn3 4 5def get_loss(name):6 if name == "cosface":7 return CosFace()8 elif name == "arcface":9 return ArcFace()10 else:11 raise ValueError()12 13 14class CosFace(nn.Module):15 def __init__(self, s=64.0, m=0.40):16 super(CosFace, self).__init__()17 self.s = s18 self.m = m19 20 def forward(self, cosine, label):21 index = torch.where(label != -1)[0]22 m_hot = torch.zeros(index.size()[0], cosine.size()[1], device=cosine.device)23 m_hot.scatter_(1, label[index, None], self.m)24 cosine[index] -= m_hot25 ret = cosine * self.s26 return ret27 28 29class ArcFace(nn.Module):30 def __init__(self, s=64.0, m=0.5):31 super(ArcFace, self).__init__()32 self.s = s33 self.m = m34 35 def forward(self, cosine: torch.Tensor, label):36 index = torch.where(label != -1)[0]37 m_hot = torch.zeros(index.size()[0], cosine.size()[1], device=cosine.device)38 m_hot.scatter_(1, label[index, None], self.m)39 cosine.acos_()40 cosine[index] += m_hot41 cosine.cos_().mul_(self.s)42 return cosine43 