基于Transformer的人脸识别系统(实时摄像头实现)
·
基于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):
# 打开摄像头
# 实时检测人脸
# 提取人脸嵌入向量
# 与数据库比对
# 显示识别结果
系统工作流程
-
初始化:
- 加载预训练模型(如果存在)
- 加载人脸数据库(如果存在)
-
人脸注册:
- 用户输入姓名
- 系统捕获5秒视频流
- 检测并提取多个人脸样本
- 计算平均嵌入向量
- 添加到数据库
-
实时识别:
- 打开摄像头
- 实时检测人脸
- 提取人脸嵌入向量
- 与数据库比对
- 显示识别结果和置信度
-
数据保存:
- 保存模型权重
- 保存人脸数据库
关键技术点
-
Vision Transformer (ViT):
- 将图像分割为16x16的小块
- 使用Transformer编码器处理图像块序列
- 提取高维特征表示
-
人脸检测:
- 使用MTCNN进行实时人脸检测
- 准确识别图像中的人脸位置
-
嵌入向量比对:
- 使用余弦相似度计算人脸相似度
- 设置阈值(0.7)判断是否匹配
-
平滑显示:
- 存储最近识别结果
- 使用多数投票和平均置信度平滑显示
使用说明
-
安装依赖:
pip install opencv-python torch torchvision facenet-pytorch matplotlib pillow -
运行系统:
python face_recognition.py -
操作指南:
- 选择"1"注册新用户
- 输入姓名后正对摄像头5秒
- 选择"2"进行实时识别
- 按"q"退出实时识别
- 选择"3"保存数据库
- 选择"4"退出系统
性能优化建议
-
模型量化:
# 训练后量化 quantized_model = torch.quantization.quantize_dynamic( model, {nn.Linear}, dtype=torch.qint8 ) -
多线程处理:
from threading import Thread # 创建处理线程 processing_thread = Thread(target=process_frame, args=(frame,)) processing_thread.start() -
模型剪枝:
# 模型剪枝示例 for module in model.modules(): if isinstance(module, nn.Linear): prune.l1_unstructured(module, name='weight', amount=0.2) -
硬件加速:
# 使用TensorRT加速 import torch_tensorrt trt_model = torch_tensorrt.compile(model, inputs=[torch_tensorrt.Input((1, 3, 224, 224))], enabled_precisions={torch.float32} )
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐



所有评论(0)