一、FID 完整原理
1. 特征提取原理
2. 分布建模原理
| 均值向量
μ \\boldsymbol{\\mu} μ |
全部样本逐维算术平均 | 特征分布的中心位置 |
| 协方差矩阵
Σ \\Sigma Σ |
偏差向量外积求和,归一化 | 特征散布范围、维度间线性相关程度 |
- 真实分布:
N
(
μ
r
,
Σ
r
)
\\mathcal{N}(\\boldsymbol{\\mu}_r, \\Sigma_r)
N(μr,Σr) - 生成分布:
N
(
μ
g
,
Σ
g
)
\\mathcal{N}(\\boldsymbol{\\mu}_g, \\Sigma_g)
N(μg,Σg)
3. FID 数学计算原理
FID 是两个多维高斯分布的 Wasserstein-2 距离平方,存在解析闭式解:
F
I
D
=
∥
μ
r
−
μ
g
∥
2
2
⏟
中心位置误差
+
Tr
(
Σ
r
+
Σ
g
−
2
Σ
r
Σ
g
)
⏟
分布形状误差
FID = \\underbrace{\\|\\boldsymbol{\\mu}_r – \\boldsymbol{\\mu}_g\\|_2^2}_{\\text{中心位置误差}} + \\underbrace{\\text{Tr}\\left(\\Sigma_r + \\Sigma_g – 2\\sqrt{\\Sigma_r \\Sigma_g}\\right)}_{\\text{分布形状误差}}
FID=中心位置误差
∥μr−μg∥22+分布形状误差
Tr(Σr+Σg−2ΣrΣg
)
第一项:中心位置误差
- 运算:两个均值向量逐维做差,差值平方求和(L2 范数平方)
- 数学含义:量化两个高斯分布中心点的欧式距离,表征全局位置偏移
- 语义解释:InceptionV3 特征编码了物体类别、光照、色彩、构图等高层语义,中心偏移 = 两批图整体语义均值存在偏差
第二项:分布形状误差
由内向外运算:
| 1 |
Σ r Σ g \\Sigma_r \\Sigma_g ΣrΣg |
协方差矩阵标准乘法 |
| 2 |
Σ r Σ g \\sqrt{\\Sigma_r \\Sigma_g} ΣrΣg |
矩阵平方根(满足
M 2 = Σ r Σ g M^2 = \\Sigma_r \\Sigma_g M2=ΣrΣg 的正定矩阵 M M M) |
| 3 |
2 Σ r Σ g 2\\sqrt{\\Sigma_r \\Sigma_g} 2ΣrΣg |
平方根矩阵整体乘 2 |
| 4 |
Σ r + Σ g \\Sigma_r + \\Sigma_g Σr+Σg |
对位相加,两组特征总散布量 |
| 5 | 减法 | 得到仅保留差异的方阵 |
| 6 |
Tr ( ⋅ ) \\text{Tr}(\\cdot) Tr(⋅) 迹运算 |
主对角线求和,2048×2048 → 单个标量 |
- 数学含义:量化两个协方差矩阵的整体差异,表征散布形态、离散程度差异
- 语义解释:协方差描述样本波动范围与维度关联,该项对应两批图多样性和纹理分布的统计偏差
4. 数值判定原理
| 0 | 完美一致 |
| < 10 | 优秀,接近真实 |
| 10~50 | 可用 |
| > 50 | 质量较差 |
5. 工程计算原理
二、完整可运行代码
依赖安装
pip install torch torchvision pillow numpy scipy
完整实现
import os
import argparse
import numpy as np
from PIL import Image
import torch
import torchvision.models as models
import torchvision.transforms as transforms
from scipy.linalg import sqrtm
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# ============================================================
# 1. 构建 InceptionV3 特征提取器
# ============================================================
def build_inception_extractor():
"""加载预训练 InceptionV3,截取 Mixed_7c 层的 2048 维特征"""
# 兼容新版 torchvision 权重加载(pretrained=True 已废弃)
weights = models.Inception_V3_Weights.IMAGENET1K_V1
inception = models.inception_v3(weights=weights).to(device)
# 关闭辅助分类分支,消除冗余计算
inception.aux_logits = False
inception.eval()
feat_output = None
def hook_fn(module, input, output):
nonlocal feat_output
feat_output = output
handle = inception.Mixed_7c.register_forward_hook(hook_fn)
def get_feature():
return feat_output
return inception, handle, get_feature
# ============================================================
# 2. 图像预处理
# ============================================================
transform = transforms.Compose([
transforms.Resize((299, 299)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]),
])
# ============================================================
# 3. 批量提取 2048 维特征
# ============================================================
def extract_all_features(img_dir, extractor, get_feat, batch_size=32):
"""批量推理提取特征,自动跳过损坏图片"""
feat_list = []
valid_suffix = ('.jpg', '.jpeg', '.png')
img_names = [f for f in os.listdir(img_dir)
if f.lower().endswith(valid_suffix)]
print(f"检测到 {len(img_names)} 张图片")
for start_idx in range(0, len(img_names), batch_size):
batch_imgs = []
end_idx = min(start_idx + batch_size, len(img_names))
for name in img_names[start_idx:end_idx]:
img_path = os.path.join(img_dir, name)
try:
img = Image.open(img_path).convert('RGB')
batch_imgs.append(transform(img))
except Exception as e:
print(f"跳过损坏图片 {name}: {e}")
continue
if not batch_imgs:
continue
img_batch = torch.stack(batch_imgs).to(device)
with torch.no_grad():
extractor(img_batch)
feat_map = get_feat() # [B, 2048, 8, 8]
feat_vec = torch.mean(feat_map, dim=[2, 3]).cpu().numpy() # [B, 2048]
feat_list.extend(feat_vec)
print(f"已处理 {end_idx}/{len(img_names)}")
return np.array(feat_list) # [N, 2048]
# ============================================================
# 4. 计算均值和协方差(有偏估计,与 pytorch-fid 对齐)
# ============================================================
def calc_mu_sigma(features):
"""从特征矩阵 [N, 2048] 计算高斯分布的均值和协方差"""
mu = np.mean(features, axis=0) # [2048]
sigma = np.cov(features, rowvar=False, bias=True) # [2048, 2048]
return mu, sigma
# ============================================================
# 5. FID 核心公式
# ============================================================
def calculate_fid(mu1, sigma1, mu2, sigma2):
"""Wasserstein-2 距离平方 = 中心偏移 + 形状差异"""
# 第一项:均值 L2 范数平方
diff_mu = mu1 – mu2
term1 = np.sum(diff_mu ** 2)
# 第二项:协方差迹部分
cov_product = sigma1 @ sigma2
cov_sqrt = sqrtm(cov_product)
if np.iscomplexobj(cov_sqrt): # 消除浮点虚数误差
cov_sqrt = cov_sqrt.real
term2 = np.trace(sigma1 + sigma2 – 2 * cov_sqrt)
return float(term1 + term2)
# ============================================================
# 主程序(支持命令行参数)
# ============================================================
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="FID 图像质量评估")
parser.add_argument("–real", type=str, required=True, help="真实图片文件夹路径")
parser.add_argument("–gen", type=str, required=True, help="生成图片文件夹路径")
parser.add_argument("–batch", type=int, default=32, help="推理批量大小")
args = parser.parse_args()
# 初始化特征提取器
model, hook_handle, get_feature = build_inception_extractor()
# 提取特征
print("===== 提取真实图片特征 =====")
feats_real = extract_all_features(args.real, model, get_feature, args.batch)
print(f"有效真实样本: {feats_real.shape[0]} 张")
print("===== 提取生成图片特征 =====")
feats_gen = extract_all_features(args.gen, model, get_feature, args.batch)
print(f"有效生成样本: {feats_gen.shape[0]} 张")
# 计算高斯参数
mu_real, sigma_real = calc_mu_sigma(feats_real)
mu_gen, sigma_gen = calc_mu_sigma(feats_gen)
# 计算 FID
fid_score = calculate_fid(mu_real, sigma_real, mu_gen, sigma_gen)
print(f"\\n===== FID = {fid_score:.4f} =====")
hook_handle.remove()
torch.cuda.empty_cache()
三、使用说明
pip install torch torchvision pillow numpy scipy
# 命令行传参
python fid_calc.py –real ./real_images –gen ./gen_images –batch 32
注意事项
| 样本量 | 建议 ≥ 50000 张(论文标准),< 10000 偏置严重,不具备对比价值 |
| 协方差计算 | bias=True 有偏估计,与 pytorch-fid 官方工具结果一致 |
| 损坏图片 | 自动跳过,不会中断程序 |
| 批量推理 | –batch 参数控制 batch size,GPU 环境建议 32~64 |
| 设备 | 自动适配 CPU / GPU |
| 图片格式 | 支持 .jpg / .jpeg / .png |
个人能力有限,有问题随时联系~
网硕互联帮助中心



评论前必须登录!
注册