news 2026/9/16 17:26:28

Sana 项目 FID 评估实战:pytorch-fid 版本演进、统计量缓存与源码级用法解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Sana 项目 FID 评估实战:pytorch-fid 版本演进、统计量缓存与源码级用法解析

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.pytools/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.02020-08-16作为 PyPI 包初始发布,pip install pytorch-fid即可安装
0.1.12020-08-16修复setup.py中的软件许可证字符串
0.2.02020-11-30改用 PyTorch DataLoader 加载图像(加速);支持更多图片扩展名;引入 Nox、lint 与单测等工程化工具
0.2.12021-10-10新增--num-workers参数;修复 Windows 下打包问题
0.3.02023-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.npz

3.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读取musigma,从而跳过特征提取。

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 版本在pretrainedweights之间做兼容转换)。

四、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.pymain()(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_activationsfid_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 中明确了依赖:numpypillowscipytorch>=1.0.1torchvision>=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.pycompute_fid.py的参数定义,后者见 compute_fid.py):

参数默认值说明
--batch-size50批大小;若大于数据总量会自动缩到数据量
--num-workersmin(8, num_cpus)DataLoader 并行进程数
--device自动选择 cuda/cpu计算设备,如cuda:0
--dims2048Inception 特征维度,取值须在BLOCK_INDEX_BY_DIM
--save-statsFalse将统计量保存为.npz(第一个路径为输入,第二个为输出)
path必填(2 个)生成图像目录或.npz统计文件路径

6.2--dims与特征层选择

与官方实现不同,pytorch-fid 允许选择 Inception 网络的不同特征层,这在数据量不足 2048 张时很有用。特征维度与网络块的映射定义在 inception.py 的BLOCK_INDEX_BY_DIM

--dims取值特征来源备注
64第一次 max pooling 后特征需全局平均池化
192第二次 max pooling 后特征需全局平均池化
768aux 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.pysave_fid_statsif __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 30000

compute_fid_embedding.sh的逻辑正是「缓存优先」:若指定img_sizesample_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/5dMixed_6b-6eMixed_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_distancefid_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)则报错,否则取实部继续计算;
  • 进入函数前会断言两组musigma的形状一致,保证维度匹配。

calculate_activation_statistics(L206-L227)则通过np.mean(act, axis=0)np.cov(act, rowvar=False)完成高斯拟合。

八、使用注意事项

  1. 与官方 TensorFlow 实现的细微差异:尽管权重一致,但图像插值实现与库后端不同会导致结果略有出入(官方 README 报告在 LSUN 上绝对误差约 0.08、相对误差约 0.0009)。如需与论文中的历史 FID 值严格对齐,应使用官方 TensorFlow 实现;在 Sana 内部对比不同检查点时则无此顾虑。
  2. 维度一致性--dims一旦改变,得到的分数量纲即改变,不能与其它维度下的分数横向比较;同时不同数据集间比较 FID 需保证预处理(如--img_size)一致。
  3. 样本量:默认 2048 维特征要求参考集与生成集样本量足够;样本过少时可考虑降维特征(64/192/768),但需接受分数可比性与视觉相关性下降。
  4. 统计量缓存的复用条件.npz统计量只有在「同一数据集、同一--dims、同一预处理(--img_size、CenterCrop)与同一采样规模」下才能安全复用,改变任一条件都应重新生成。
  5. 权重文件准备:在 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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/16 17:25:57

企业内容管理转型:从PDF到结构化数据的实践指南

1. PDF在企业内容管理中的传统地位PDF格式自1993年由Adobe推出以来&#xff0c;已经成为企业文档交换的事实标准。它的跨平台一致性、固定布局特性和广泛的阅读器支持&#xff0c;使其在合同签署、技术文档发布等场景中长期占据主导地位。我曾参与过多个大型企业的文档管理系统…

作者头像 李华
网站建设 2026/9/16 17:25:46

抖音批量下载:从单个视频到整页作品,一个工具就能搞定

抖音批量下载&#xff1a;从单个视频到整页作品&#xff0c;一个工具就能搞定 【免费下载链接】douyin-downloader A practical Douyin downloader for both single-item and profile batch downloads, with progress display, retries, SQLite deduplication, and browser fal…

作者头像 李华
网站建设 2026/9/16 17:25:09

camofox-browser /snapshot端点深度解析:includeScreenshot与offset分页

camofox-browser /snapshot端点深度解析&#xff1a;includeScreenshot与offset分页 【免费下载链接】camofox-browser Stealth headless browser for AI agents — bypass Cloudflare, bot detection, and anti-scraping. Drop-in Puppeteer/Playwright replacement. 项目地…

作者头像 李华
网站建设 2026/9/16 17:24:24

Harbor企业级镜像仓库从零部署与排错实战

1. 为什么需要自己搭一个 Harbor&#xff1f;——从“Docker Hub 被限速”说起Harbor 不是 Docker 的替代品&#xff0c;而是 Docker 生态里真正能让你把镜像“管起来”的那把锁。我第一次在客户现场踩坑&#xff0c;就是因为没提前搭 Harbor&#xff1a;开发团队用docker push…

作者头像 李华