ArcFace损失函数公式解析与实战:从原理到PyTorch实现

1次阅读
没有评论

共计 2325 个字符,预计需要花费 6 分钟才能阅读完成。

image.webp

背景痛点:为什么需要 ArcFace?

在人脸识别任务中,传统 Softmax Loss 存在两个主要问题:

  • 类内方差大 :同一人的不同照片(如光照、角度变化)在特征空间(feature space) 中分布分散
  • 类间相似度高:不同人的面部特征容易相互重叠,导致分类边界模糊

这就像试图用钝刀切蛋糕——虽然能勉强分开,但边缘总是粘连。ArcFace 通过引入角度间隔 (angular margin) 机制,让同类样本更紧凑、异类更疏远。

公式解析:ArcFace 的数学本质

ArcFace 的核心公式可分解为四步:

  1. 特征与权重归一化
    $$ \begin{aligned}
    W_j &\leftarrow \frac{W_j}{|W_j|} \
    x_i &\leftarrow \frac{x_i}{|x_i|}
    \end{aligned} $$
    这相当于将特征和分类权重都投影到单位超球面上

  2. 计算余弦相似度
    $$ \cos\theta_j = W_j^T x_i $$
    此时 θ 就是特征向量与类别权重的夹角

  3. 添加角度间隔
    $$ \cos(\theta_{y_i} + m) $$
    其中 m 是预设的 margin 值(通常 0.3~0.5),相当于在决策边界上挖走一块区域

  4. 重新缩放计算概率
    $$ L_i = -\log\left(\frac{e^{s\cdot\cos(\theta_{y_i}+m)}}{e^{s\cdot\cos(\theta_{y_i}+m)} + \sum_{j\neq y_i} e^{s\cdot\cos\theta_j}}\right) $$
    缩放因子 s(通常取 64)用于放大梯度信号

PyTorch 实现详解

import torch
import torch.nn as nn
import torch.nn.functional as F

class ArcFace(nn.Module):
    def __init__(self, feat_dim, num_classes, margin=0.3, scale=64):
        super().__init__()
        self.W = nn.Parameter(torch.randn(feat_dim, num_classes))
        self.margin = margin
        self.scale = scale

    def forward(self, x, labels):
        # 1. 权重归一化 (num_classes, feat_dim)
        W_norm = F.normalize(self.W, dim=0) 

        # 2. 特征归一化 (batch_size, feat_dim)
        x_norm = F.normalize(x, dim=1)

        # 3. 计算余弦相似度 (batch_size, num_classes)
        cos_theta = torch.mm(x_norm, W_norm)

        # 4. 提取真实类别的 cos 值
        one_hot = F.one_hot(labels, num_classes=self.W.shape[1])
        cos_theta_yi = cos_theta[one_hot.bool()].view(-1, 1)

        # 5. 计算 cos(theta + margin)
        cos_theta_m = torch.cos(torch.acos(cos_theta_yi) + self.margin)

        # 6. 替换原矩阵中的对应值
        output = cos_theta.scatter(1, labels.view(-1, 1), cos_theta_m)

        # 7. 缩放后计算交叉熵
        loss = F.cross_entropy(self.scale * output, labels)
        return loss

关键维度说明:
– 输入特征 x: (batch_size, feat_dim)
– 权重矩阵 W: (feat_dim, num_classes)
– cos_theta: (batch_size, num_classes)
– one_hot: (batch_size, num_classes)

对比实验:t-SNE 可视化

在 CASIA-WebFace 数据集上的对比:

  • Softmax Loss
    特征呈模糊的星型分布,类间边界不清晰

  • ArcFace(m=0.5)
    同类样本紧密聚集,不同类形成明显间隔

ArcFace 损失函数公式解析与实战:从原理到 PyTorch 实现(注:此为示意图,实际效果需运行代码验证)

避坑指南

  1. Margin 值选择
  2. 过大(>0.8)会导致训练震荡,建议从 0.3 开始逐步增加
  3. 可通过验证集准确率变化判断是否合适

  4. 特征归一化缺失

  5. 未归一化时,cos_theta 可能超出 [-1,1] 范围,导致 acos()计算报错
  6. 解决方案:检查 F.normalize 是否应用到所有特征和权重

  7. 小数据集过拟合

  8. 可尝试:
    • 减小 margin 值
    • 增加权重衰减(weight decay)
    • 使用标签平滑(label smoothing)

延伸思考:与 Triplet Loss 结合

工业级人脸识别系统常组合使用:
ArcFace:保证全局类别可分性
Triplet Loss:优化局部样本对关系

实现方式:

def forward(self, x, labels):
    arc_loss = self.arcface(x, labels)
    triplet_loss = self.triplet(x, labels)
    return arc_loss + 0.1 * triplet_loss  # 需调权重

这种组合能同时利用类别信息和样本间相对关系,通常能提升 1 -2% 的识别准确率。

结语

ArcFace 通过几何直观的角度间隔,优雅地解决了人脸识别中的特征判别问题。建议读者:
1. 从 margin=0.3、scale=64 的基础配置开始实验
2. 重点关注特征归一化的实现细节
3. 可视化训练过程中的特征分布变化

完整的可运行代码已上传至 GitHub 仓库(虚构链接):
https://github.com/example/arcface-pytorch

正文完
 0
评论(没有评论)