如何在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()
使用流程
- 训练前调用
compute_real_stats得到真实数据的mu_real和sigma_real; - 每次迭代生成一批数据,计算其
mu_gen和sigma_gen; - 调用
calculate_fid得到当前迭代的FID分数。
内容的提问来源于stack exchange,提问作者Zahra Reyhanian
相关产品推荐
相关产品推荐

