一、FID 完整原理1. 特征提取原理使用 ImageNet 预训练的InceptionV3作为固定特征提取器移除网络末端分类层与 Softmax截取倒数第二层Mixed_7c特征图对特征图执行全局平均池化将任意输入图像映射为固定的2048 维实数特征向量分别对真实图和生成图批量提取各自构成独立特征样本集2. 分布建模原理海量图像经 InceptionV3 得到的 2048 维特征统计分布近似服从多维高斯分布多维高斯分布由两个参数唯一确定参数计算方式表征含义均值向量μ \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_gM2ΣrΣg的正定矩阵M MM32 Σ r Σ g 2\sqrt{\Sigma_r \Sigma_g}2ΣrΣg平方根矩阵整体乘 24Σ r Σ g \Sigma_r \Sigma_gΣrΣg对位相加两组特征总散布量5减法得到仅保留差异的方阵6Tr ( ⋅ ) \text{Tr}(\cdot)Tr(⋅)迹运算主对角线求和2048×2048 → 单个标量数学含义量化两个协方差矩阵的整体差异表征散布形态、离散程度差异语义解释协方差描述样本波动范围与维度关联该项对应两批图多样性和纹理分布的统计偏差4. 数值判定原理FID 为非负实数两组分布完全一致时两项均为 0FID 0FID 单调递增对应两组图像特征分布差异扩大FID含义0完美一致 10优秀接近真实10~50可用 50质量较差5. 工程计算原理批量读取图像统一缩放至 299×299InceptionV3 标准输入批量前向推理提取 2048 维特征缓存全部特征样本基于缓存特征分别求解真实集、生成集的均值向量与协方差矩阵代入 Wasserstein-2 距离闭式公式输出标量 FID二、完整可运行代码依赖安装pipinstalltorch torchvision pillow numpy scipy完整实现importosimportargparseimportnumpyasnpfromPILimportImageimporttorchimporttorchvision.modelsasmodelsimporttorchvision.transformsastransformsfromscipy.linalgimportsqrtm devicetorch.device(cudaiftorch.cuda.is_available()elsecpu)# # 1. 构建 InceptionV3 特征提取器# defbuild_inception_extractor():加载预训练 InceptionV3截取 Mixed_7c 层的 2048 维特征# 兼容新版 torchvision 权重加载pretrainedTrue 已废弃weightsmodels.Inception_V3_Weights.IMAGENET1K_V1 inceptionmodels.inception_v3(weightsweights).to(device)# 关闭辅助分类分支消除冗余计算inception.aux_logitsFalseinception.eval()feat_outputNonedefhook_fn(module,input,output):nonlocalfeat_output feat_outputoutput handleinception.Mixed_7c.register_forward_hook(hook_fn)defget_feature():returnfeat_outputreturninception,handle,get_feature# # 2. 图像预处理# transformtransforms.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 维特征# defextract_all_features(img_dir,extractor,get_feat,batch_size32):批量推理提取特征自动跳过损坏图片feat_list[]valid_suffix(.jpg,.jpeg,.png)img_names[fforfinos.listdir(img_dir)iff.lower().endswith(valid_suffix)]print(f检测到{len(img_names)}张图片)forstart_idxinrange(0,len(img_names),batch_size):batch_imgs[]end_idxmin(start_idxbatch_size,len(img_names))fornameinimg_names[start_idx:end_idx]:img_pathos.path.join(img_dir,name)try:imgImage.open(img_path).convert(RGB)batch_imgs.append(transform(img))exceptExceptionase:print(f跳过损坏图片{name}:{e})continueifnotbatch_imgs:continueimg_batchtorch.stack(batch_imgs).to(device)withtorch.no_grad():extractor(img_batch)feat_mapget_feat()# [B, 2048, 8, 8]feat_vectorch.mean(feat_map,dim[2,3]).cpu().numpy()# [B, 2048]feat_list.extend(feat_vec)print(f已处理{end_idx}/{len(img_names)})returnnp.array(feat_list)# [N, 2048]# # 4. 计算均值和协方差有偏估计与 pytorch-fid 对齐# defcalc_mu_sigma(features):从特征矩阵 [N, 2048] 计算高斯分布的均值和协方差munp.mean(features,axis0)# [2048]sigmanp.cov(features,rowvarFalse,biasTrue)# [2048, 2048]returnmu,sigma# # 5. FID 核心公式# defcalculate_fid(mu1,sigma1,mu2,sigma2):Wasserstein-2 距离平方 中心偏移 形状差异# 第一项均值 L2 范数平方diff_mumu1-mu2 term1np.sum(diff_mu**2)# 第二项协方差迹部分cov_productsigma1 sigma2 cov_sqrtsqrtm(cov_product)ifnp.iscomplexobj(cov_sqrt):# 消除浮点虚数误差cov_sqrtcov_sqrt.real term2np.trace(sigma1sigma2-2*cov_sqrt)returnfloat(term1term2)# # 主程序支持命令行参数# if__name____main__:parserargparse.ArgumentParser(descriptionFID 图像质量评估)parser.add_argument(--real,typestr,requiredTrue,help真实图片文件夹路径)parser.add_argument(--gen,typestr,requiredTrue,help生成图片文件夹路径)parser.add_argument(--batch,typeint,default32,help推理批量大小)argsparser.parse_args()# 初始化特征提取器model,hook_handle,get_featurebuild_inception_extractor()# 提取特征print( 提取真实图片特征 )feats_realextract_all_features(args.real,model,get_feature,args.batch)print(f有效真实样本:{feats_real.shape[0]}张)print( 提取生成图片特征 )feats_genextract_all_features(args.gen,model,get_feature,args.batch)print(f有效生成样本:{feats_gen.shape[0]}张)# 计算高斯参数mu_real,sigma_realcalc_mu_sigma(feats_real)mu_gen,sigma_gencalc_mu_sigma(feats_gen)# 计算 FIDfid_scorecalculate_fid(mu_real,sigma_real,mu_gen,sigma_gen)print(f\n FID {fid_score:.4f})hook_handle.remove()torch.cuda.empty_cache()三、使用说明pipinstalltorch torchvision pillow numpy scipy# 命令行传参python fid_calc.py--real./real_images--gen./gen_images--batch32注意事项注意点说明样本量建议 ≥ 50000 张论文标准 10000 偏置严重不具备对比价值协方差计算biasTrue有偏估计与pytorch-fid官方工具结果一致损坏图片自动跳过不会中断程序批量推理--batch参数控制 batch sizeGPU 环境建议 32~64设备自动适配 CPU / GPU图片格式支持.jpg/.jpeg/.png个人能力有限有问题随时联系~