共计 3251 个字符,预计需要花费 9 分钟才能阅读完成。
背景介绍
CelebA(CelebFaces Attributes Dataset)是计算机视觉领域广泛使用的人脸数据集,包含 202,599 张名人图像,每张图像标注了 40 种二元属性(如是否戴眼镜、是否微笑等)。这个数据集在人脸识别、人脸生成、属性分类等任务中扮演着重要角色,是许多论文和项目的基准数据集。

下载避坑指南
官方源与镜像源对比
CelebA 数据集最初由香港中文大学发布,官方下载地址通常速度较慢,特别是对于国内用户。推荐使用以下镜像源:
- 官方源:http://mmlab.ie.cuhk.edu.hk/projects/CelebA.html
- 百度云镜像(国内推荐):https://pan.baidu.com/s/1CRxxhoQ97A5qbsKO7iaAJg
- Google Drive 镜像:https://drive.google.com/drive/folders/0B7EVK8r0v71pWEZsZE9oNnFzTm8
使用 wget 进行断点续传
对于大文件下载,建议使用 wget 的 -c 参数实现断点续传:
wget -c http://mmlab.ie.cuhk.edu.hk/projects/CelebA/CelebA.zip
文件完整性校验
下载完成后,务必验证文件完整性。官方提供的 MD5 校验码:
# 计算文件的 MD5 值
md5sum CelebA.zip
# 应该得到的结果
d1e0d838f0d9a0d8b36a08a3e0f7b9d0 CelebA.zip
预处理实战
关键点对齐与裁剪
CelebA 提供了面部关键点坐标,我们可以使用 OpenCV 进行对齐和裁剪:
import cv2
import numpy as np
def align_face(img, landmarks):
"""
使用相似变换对齐人脸
:param img: 输入图像
:param landmarks: 5 个关键点坐标
:return: 对齐后的人脸图像
"""
# 标准化的目标关键点位置(根据 DeepFaceLab)dst_points = np.array([[30.2946, 51.6963],
[65.5318, 51.5014],
[48.0252, 71.7366],
[33.5493, 92.3655],
[62.7299, 92.2041]
], dtype=np.float32)
# 计算变换矩阵
transform = cv2.estimateAffinePartial2D(landmarks, dst_points)[0]
# 应用变换
aligned_face = cv2.warpAffine(img, transform, (96, 112), flags=cv2.INTER_LINEAR)
return aligned_face
属性标注解析
CelebA 的标注文件 list_attr_celeba.txt 格式需要特别注意:
def parse_attributes(file_path):
attributes = {}
with open(file_path, 'r') as f:
# 跳过前两行(文件头和数量声明)_ = f.readline()
_ = f.readline()
for line in f:
parts = line.split()
img_name = parts[0]
attrs = [1 if int(x) == 1 else 0 for x in parts[1:]]
attributes[img_name] = attrs
return attributes
构建 PyTorch Dataset
这是一个完整的 PyTorch Dataset 实现:
from torch.utils.data import Dataset
from PIL import Image
class CelebADataset(Dataset):
def __init__(self, img_dir, attr_path, transform=None):
self.img_dir = img_dir
self.transform = transform
self.attributes = self._load_attributes(attr_path)
self.image_names = list(self.attributes.keys())
def _load_attributes(self, attr_path):
# 实现同上 parse_attributes 函数
pass
def __len__(self):
return len(self.image_names)
def __getitem__(self, idx):
img_name = self.image_names[idx]
img_path = os.path.join(self.img_dir, img_name)
# 使用 PIL 打开图像
image = Image.open(img_path)
# 获取属性
attrs = self.attributes[img_name]
if self.transform:
image = self.transform(image)
return image, torch.tensor(attrs, dtype=torch.float32)
性能优化
多进程数据加载
PyTorch 的 DataLoader 原生支持多进程:
from torch.utils.data import DataLoader
dataset = CelebADataset(img_dir="img_align_celeba",
attr_path="list_attr_celeba.txt",
transform=transforms.ToTensor())
dataloader = DataLoader(dataset, batch_size=32,
shuffle=True, num_workers=4)
LMDB 格式转换
对于超大规模数据,LMDB 是更高效的选择:
import lmdb
import pickle
def convert_to_lmdb(img_dir, attr_path, output_path):
"""将 CelebA 转换为 LMDB 格式"""
env = lmdb.open(output_path, map_size=1099511627776)
attributes = parse_attributes(attr_path)
with env.begin(write=True) as txn:
for idx, img_name in enumerate(attributes.keys()):
img_path = os.path.join(img_dir, img_name)
with open(img_path, 'rb') as f:
img_data = f.read()
# 存储图像数据和属性
data = {
'image': img_data,
'attributes': attributes[img_name]
}
txn.put(str(idx).encode(), pickle.dumps(data))
常见问题排查
编码错误
CelebA 的标注文件可能包含特殊字符,建议指定编码:
with open("list_attr_celeba.txt", "r", encoding="utf-8", errors="ignore") as f:
# 处理文件内容
内存不足
处理大图像时,可以考虑分块处理:
# 分块处理图像
for i in range(0, len(image_names), chunk_size):
chunk = image_names[i:i+chunk_size]
process_images(chunk)
思考题
CelebA 数据集中存在性别不平衡问题(男性样本多于女性)。设计数据增强策略时,可以考虑:
- 对女性样本应用更多的增强变换(旋转、颜色抖动等)
- 使用过采样技术增加女性样本数量
- 在损失函数中引入类别权重
- 使用生成对抗网络(GAN)生成更多女性样本
你还能想到其他解决方法吗?欢迎在评论区分享你的想法。
完整的 Colab Notebook 实现可以参考:CelebA 预处理 Notebook
正文完
