Sana 项目 FID 评估实战:pytorch-fid 版本演进、统计量缓存与源码级用法解析
【免费下载链接】SanaSANA: Efficient High-Resolution Image Synthesis with Linear Diffusion Transformer项目地址: https://gitcode.com/GitHub_Trending/sana/Sana
本篇技术指南以本仓库tools/metrics/pytorch-fid/下引入的 pytorch-fid(Fréchet Inception Distance,FID)实现及其CHANGELOG.md为骨架,系统梳理该指标工具从 0.1.0 到 0.3.0 的版本演进脉络(尤其是--save-stats统计量缓存、--num-workers并行加载等关键能力),并结合 Sana 仓库内的定制脚本(如tools/metrics/pytorch-fid/compute_fid.py、tools/metrics/compute_fid_embedding.sh)讲解如何在文生图(T2I)模型评测中落地使用。读完本文,你将掌握 FID 的底层计算原理、完整命令行用法,以及如何在 Sana 的推理评测流水线中复用缓存统计量、批量评估多个检查点并上报结果。
一、FID 指标与 pytorch-fid:从论文到开源实现
FID(Fréchet Inception Distance)是衡量两个图像数据集之间相似度的指标,由 Martin Heusel 等人在论文 "GANs Trained by a Two-Time Scale Update Rule Converge to a Local Nash Equilibrium" 中提出。其核心思想是:用 Inception 网络提取图像特征,将特征分布近似为两个多元高斯分布,再计算这两个高斯分布之间的 Fréchet 距离。实验表明 FID 与人类对图像质量的判断相关性良好,是评估生成模型(尤其是 GAN,以及扩散模型)样本质量的事实标准之一。
pytorch-fid 是官方 TensorFlow 实现(bioinf-jku/TTUR)到 PyTorch 的移植版本,权重与模型结构和官方实现完全一致。在本仓库中,它被集成到 Sana 的文生图评估工具链中:docs/metrics_toolkit.md明确将 FID 列为支持的评测指标之一(与 CLIP-Score、GenEval、DPG-Bench、ImageReward 并列),评估数据使用 MJHQ-30K 数据集。其对应的版本变更记录即本篇文章的主干文档 tools/metrics/pytorch-fid/CHANGELOG.md。
二、CHANGELOG 全览:版本演进时间线
pytorch-fid 自 2020 年发布以来经历了五个版本,CHANGELOG.md记录了每个版本的 Added / Fixed 明细,演进主线是「加载性能优化 → 并行能力 → 统计量缓存」:
| 版本 | 发布日期 | 核心变更 |
|---|---|---|
| 0.1.0 | 2020-08-16 | 作为 PyPI 包初始发布,pip install pytorch-fid即可安装 |
| 0.1.1 | 2020-08-16 | 修复setup.py中的软件许可证字符串 |
| 0.2.0 | 2020-11-30 | 改用 PyTorch DataLoader 加载图像(加速);支持更多图片扩展名;引入 Nox、lint 与单测等工程化工具 |
| 0.2.1 | 2021-10-10 | 新增--num-workers参数;修复 Windows 下打包问题 |
| 0.3.0 | 2023-01-05 | 新增--save-stats统计量缓存;修复 Windows CPU 探测与 torchvision 0.13 兼容问题 |
从变更记录可以看出,每个版本都围绕「让大规模 FID 计算更快、更稳、更省算力」展开,其中--save-stats是 0.3.0 引入的最具实战价值的能力。
三、0.3.0:--save-stats统计量缓存
3.1 功能背景
在文生图模型评测中,一个典型场景是「多个模型 / 多个检查点反复与同一参考数据集比较」。如果每次都重新抽取参考数据集的 Inception 特征并计算均值与协方差,会造成大量重复计算。0.3.0 引入的--save-stats参数允许一次性计算数据集统计量并保存为.npz文件,后续 FID 计算直接加载该文件,无需重新计算参考集统计量。
3.2 命令用法
python -m pytorch_fid --save-stats path/to/dataset path/to/outputfile第一个路径为输入数据集目录,第二个路径为输出的.npz文件。生成的.npz文件可以在后续 FID 计算中替代原始数据集路径,例如:
python -m pytorch_fid path/to/generated_images path/to/outputfile.npz3.3 源码实现印证
在上游实现 tools/metrics/pytorch-fid/src/pytorch_fid/fid_score.py 中:
save_fid_stats(L259-L275)会先校验输入路径存在、输出文件不存在,然后构建 Inception 模型、计算统计量,最后通过np.savez_compressed(paths[1], mu=m1, sigma=s1)将均值mu与协方差sigma压缩保存;compute_statistics_of_path(L230-L239)在路径以.npz结尾时直接np.load读取mu与sigma,从而跳过特征提取。
3.4 0.3.0 的两项修复
- Windows CPU 探测修复:不再使用
os.sched_getaffinity获取可用 CPU 数量(该 API 在 Windows 上不可用),改为在其抛AttributeError时回退到os.cpu_count()(见fid_score.pyL286-L295)。 - torchvision 0.13 兼容修复:不再使用 Inception 模型的
pretrained参数(该参数在 torchvision 0.13 中已被弃用),改为通过weights参数指定权重(见tools/metrics/pytorch-fid/src/pytorch_fid/inception.py中的_inception_v3包装函数,它会根据 torchvision 版本在pretrained与weights之间做兼容转换)。
四、0.2.x:数据加载与并行化改进
4.1 0.2.1:--num-workers参数
0.2.1 新增--num-workers,用于指定 DataLoader 的进程数。其默认值为 8,若可用 CPU 数少于 8 则取可用 CPU 数,即min(8, num_cpus)。这一逻辑在fid_score.py的main()(L286-L297)中有完整实现,Sana 仓库的定制版 tools/metrics/pytorch-fid/compute_fid.py 也保留了一致逻辑(L217-L225)。
4.2 0.2.0:DataLoader 加载与扩展名支持
0.2.0 的核心改动是用 PyTorch DataLoader 加载图像(替代原先一次性读入全部图像数组的方式),在大规模数据集上显著提速,并通过多 worker 并行实现吞吐提升。同时扩展了支持的图片格式集合。在源码中,支持格式定义在IMAGE_EXTENSIONS:
IMAGE_EXTENSIONS = {"bmp", "jpg", "jpeg", "pgm", "png", "ppm", "tif", "tiff", "webp"}get_activations(fid_score.pyL98-L150)将图片路径列表包装为ImagePathDataset并送入DataLoader(batch_size=batch_size, shuffle=False, drop_last=False, num_workers=num_workers),逐 batch 前向得到特征后写入预分配的pred_arr数组。
此外 0.2.0 还引入了 Nox 工具链、lint 与单元测试支持;0.2.1 修复了 Windows 下的包配置问题,并在 setup.py 中明确了依赖:numpy、pillow、scipy、torch>=1.0.1、torchvision>=0.2.2。
五、0.1.x:初始发布
0.1.0 作为 PyPI 包发布,安装方式为pip install pytorch-fid。0.1.1 仅修复了setup.py中的许可证字符串问题(软件许可证为 Apache License 2.0)。这一阶段奠定了「计算两个目录图像的 FID」这一最基础用法的基础:
python -m pytorch_fid path/to/dataset1 path/to/dataset2模块入口由src/pytorch_fid/__main__.py提供,它直接调用pytorch_fid.fid_score.main()。
六、在 Sana 仓库中的实际集成与定制
6.1 基础用法与常用参数
直接对比两个图像目录的 FID(上游标准用法):
python -m pytorch_fid path/to/dataset1 path/to/dataset2 # 指定 GPU 运行 python -m pytorch_fid --device cuda:0 path/to/dataset1 path/to/dataset2核心参数如下(源自fid_score.py与compute_fid.py的参数定义,后者见 compute_fid.py):
| 参数 | 默认值 | 说明 |
|---|---|---|
--batch-size | 50 | 批大小;若大于数据总量会自动缩到数据量 |
--num-workers | min(8, num_cpus) | DataLoader 并行进程数 |
--device | 自动选择 cuda/cpu | 计算设备,如cuda:0 |
--dims | 2048 | Inception 特征维度,取值须在BLOCK_INDEX_BY_DIM中 |
--save-stats | False | 将统计量保存为.npz(第一个路径为输入,第二个为输出) |
path | 必填(2 个) | 生成图像目录或.npz统计文件路径 |
6.2--dims与特征层选择
与官方实现不同,pytorch-fid 允许选择 Inception 网络的不同特征层,这在数据量不足 2048 张时很有用。特征维度与网络块的映射定义在 inception.py 的BLOCK_INDEX_BY_DIM:
--dims取值 | 特征来源 | 备注 |
|---|---|---|
| 64 | 第一次 max pooling 后特征 | 需全局平均池化 |
| 192 | 第二次 max pooling 后特征 | 需全局平均池化 |
| 768 | aux classifier 前特征 | 需全局平均池化 |
| 2048 | 最终平均池化特征(pool3) | 默认值 |
注意:改变维度会改变 FID 的数值量纲,不同维度下的分数不可相互比较,且低维特征分数可能与视觉质量的关联性变弱。当输出特征图仍有空间尺寸时,代码会通过adaptive_avg_pool2d先做全局平均池化再估计均值与协方差(见fid_score.pyL139-L144)。
6.3 Sana 定制版compute_fid.py的扩展能力
Sana 仓库将上游脚本扩展为更适合文生图评测的形态,主要扩展点包括:
- 三种输入类型:
.npz(直接加载统计量)、.json(按 meta 文件解析图像路径,支持 MJHQ-30K 的分目录结构)、普通图片目录(按IMAGE_EXTENSIONS收集文件); - 评测规模控制:
--sample_nums(默认 30000)限定采样数量,--img_size(默认 512)控制 Resize + CenterCrop 的预处理尺寸,与 MJHQ-30K 的 512/1024 分辨率评测对齐; - 批量评测:
--exp_name(默认Sana)标记实验名,--txt_path用于把 FID 结果写入<exp_name>_sample<sample_nums>.txt,避免重复计算同一实验; - 结果上报:
--log_fid、--report_to(tensorboard / wandb / comet_ml)、--tracker_pattern、--suffix_label等参数配合 tools/metrics/utils.py 中的tracker()函数,将 FID 结果以「step 为横轴、FID 为纵轴」的曲线形式记录到 wandb; - 统计量专用模式:
--stat与--save-stats配合,仅计算并保存参考集统计量,不执行 FID 计算(见compute_fid.py中save_fid_stats与if __name__ == "__main__"分支 L304-L327)。
一个典型的两段式用法(参考 tools/metrics/compute_fid_embedding.sh):
# 第一步:为参考集(MJHQ-30K)保存 FID 嵌入统计量 CUDA_VISIBLE_DEVICES=0 python tools/metrics/pytorch-fid/compute_fid.py \ --img_size 256 --path data/test/PG-eval-data/MJHQ-30K/meta_data.json \ --img_path data/test/PG-eval-data/MJHQ-30K/imgs \ --stat --sample_nums 30000 \ data/test/PG-eval-data/MJHQ-30K/MJHQ_30K_256px_fid_embeddings_30000.npz # 第二步:用已缓存的 npz 与生成图像目录计算 FID python tools/metrics/pytorch-fid/compute_fid.py --img_size 256 \ --path MJHQ_30K_256px_fid_embeddings_30000.npz data/test/PG-eval-data/MJHQ-30K/meta_data.json \ --exp_name your_exp --txt_path output/your_job --img_path output/your_job/vis \ --sample_nums 30000compute_fid_embedding.sh的逻辑正是「缓存优先」:若指定img_size与sample_nums对应的.npz参考统计量不存在,则先保存;随后按单实验或 txt 文件批量启动 FID 计算(最多并行 8 个 GPU 任务),最后统一将结果上报到 wandb。--exp_name支持.txt文件,每行一个实验目录,配合asset/model_paths.txt可批量评估一批检查点。
6.4 与整体评测流水线的衔接
在 Sana 中,FID 评测通常不是孤立运行的,而是「推理 + 评测」一体:scripts/bash_run_inference_metric.sh接收配置文件与模型路径列表,先调用推理脚本生成图像,再计算 FID / CLIP-Score,最后上传指标到 wandb。其关键默认参数包括:--img_size默认 512、--sample_nums默认 30000、采样算法默认flow_dpm-solver、--step默认 20、参考集 meta 文件为data/test/PG-eval-data/MJHQ-30K/meta_data.json。评测结果按 docs/metrics_toolkit.md 约定的目录树组织在output/your_job_name/下(checkpoints/、vis/、metrics/等子目录)。
七、底层原理:Inception 特征与 Fréchet 距离计算
7.1 FID 专用的 InceptionV3
pytorch-fid 使用的 Inception 模型与 torchvision 自带版本结构略有不同(如Mixed_5b/5c/5d、Mixed_6b-6e、Mixed_7b/7c被替换为 FID 专用的FIDInceptionA/C/E块,其中FIDInceptionE_2使用 max pooling 而非 average pooling),因此必须加载官方 FID 权重才能得到可比较的分数。这一点在 inception.py 的fid_inception_v3()(L189-L215)中体现。
需要注意的是:上游实现从 URL 下载权重(FID_WEIGHTS_URL),而 Sana 仓库内的版本将下载逻辑注释掉,改为从本地路径加载:
inception.load_state_dict( torch.load("output/pretrained_models/pt_inception-2015-12-05-6726825d.pth", map_location="cpu") )因此,在 Sana 仓库中运行 FID 评测前,需要先将该 Inception 权重文件放到output/pretrained_models/目录下,否则模型加载会失败。
7.2 Fréchet 距离公式与数值稳定性
calculate_frechet_distance(fid_score.pyL153-L203)实现了 Fréchet 距离的数值稳定版本,其数学形式为:
d^2 = ||mu1 - mu2||^2 + Tr(C1 + C2 - 2 * sqrt(C1 * C2))即两个高斯分布(由特征均值mu与协方差sigma刻画)之间的 Wasserstein-2 距离。实现中包含以下数值处理细节:
- 使用
scipy.linalg.sqrtm计算协方差乘积的矩阵平方根; - 当乘积近似奇异(
covmean出现非有限值)时,向两个协方差矩阵的对角线加上eps = 1e-6再重算; - 当矩阵平方根因数值误差出现微小虚部时,若虚部超过阈值(对角元素虚部与 0 的误差大于
1e-3)则报错,否则取实部继续计算; - 进入函数前会断言两组
mu、sigma的形状一致,保证维度匹配。
calculate_activation_statistics(L206-L227)则通过np.mean(act, axis=0)与np.cov(act, rowvar=False)完成高斯拟合。
八、使用注意事项
- 与官方 TensorFlow 实现的细微差异:尽管权重一致,但图像插值实现与库后端不同会导致结果略有出入(官方 README 报告在 LSUN 上绝对误差约 0.08、相对误差约 0.0009)。如需与论文中的历史 FID 值严格对齐,应使用官方 TensorFlow 实现;在 Sana 内部对比不同检查点时则无此顾虑。
- 维度一致性:
--dims一旦改变,得到的分数量纲即改变,不能与其它维度下的分数横向比较;同时不同数据集间比较 FID 需保证预处理(如--img_size)一致。 - 样本量:默认 2048 维特征要求参考集与生成集样本量足够;样本过少时可考虑降维特征(64/192/768),但需接受分数可比性与视觉相关性下降。
- 统计量缓存的复用条件:
.npz统计量只有在「同一数据集、同一--dims、同一预处理(--img_size、CenterCrop)与同一采样规模」下才能安全复用,改变任一条件都应重新生成。 - 权重文件准备:在 Sana 仓库中运行 FID 前,请确认
output/pretrained_models/pt_inception-2015-12-05-6726825d.pth已就位。
九、引用与许可
若在研究中使用了 pytorch-fid,可按其 README 提供的 BibTeX 条目引用(作者 Maximilian Seitzer,版本 0.3.0)。该实现与原始 JKU Linz 实现均遵循 Apache License 2.0;FID 指标原始出处为 Heusel 等人 2017 年的 GAN 论文。
十、进一步阅读
- FID 上游实现与标准用法说明:安装、基础用法、
--dims详解 - FID 版本变更记录:各版本 Added / Fixed 明细
- Sana 定制 FID 脚本:支持 json / npz / 目录三种输入与 wandb 上报
- FID 嵌入保存与批量评测脚本:MJHQ-30K 统计量缓存与多 GPU 并行评测
- 评测工具链总览:FID / CLIP-Score / GenEval / DPG-Bench / ImageReward 的一体化评测说明
- 推理评测流水线脚本:推理 + 评测 + 日志上报的完整入口
【免费下载链接】SanaSANA: Efficient High-Resolution Image Synthesis with Linear Diffusion Transformer项目地址: https://gitcode.com/GitHub_Trending/sana/Sana
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考