尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

使用 pytorch-fid 评估 Sana 文生图模型的图像质量:FID 原理、命令行用法与仓库内一键评测实践

发布时间:2026/9/16 13:04:50

资讯中心
01
ARTICLE

使用 pytorch-fid 评估 Sana 文生图模型的图像质量:FID 原理、命令行用法与仓库内一键评测实践

使用 pytorch-fid 评估 Sana 文生图模型的图像质量:FID 原理、命令行用法与仓库内一键评测实践
使用 pytorch-fid 评估 Sana 文生图模型的图像质量FID 原理、命令行用法与仓库内一键评测实践【免费下载链接】SanaSANA: Efficient High-Resolution Image Synthesis with Linear Diffusion Transformer项目地址: https://gitcode.com/GitHub_Trending/sana/SanaFIDFréchet Inception Distance是衡量两组图像分布相似度的经典指标常用于评估生成模型尤其是 GAN 与扩散模型输出样本的视觉质量。本文以 Sana 仓库内置的 pytorch-fid 工具包 为主体系统讲解 FID 的数学原理、pip/本地安装方式、CLI 用法、特征层选择技巧并结合仓库中的 compute_fid.py、compute_fid_embedding.sh 与指标评测管线介绍如何在 Sana 的 MJHQ-30K 评测流程中一键计算并上报 FID。读完本文你将掌握从两张图片文件夹到模型迭代曲线的完整 FID 评测能力。FID 是什么两个高斯分布之间的 Fréchet 距离FID 由 Martin Heusel 等人于 2017 年提出论文《GANs Trained by a Two Time-Scale Update Rule Converge to a Local Nash Equilibrium》最初用于量化 GAN 生成样本与真实图像在特征空间的分布差异。其核心思路是用预训练的 Inception 网络分别提取真实图像集与生成图像集的特征假设两组特征各自服从多元高斯分布分别估计均值向量 μ 与协方差矩阵 Σ计算两个高斯分布之间的 Fréchet 距离即 FID 分数。FID 与人类对图像质量的主观判断具有较好的相关性因此被广泛用作图像生成模型的评估指标。更详细的独立评估可参考论文《Are GANs Created Equal? A Large-Scale Study》。Sana 仓库将 pytorch-fid 以完整源码形式内置于 tools/metrics/pytorch-fid/其包结构如下src/pytorch_fid/fid_score.pyFID 计算主流程特征提取、统计量估计、距离计算src/pytorch_fid/inception.py为 FID 定制的 InceptionV3 模型与 TensorFlow 官方实现权重、结构完全一致src/pytorch_fid/main.pypython -m pytorch_fid入口setup.py打包配置同时注册pytorch-fid命令行脚本。核心算法calculate_frechet_distance在 fid_score.py 中Fréchet 距离按以下公式计算d^2 ||mu1 - mu2||^2 Tr(C1 C2 - 2 * sqrt(C1 * C2))实现上有两个值得注意的数值稳定性处理当协方差矩阵乘积接近奇异linalg.sqrtm结果非有限时会向两个协方差矩阵的对角线加上eps 1e-6再重新计算平方根由于数值误差sqrtm结果可能带有极小的虚部实现会检查虚部幅度阈值为1e-3随后丢弃虚部只保留实部避免最终距离出现复数。统计量估计则由calculate_activation_statistics完成对提取到的特征矩阵按行求均值mu np.mean(act, axis0)并计算协方差sigma np.cov(act, rowvarFalse)见 fid_score.py。安装方式一从 PyPI 安装按官方 README可直接安装发布版pip install pytorch-fid依赖项包括python3、pytorch、torchvision、pillow、numpy、scipy。其中 torch/torchvision 需满足torch1.0.1、torchvision0.2.2见 setup.py。方式二在 Sana 仓库内本地安装仓库已包含完整源码可直接以可编辑模式安装便于与 Sana 其他评测脚本共享环境pip install -e tools/metrics/pytorch-fid安装后src下的pytorch_fid包会被识别package_dir{: src}同时注册pytorch-fid pytorch_fid.fid_score:main控制台脚本见 setup.py。基础用法计算两组图像之间的 FID将两组图像分别放入两个文件夹即可通过命令行计算它们之间的 FIDpython -m pytorch_fid path/to/dataset1 path/to/dataset2若要在 GPU 上运行使用--device cuda:N指定 GPU 编号python -m pytorch_fid --device cuda:0 path/to/dataset1 path/to/dataset2在 fid_score.py 中定义的主要参数包括参数默认值说明path必填两个路径图像目录或.npz统计文件路径--batch-size50批大小建议按硬件显存调整--num-workersmin(8, num_cpus)数据加载进程数--device自动选择 cuda/cpu如cuda、cuda:0、cpu--dims2048Inception 特征维度详见下节--save-statsFalse将第一个路径的统计量保存为.npz数据加载方面工具支持bmp/jpg/jpeg/pgm/png/ppm/tif/tiff/webp共 9 种图像格式读取时统一convert(RGB)并做 ToTensor 归一化见 fid_score.py。若batch_size大于图片总数会打印警告并自动将批大小调整为数据量。选择不同的特征层--dims与官方 TensorFlow 实现不同pytorch-fid 允许通过--dims N选择 Inception 网络的不同特征层。当选择低层特征仍带空间维度时工具会先通过全局平均池化adaptive_avg_pool2d把特征压成向量再估计均值与协方差见 fid_score.py。可选维度如下映射关系定义在 inception.py 的BLOCK_INDEX_BY_DIM--dims对应网络块说明64Block 0第一次 max pooling 后的特征192Block 1第二次 max pooling 后的特征768Block 2送入 aux classifier 之前的特征2048Block 3最终平均池化特征默认使用低维特征的一个典型场景是当待比较的数据集不足 2048 张图片时可改用低维特征以获得更稳定的统计估计。但需要注意不同维度下的 FID 数值量级不同不能跨维度比较低层特征得到的分数可能不再与视觉质量强相关如果要在论文中与其他已发表分数对比务必使用默认的 2048 维pool3层特征。与 TensorFlow 官方实现保持一致的关键设计为确保分数可比inception.py中做了两个重要工作权重完全一致使用与官方 TensorFlow 实现相同的预训练权重pt_inception-2015-12-05-6726825d.pth。在 Sana 仓库中权重从本地路径output/pretrained_models/pt_inception-2015-12-05-6726825d.pth加载见 inception.py运行时请确保该权重文件已就位结构补丁通过FIDInceptionA/C/E_1/E_2四个子类复刻 TensorFlow 版 Inception 的差异结构包括平均池化不计入 padding 零值count_include_padFalse以及Mixed_7c分支使用 max pooling 而非 average pooling 等细节见 inception.py。因此README 也明确指出由于图像插值实现与库后端存在差异pytorch-fid 与官方 TensorFlow 实现的 FID 结果仍有细微差别在 LSUN 上使用 ProGAN 生成图像测试绝对误差约 0.08、相对误差约 0.0009。若论文需要与既有文献中的 FID 分数严格可比建议改用官方 TensorFlow 实现。生成可复用的.npz统计档案实际评测中常见需求是多个模型与同一个参考数据集比较。为避免每次都对参考数据集重复提取特征可以使用--save-stats把参考数据的统计量一次性保存为.npz档案python -m pytorch_fid --save-stats path/to/dataset path/to/outputfile生成的.npz文件中包含mu与sigma两个数组np.savez_compressed压缩存储见 fid_score.py。此后该档案可直接替代原始数据集路径参与比较python -m pytorch_fid path/to/generated_dataset path/to/outputfile.npz在 compute_statistics_of_path 中路径以.npz结尾时直接读取mu/sigma否则按目录扫描图像重新计算——这就是档案复用的底层机制。在 Sana 仓库中的实战MJHQ-30K 一键 FID 评测Sana 仓库将 pytorch-fid 深度集成进了指标评测工具链核心入口是 tools/metrics/compute_fid.py。与上游版本相比它扩展了以下能力支持 JSON 格式的样本清单路径可以是.npz统计文件、图像目录也可以是meta_data.json如 MJHQ-30K 的标注文件按sample_nums截取样本并拼接实际图片路径见 compute_fid.py可配置输入分辨率--img_size控制 Resize CenterCrop 的尺寸默认 512Sana 常用 256/512/1024使参考档案与生成图对齐在线日志上报--report_to wandb--tracker_pattern epoch_step可将 FID 按训练步数绘制成曲线见 tools/metrics/utils.py 的tracker函数结果缓存每个实验的 FID 会写入{exp_name}_sample{n}.txt重复评测时直接读取已有结果见 compute_fid.py批处理通过--stat先保存参考集档案、再批量计算多组生成样本避免重复计算。一键脚本compute_fid_embedding.shtools/metrics/compute_fid_embedding.sh 封装了完整的评测流程若参考档案MJHQ_30K_{img_size}px_fid_embeddings_{sample_nums}.npz不存在则先从 MJHQ-30K 参考图像生成它默认--img_size 512、--sample_nums 30000对每个实验目录单个目录或 model_paths.txt 中的列表计算 FID支持最多 8 个 GPU 并行处理多个实验按 GPU 轮询分配、每 8 个任务wait一次计算完成后可自动上报 wandb--log_fid true默认项目名t2i-evit-baseline。从模型权重到 FID 曲线的完整评测链Sana 提供了更上层的评测入口 scripts/bash_run_inference_metric.sh调用链为bash_run_inference_metric.sh └─ infer_metric_run_inference_metric.sh ├─ scripts/inference.py用配置与权重生成图像到 output/job/vis/ └─ tools/metrics/compute_fid_embedding.sh └─ tools/metrics/pytorch-fid/compute_fid.py典型用法单个权重文件或一组权重清单bash scripts/bash_run_inference_metric.sh \ configs/sana_config/1024ms/Sana_1600M_img1024.yaml \ output/Sana_1600M_1024px/checkpoints/Sana_1600M_1024px.pth其中可调的 FID 相关参数包括--img_size256/512/1024、--sample_nums1000/2500/5000/10000/30000、--fid_suffix_labelwandb 曲线后缀默认30K_bs50_Flow_DPM20、--tracker_pattern与--tracker_project_name等。评测数据准备与目录结构约定可参考 docs/metrics_toolkit.mdMJHQ-30K 图像解压到data/test/PG-eval-data/MJHQ-30K/imgs/下的 10 个类别子目录animals、art、fashion、food、indoor、landscape、logo、people、plants、vehicles评测产物统一存放在output/job_name/metrics/下。引用与许可如果你在研究中使用本仓库的 pytorch-fid 实现建议引用BibTeXmisc{Seitzer2020FID, author{Maximilian Seitzer}, title{{pytorch-fid: FID Score for PyTorch}}, month{August}, year{2020}, note{Version 0.3.0}, }许可说明pytorch-fid 实现本身采用 Apache License 2.0FID 指标由 Heusel、Ramsauer、Unterthiner、Nessler 与 Hochreiter 在《GANs Trained by a Two Time-Scale Update Rule Converge to a Local Nash Equilibrium》中提出官方 TensorFlow 实现由奥地利林茨大学生物信息学研究所JKU Linz发布同样采用 Apache License 2.0。在 Sana 仓库中相关评测脚本 scripts/inference.py、tools/metrics/compute_clipscore.sh 等共同构成了 FID、CLIP-Score、GenEval、DPG-Bench、ImageReward 的完整指标工具链FID 是其中评估生成分布与真实分布距离的核心一环。【免费下载链接】SanaSANA: Efficient High-Resolution Image Synthesis with Linear Diffusion Transformer项目地址: https://gitcode.com/GitHub_Trending/sana/Sana创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。