CoolFace
Apppublic

RabbitRUI/ruispace

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
losses.py43 linesDownload Raw Back to arcface_torch
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