deepinsight / deepinsight/insightface

arcface loss 为都求 cos(target_logit + self.m2) 不是只有yi 求 cos(target_logit + self.m2) , 而非本类求 cos(target_logit)?

Open
#2,513 2 comments 1 reaction 0 assignees View on GitHub
Dominant language
Python
Stars
29.7k
Forks
6.1k
PR merge metrics
No merged PRs in 30d

Description

![v2-7e2870d20bc82f7ba3862c1be68fa2af_720w](https://github.com/deepinsight/insightface/assets/32585434/fd177da7-19f6-4a29-b431-27e3ef38d19b)
if self.m1 == 1.0 and self.m3 == 0.0:
with torch.no_grad():
target_logit.arccos_()
logits.arccos_()
final_target_logit = target_logit + self.m2
#cos(target_logit + self.m2) 不是只有yi 求 cos(target_logit + self.m2) , 而非本类求 cos(target_logit)?
logits[index_positive, labels[index_positive].view(-1)] = final_target_logit
logits.cos_()

这两种实现方式差距很大啊

class ArcMarginProduct(nn.Module):
r"""Implement of large margin arc distance: :
Args:
in_features: size of each input sample
out_features: size of each output sample
s: norm of input feature
m: additive angular margin
cos(theta + m)
"""
def __init__(self, in_features, out_features, s=30.0, m=0.50, easy_margin=False):
super(ArcMarginProduct, self).__init__()

self.in_features = in_features # 特征输入通道数
self.out_features = out_features # 特征输出通道数
self.s = s # 输入特征范数 ||x_i||
self.m = m # 加性角度边距 m (additive angular margin)
self.weight = Parameter(torch.FloatTensor(out_features, in_features)) # FC 权重
nn.init.xavier_uniform_(self.weight) # Xavier 初始化 FC 权重

self.easy_margin = easy_margin
self.cos_m = math.cos(m)
self.sin_m = math.sin(m)
self.th = math.cos(math.pi - m)
self.mm = math.sin(math.pi - m) * m

def forward(self, input, label):
# --------------------------- cos(theta) & phi(theta) ---------------------------
# 分别归一化输入特征 xi 和 FC 权重 W, 二者点乘得到 cosθ, 即预测值 Logit
cosine = F.linear(F.normalize(input), F.normalize(self.weight))
# 由 cosθ 计算相应的 sinθ
sine = torch.sqrt(1.0 - torch.pow(cosine, 2))
# 展开计算 cos(θ+m) = cosθ*cosm - sinθ*sinm, 其中包含了 Target Logit (cos(θyi+ m)) (由于输入特征 xi 的非真实类也参与了计算, 最后计算新 Logit 时需使用 One-Hot 区别)
phi = cosine * self.cos_m - sine * self.sin_m
# 是否松弛约束??
if self.easy_margin:
phi = torch.where(cosine > 0, phi, cosine)
else:
phi = torch.where(cosine > self.th, phi, cosine - self.mm)

# --------------------------- convert label to one-hot ---------------------------
# one_hot = torch.zeros(cosine.size(), requires_grad=True, device='cuda')
# 将 labels 转换为独热编码, 用于区分是否为输入特征 xi 对应的真实类别 yi
one_hot = torch.zeros(cosine.size(), device='cuda')
one_hot.scatter_(1, label.view(-1, 1).long(), 1)

# -------------torch.where(out_i = {x_i if condition_i else y_i) -------------
# 计算新 Logit
# - 只有输入特征 xi 对应的真实类别 yi (one_hot=1) 采用新 Target Logit cos(θ_yi + m)
# - 其余并不对应输入特征 xi 的真实类别的类 (one_hot=0) 则仍保持原 Logit cosθ_j
output = (one_hot * phi) + ((1.0 - one_hot) * cosine) # can use torch.where if torch.__version__ > 0.4
# 使用 s rescale 放缩新 Logit, 以馈入传统 Softmax Loss 计算
output *= self.s

return output

Contributor guide

No contributing guide indexed for this repository

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.