You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在MNIST数据集的GAN迭代中快速计算FID分数?

快速计算FID分数的最优方案

针对MNIST这类小数据集的GAN训练场景,要实现高效的每迭代FID计算,核心是简化特征提取流程、避免冗余计算、优化矩阵操作,具体方案如下:

一、替换重型特征提取模型

放弃默认的InceptionV3(针对RGB大图设计,且MNIST灰度图需额外格式转换,冗余度极高),改用预训练LeNet-5或自定义轻量CNN作为特征提取器:

  • LeNet-5专为MNIST这类手写数字任务设计,特征提取效率比InceptionV3高10倍以上;
  • 若需和标准FID对齐,可截取InceptionV3的前几层(如pool3输出),但远轻于完整模型。

二、预计算真实数据统计量

真实数据是固定的,只需要计算一次均值和协方差,将结果存储在GPU显存中,后续迭代仅需计算生成数据的特征统计量:

def compute_real_stats(dataloader, feature_extractor, device):
    feature_extractor.eval()
    features = []
    with torch.no_grad():
        for imgs, _ in dataloader:
            imgs = imgs.to(device)
            feat = feature_extractor(imgs).flatten(1)
            features.append(feat)
    features = torch.cat(features, dim=0)
    mu_real = torch.mean(features, dim=0)
    sigma_real = torch.cov(features.T)
    return mu_real, sigma_real

三、优化生成数据的特征计算

  • 每次迭代生成批量数据,直接批量输入特征提取器,避免单样本计算的开销;
  • 确保生成器、特征提取器全程在GPU上运行,减少CPU-GPU数据传输的耗时;
  • 特征提取器保持eval模式,关闭dropout、batchnorm的训练逻辑,避免额外计算和随机波动。

四、加速协方差矩阵运算

FID的核心是计算Fréchet距离,可通过以下技巧优化:

  • 用torch.linalg.cholesky分解替代直接计算协方差逆矩阵,降低计算复杂度;
  • 给协方差矩阵添加极小的单位矩阵(1e-6 * torch.eye),避免数值不稳定;
  • 全程使用float32计算,MNIST场景下不需要更高精度,大幅减少计算时间。

五、极简版FID计算实现

自己实现核心逻辑比调用通用库(如torcheval)更高效,以下是轻量化的FID计算函数:

def calculate_fid(mu1, sigma1, mu2, sigma2, eps=1e-6):
    mu1 = mu1.flatten()
    mu2 = mu2.flatten()
    # 避免协方差矩阵奇异
    sigma1 = sigma1 + eps * torch.eye(sigma1.shape[0], device=sigma1.device)
    sigma2 = sigma2 + eps * torch.eye(sigma2.shape[0], device=sigma2.device)
    
    # 计算协方差矩阵的平方根
    covmean, _ = torch.linalg.sqrtm(sigma1 @ sigma2, out=None)
    if not torch.isfinite(covmean).all():
        offset = torch.eye(sigma1.shape[0], device=sigma1.device) * eps
        covmean = torch.linalg.sqrtm((sigma1 + offset) @ (sigma2 + offset))
    
    # 处理复数结果
    if covmean.is_complex():
        covmean = covmean.real
    
    # 计算Fréchet距离
    fid_score = diff @ diff + torch.trace(sigma1) + torch.trace(sigma2) - 2 * torch.trace(covmean)
    return fid_score.item()

使用流程

  1. 训练前调用compute_real_stats得到真实数据的mu_real和sigma_real;
  2. 每次迭代生成一批数据,计算其mu_gen和sigma_gen;
  3. 调用calculate_fid得到当前迭代的FID分数。

内容的提问来源于stack exchange,提问作者Zahra Reyhanian

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.03 16:32:22