基于Transformer架构的人脸识别系统,它使用电脑摄像头实时捕获视频流,并进行人脸识别。这个实现结合了现代Transformer架构和传统人脸检测技术。

import cv2
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torchvision import transforms
from facenet_pytorch import MTCNN
import matplotlib.pyplot as plt
from PIL import Image
import time
import os
import pickle
from collections import defaultdict

# 设备配置
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
print(f"使用设备: {device}")

# 人脸检测器
mtcnn = MTCNN(keep_all=True, device=device)

# 人脸识别模型 - Vision Transformer (ViT)
class FaceViT(nn.Module):
    def __init__(self, num_classes, embedding_dim=128, num_heads=8, num_layers=6):
        super(FaceViT, self).__init__()
        
        # 图像预处理
        self.patch_size = 16
        self.patch_dim = 3 * self.patch_size**2
        
        # 位置编码
        self.position_embedding = nn.Parameter(torch.randn(1, (224//self.patch_size)**2 + 1, embedding_dim))
        
        # Transformer编码器
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=embedding_dim, 
            nhead=num_heads,
            dim_feedforward=embedding_dim*4,
            dropout=0.1,
            activation='gelu'
        )
        self.transformer_encoder = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)
        
        # 分类头
        self.classifier = nn.Sequential(
            nn.Linear(embedding_dim, embedding_dim),
            nn.ReLU(),
            nn.Linear(embedding_dim, num_classes)
        )
        
        # 嵌入层
        self.embedding_layer = nn.Linear(self.patch_dim, embedding_dim)
        
        # 分类标记
        self.cls_token = nn.Parameter(torch.randn(1, 1, embedding_dim))
    
    def forward(self, x):
        # 将图像分割为小块
        B, C, H, W = x.shape
        x = x.unfold(2, self.patch_size, self.patch_size).unfold(3, self.patch_size, self.patch_size)
        x = x.contiguous().view(B, C, -1, self.patch_size, self.patch_size)
        x = x.permute(0, 2, 3, 4, 1).contiguous().view(B, -1, self.patch_dim)
        
        # 嵌入转换
        x = self.embedding_layer(x)
        
        # 添加分类标记
        cls_tokens = self.cls_token.expand(B, -1, -1)
        x = torch.cat((cls_tokens, x), dim=1)
        
        # 添加位置编码
        x = x + self.position_embedding
        
        # Transformer编码
        x = self.transformer_encoder(x)
        
        # 使用分类标记进行分类
        cls_output = x[:, 0]
        
        # 分类
        logits = self.classifier(cls_output)
        
        # 嵌入向量(用于人脸比对)
        embeddings = F.normalize(cls_output, p=2, dim=1)
        
        return logits, embeddings

# 人脸数据库类
class FaceDatabase:
    def __init__(self, threshold=0.7):
        self.database = defaultdict(list)
        self.threshold = threshold
        self.embeddings = []
        self.labels = []
    
    def add_face(self, embedding, label):
        """添加新的人脸到数据库"""
        self.database[label].append(embedding)
        self.embeddings.append(embedding)
        self.labels.append(label)
    
    def recognize(self, embedding):
        """识别人脸"""
        if not self.embeddings:
            return "Unknown", 0.0
        
        # 计算相似度
        similarities = []
        for db_emb in self.embeddings:
            sim = F.cosine_similarity(embedding, db_emb, dim=0).item()
            similarities.append(sim)
        
        # 找到最相似的人脸
        max_idx = np.argmax(similarities)
        max_sim = similarities[max_idx]
        
        if max_sim > self.threshold:
            return self.labels[max_idx], max_sim
        else:
            return "Unknown", max_sim
    
    def save(self, path):
        """保存数据库"""
        with open(path, 'wb') as f:
            pickle.dump({
                'database': dict(self.database),
                'threshold': self.threshold,
                'embeddings': self.embeddings,
                'labels': self.labels
            }, f)
    
    def load(self, path):
        """加载数据库"""
        with open(path, 'rb') as f:
            data = pickle.load(f)
            self.database = defaultdict(list, data['database'])
            self.threshold = data['threshold']
            self.embeddings = data['embeddings']
            self.labels = data['labels']

# 人脸注册函数
def register_face(model, database, name, capture_time=5):
    """注册新的人脸"""
    cap = cv2.VideoCapture(0)
    if not cap.isOpened():
        print("无法打开摄像头")
        return
    
    print(f"正在注册: {name}")
    print("请正对摄像头...")
    
    embeddings = []
    start_time = time.time()
    
    while time.time() - start_time < capture_time:
        ret, frame = cap.read()
        if not ret:
            print("无法获取帧")
            break
        
        # 检测人脸
        frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
        boxes, _ = mtcnn.detect(frame_rgb)
        
        if boxes is not None:
            for box in boxes:
                x1, y1, x2, y2 = map(int, box)
                cv2.rectangle(frame, (x1, y1), (x2, y2), (0, 255, 0), 2)
                
                # 提取人脸
                face = frame_rgb[y1:y2, x1:x2]
                if face.size == 0:
                    continue
                
                # 预处理
                face_pil = Image.fromarray(face)
                face_tensor = transforms.Compose([
                    transforms.Resize((224, 224)),
                    transforms.ToTensor(),
                    transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
                ])(face_pil).unsqueeze(0).to(device)
                
                # 获取嵌入向量
                with torch.no_grad():
                    _, embedding = model(face_tensor)
                    embeddings.append(embedding.cpu())
        
        cv2.putText(frame, f"注册: {name} | {int(capture_time - (time.time() - start_time))}秒", 
                   (10, 30), cv2.FONT_HERSHEY_SIMPLEX, 1, (0, 255, 0), 2)
        cv2.imshow('Face Registration', frame)
        
        if cv2.waitKey(1) & 0xFF == ord('q'):
            break
    
    cap.release()
    cv2.destroyAllWindows()
    
    if embeddings:
        # 计算平均嵌入向量
        avg_embedding = torch.mean(torch.stack(embeddings), dim=0)
        database.add_face(avg_embedding, name)
        print(f"成功注册: {name}")
        return True
    else:
        print("未检测到人脸,注册失败")
        return False

# 实时人脸识别函数
def realtime_recognition(model, database):
    """实时人脸识别"""
    cap = cv2.VideoCapture(0)
    if not cap.isOpened():
        print("无法打开摄像头")
        return
    
    print("实时人脸识别中... 按 'q' 退出")
    
    # 存储最近识别结果以平滑显示
    recognition_history = []
    max_history = 5
    
    while True:
        ret, frame = cap.read()
        if not ret:
            print("无法获取帧")
            break
        
        # 检测人脸
        frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
        boxes, _ = mtcnn.detect(frame_rgb)
        
        if boxes is not None:
            for box in boxes:
                x1, y1, x2, y2 = map(int, box)
                
                # 提取人脸
                face = frame_rgb[y1:y2, x1:x2]
                if face.size == 0:
                    continue
                
                # 预处理
                face_pil = Image.fromarray(face)
                face_tensor = transforms.Compose([
                    transforms.Resize((224, 224)),
                    transforms.ToTensor(),
                    transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
                ])(face_pil).unsqueeze(0).to(device)
                
                # 识别
                with torch.no_grad():
                    _, embedding = model(face_tensor)
                    name, confidence = database.recognize(embedding.cpu())
                
                # 更新识别历史
                recognition_history.append((name, confidence))
                if len(recognition_history) > max_history:
                    recognition_history.pop(0)
                
                # 计算最近识别结果的平均值
                if recognition_history:
                    names, confidences = zip(*recognition_history)
                    avg_confidence = sum(confidences) / len(confidences)
                    
                    # 选择出现次数最多的名字
                    from collections import Counter
                    most_common = Counter(names).most_common(1)
                    if most_common:
                        display_name = most_common[0][0]
                    else:
                        display_name = "Unknown"
                else:
                    display_name = "Unknown"
                    avg_confidence = 0.0
                
                # 绘制结果
                color = (0, 255, 0) if display_name != "Unknown" else (0, 0, 255)
                cv2.rectangle(frame, (x1, y1), (x2, y2), color, 2)
                cv2.putText(frame, f"{display_name} ({avg_confidence:.2f})", 
                           (x1, y1 - 10), cv2.FONT_HERSHEY_SIMPLEX, 0.9, color, 2)
        
        cv2.imshow('Real-time Face Recognition', frame)
        
        if cv2.waitKey(1) & 0xFF == ord('q'):
            break
    
    cap.release()
    cv2.destroyAllWindows()

# 主函数
def main():
    # 初始化模型
    num_classes = 10  # 初始类别数,可以动态增加
    model = FaceViT(num_classes=num_classes).to(device)
    
    # 加载预训练权重(如果有)
    model_path = "face_vit_model.pth"
    if os.path.exists(model_path):
        model.load_state_dict(torch.load(model_path, map_location=device))
        print("加载预训练模型")
    
    # 初始化人脸数据库
    database = FaceDatabase(threshold=0.7)
    db_path = "face_database.pkl"
    if os.path.exists(db_path):
        database.load(db_path)
        print(f"加载人脸数据库,包含 {len(database.labels)} 个人脸")
    
    # 主菜单
    while True:
        print("\n人脸识别系统")
        print("1. 注册新用户")
        print("2. 实时人脸识别")
        print("3. 保存数据库")
        print("4. 退出")
        
        choice = input("请选择操作: ")
        
        if choice == '1':
            name = input("请输入姓名: ")
            if name in database.labels:
                print(f"{name} 已存在,是否重新注册? (y/n)")
                if input().lower() != 'y':
                    continue
            
            success = register_face(model, database, name)
            if success:
                # 保存模型和数据库
                torch.save(model.state_dict(), model_path)
                database.save(db_path)
        
        elif choice == '2':
            realtime_recognition(model, database)
        
        elif choice == '3':
            torch.save(model.state_dict(), model_path)
            database.save(db_path)
            print("模型和数据库已保存")
        
        elif choice == '4':
            print("退出系统")
            break
        
        else:
            print("无效选择,请重新输入")

if __name__ == "__main__":
    main()

系统架构详解

1. Vision Transformer (ViT) 模型

class FaceViT(nn.Module):
    def __init__(self, num_classes, embedding_dim=128, num_heads=8, num_layers=6):
        # 图像分块处理
        self.patch_size = 16
        self.patch_dim = 3 * self.patch_size**2
        
        # 位置编码
        self.position_embedding = nn.Parameter(...)
        
        # Transformer编码器
        self.transformer_encoder = nn.TransformerEncoder(...)
        
        # 分类头
        self.classifier = nn.Sequential(...)
        
        # 嵌入层
        self.embedding_layer = nn.Linear(...)
        
        # 分类标记
        self.cls_token = nn.Parameter(...)
    
    def forward(self, x):
        # 图像分块
        # 嵌入转换
        # 添加分类标记和位置编码
        # Transformer编码
        # 分类和嵌入向量提取

2. 人脸数据库管理

class FaceDatabase:
    def __init__(self, threshold=0.7):
        self.database = defaultdict(list)
        self.threshold = threshold
        self.embeddings = []
        self.labels = []
    
    def add_face(self, embedding, label):
        # 添加新人脸
    
    def recognize(self, embedding):
        # 识别人脸
    
    def save(self, path):
        # 保存数据库
    
    def load(self, path):
        # 加载数据库

3. 主要功能模块

人脸注册
def register_face(model, database, name, capture_time=5):
    # 打开摄像头
    # 捕获多帧图像
    # 检测人脸
    # 提取人脸嵌入向量
    # 计算平均嵌入向量
    # 添加到数据库
实时识别
def realtime_recognition(model, database):
    # 打开摄像头
    # 实时检测人脸
    # 提取人脸嵌入向量
    # 与数据库比对
    # 显示识别结果

系统工作流程

  1. ​​初始化​​:

    • 加载预训练模型(如果存在)
    • 加载人脸数据库(如果存在)
  2. ​​人脸注册​​:

    • 用户输入姓名
    • 系统捕获5秒视频流
    • 检测并提取多个人脸样本
    • 计算平均嵌入向量
    • 添加到数据库
  3. ​​实时识别​​:

    • 打开摄像头
    • 实时检测人脸
    • 提取人脸嵌入向量
    • 与数据库比对
    • 显示识别结果和置信度
  4. ​​数据保存​​:

    • 保存模型权重
    • 保存人脸数据库

关键技术点

  1. ​​Vision Transformer (ViT)​​:

    • 将图像分割为16x16的小块
    • 使用Transformer编码器处理图像块序列
    • 提取高维特征表示
  2. ​​人脸检测​​:

    • 使用MTCNN进行实时人脸检测
    • 准确识别图像中的人脸位置
  3. ​​嵌入向量比对​​:

    • 使用余弦相似度计算人脸相似度
    • 设置阈值(0.7)判断是否匹配
  4. ​​平滑显示​​:

    • 存储最近识别结果
    • 使用多数投票和平均置信度平滑显示

使用说明

  1. ​​安装依赖​​:

    pip install opencv-python torch torchvision facenet-pytorch matplotlib pillow
  2. ​​运行系统​​:

    python face_recognition.py
  3. ​​操作指南​​:

    • 选择"1"注册新用户
    • 输入姓名后正对摄像头5秒
    • 选择"2"进行实时识别
    • 按"q"退出实时识别
    • 选择"3"保存数据库
    • 选择"4"退出系统

性能优化建议

  1. ​​模型量化​​:

    # 训练后量化
    quantized_model = torch.quantization.quantize_dynamic(
        model, {nn.Linear}, dtype=torch.qint8
    )
  2. ​​多线程处理​​:

    from threading import Thread
    
    # 创建处理线程
    processing_thread = Thread(target=process_frame, args=(frame,))
    processing_thread.start()
  3. ​​模型剪枝​​:

    # 模型剪枝示例
    for module in model.modules():
        if isinstance(module, nn.Linear):
            prune.l1_unstructured(module, name='weight', amount=0.2)
  4. ​​硬件加速​​:

    # 使用TensorRT加速
    import torch_tensorrt
    trt_model = torch_tensorrt.compile(model, 
        inputs=[torch_tensorrt.Input((1, 3, 224, 224))],
        enabled_precisions={torch.float32}
    )

Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐