云计算百科
云计算领域专业知识百科平台

图像评估方法 FID 的具体实现原理与代码实现

一、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μg22+分布形状误差

    Tr(Σr+Σg2Σ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. 数值判定原理

  • FID 为非负实数,两组分布完全一致时两项均为 0,FID = 0
  • FID 单调递增对应两组图像特征分布差异扩大
  • FID含义
    0 完美一致
    < 10 优秀,接近真实
    10~50 可用
    > 50 质量较差

    5. 工程计算原理

  • 批量读取图像,统一缩放至 299×299(InceptionV3 标准输入)
  • 批量前向推理提取 2048 维特征,缓存全部特征样本
  • 基于缓存特征分别求解真实集、生成集的均值向量与协方差矩阵
  • 代入 Wasserstein-2 距离闭式公式,输出标量 FID

  • 二、完整可运行代码

    依赖安装

    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

    个人能力有限,有问题随时联系~

    赞(0)
    未经允许不得转载:网硕互联帮助中心 » 图像评估方法 FID 的具体实现原理与代码实现
    分享到: 更多 (0)

    评论 抢沙发

    评论前必须登录!