You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch重训练MTCNN模型实现三个人脸类别概率输出

首先要明确一个核心问题:MTCNN是专门做人脸检测、关键点定位的模型,本身结构设计就不是为了人脸分类任务,你现有代码里只解冻onet.dense6_3.bias的逻辑完全没用——这一层是MTCNN输出人脸检测置信度的层,输出维度只有1,根本没法直接输出3分类结果,硬改MTCNN做分类训练效率和精度都会很差。

标准实现流程是用MTCNN负责人脸检测对齐,搭配facenet_pytorch自带的InceptionResnetV1(FaceNet)做人脸特征提取,加一个3分类头完成任务,最终输出的就是你要的三个类别概率张量,具体实现如下:

完整实现步骤

1. 基础配置与数据集加载

你原来的collate_fn只取第一个元素会丢失batch数据,这里自定义数据集类,直接在加载阶段用MTCNN完成人脸对齐,自动拆分训练/验证集:

import torch
import torch.nn as nn
import torch.optim as optim
from facenet_pytorch import InceptionResnetV1, MTCNN
from torch.utils.data import DataLoader, random_split
from torchvision import datasets, transforms
from tqdm import tqdm
import os

# 基础超参数
workers = 0 if os.name == 'nt' else 4
device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')
CLASS_NUM = 3
BATCH_SIZE = 16
EPOCHS = 10
LR = 1e-4

# 初始化MTCNN,仅用于人脸检测对齐
mtcnn = MTCNN(
    image_size=160, margin=20, min_face_size=20,
    thresholds=[0.6, 0.7, 0.7], factor=0.709, post_process=True,
    device=device
)

# 自定义人脸数据集类
class FaceDataset(torch.utils.data.Dataset):
    def __init__(self, root_dir):
        self.dataset = datasets.ImageFolder(root_dir)
        self.class_to_idx = self.dataset.class_to_idx
        self.idx_to_class = {i:c for c,i in self.class_to_idx.items()}

    def __len__(self):
        return len(self.dataset)

    def __getitem__(self, idx):
        img, label = self.dataset[idx]
        # 对齐人脸,未检测到人脸时直接resize原图避免训练中断
        img_aligned = mtcnn(img)
        if img_aligned is None:
            img_aligned = transforms.Resize((160,160))(transforms.ToTensor()(img))
        return img_aligned, label

# 加载数据集,按8:2拆分训练/验证集
dataset = FaceDataset('data/images/')
train_size = int(0.8 * len(dataset))
val_size = len(dataset) - train_size
train_set, val_set = random_split(dataset, [train_size, val_size])
train_loader = DataLoader(train_set, batch_size=BATCH_SIZE, shuffle=True, num_workers=workers)
val_loader = DataLoader(val_set, batch_size=BATCH_SIZE, shuffle=False, num_workers=workers)

2. 构建三分类模型

加载在大规模人脸数据集上预训练的InceptionResnetV1权重,替换最后一层分类头为3分类输出,小数据集下冻结前面的特征提取层只训分类头,能有效避免过拟合:

# 加载预训练FaceNet模型,直接指定分类头输出维度为3
model = InceptionResnetV1(
    classify=True,
    pretrained='vggface2',
    num_classes=CLASS_NUM
).to(device)

# 冻结特征提取层,仅训练最后分类层
for name, param in model.named_parameters():
    param.requires_grad = True if name.startswith('logits') else False

# 定义损失、优化器、概率激活层
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(filter(lambda p: p.requires_grad, model.parameters()), lr=LR)
softmax = nn.Softmax(dim=1) # 输出转成0-1之间的概率值,和为1

3. 训练循环

for epoch in range(EPOCHS):
    model.train()
    train_loss, train_correct = 0.0, 0
    for imgs, labels in tqdm(train_loader):
        imgs, labels = imgs.to(device), labels.to(device)
        optimizer.zero_grad()
        outputs = model(imgs)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()

        train_loss += loss.item() * imgs.size(0)
        _, preds = torch.max(outputs, 1)
        train_correct += torch.sum(preds == labels.data)

    # 验证阶段
    model.eval()
    val_loss, val_correct = 0.0, 0
    with torch.no_grad():
        for imgs, labels in val_loader:
            imgs, labels = imgs.to(device), labels.to(device)
            outputs = model(imgs)
            loss = criterion(outputs, labels)
            val_loss += loss.item() * imgs.size(0)
            _, preds = torch.max(outputs, 1)
            val_correct += torch.sum(preds == labels.data)

    # 打印指标
    print(f'Epoch {epoch+1}/{EPOCHS}:')
    print(f'Train Loss: {train_loss/train_size:.4f} Acc: {train_correct.double()/train_size:.4f}')
    print(f'Val Loss: {val_loss/val_size:.4f} Acc: {val_correct.double()/val_size:.4f}\n')

# 保存训练好的权重
torch.save(model.state_dict(), 'face_3class.pth')

4. 推理使用

训练完成后,输入图像经过MTCNN对齐、模型推理、Softmax激活后,输出就是你要的[prob1, prob2, prob3]格式张量,顺序和数据集文件夹名排序一致(即faces1、faces2、faces3对应索引0、1、2):

# 加载训练好的模型
model.load_state_dict(torch.load('face_3class.pth'))
model.eval()

def predict(img):
    img_aligned = mtcnn(img)
    if img_aligned is None:
        return [0.0, 0.0, 0.0] # 未检测到人脸返回全0
    with torch.no_grad():
        logits = model(img_aligned.unsqueeze(0).to(device))
        probs = softmax(logits).squeeze(0).cpu().tolist()
    return probs

补充说明:如果单类样本量超过50张,可以解冻模型最后2-3个卷积块一起微调,学习率调低到1e-5即可,精度会进一步提升。如果你一定要基于MTCNN改分类头,只需要把ONet的最后一层dense6_3替换成输出维度为3的全连接层即可,但MTCNN的特征是为人脸检测优化的,分类效果会远差于上面的FaceNet方案。

内容的提问来源于stack exchange,提问作者Norhther

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.30 20:24:19