同一向量计算Fréchet Inception Distance(FID)得分不为零的问题排查
Fréchet Inception Distance(FID)计算异常问题排查
问题描述
实现Fréchet Inception Distance(FID)得分计算时,使用同一特征向量act1计算得分,预期结果应为0,但实际得到了极大的异常数值。以下是代码及运行结果:
代码实现
# example of calculating the frechet inception distance import numpy from numpy import cov from numpy import trace from numpy import iscomplexobj from numpy.random import random from scipy.linalg import sqrtm # calculate frechet inception distance def calculate_fid(act1, act2): # calculate mean and covariance statistics mu1, sigma1 = act1.mean(axis=0), cov(act1, rowvar=False) mu2, sigma2 = act2.mean(axis=0), cov(act2, rowvar=False) # calculate sum squared difference between means ssdiff = numpy.sum((mu1 - mu2)**2.0) # calculate sqrt of product between cov covmean = sqrtm(sigma1.dot(sigma2)) # check and correct imaginary numbers from sqrt if iscomplexobj(covmean): covmean = covmean.real # calculate score fid = ssdiff + trace(sigma1 + sigma2 - 2.0 * covmean) return fid # define two collections of activations act1 = random(10*2048) act1 = act1.reshape((10,2048)) act2 = random(10*2048) act2 = act2.reshape((10,2048)) # fid between act1 and act1 fid = calculate_fid(act1, act1) print('FID (same): %.3f' % fid) # fid between act1 and act2 fid = calculate_fid(act1, act2) print('FID (different): %.3f' % fid)
运行结果
FID (same): -66113130760175032991744.000 FID (different): -55213970774324510299478046898216203619608871777363092441300193790394368.000
问题原因
核心问题是样本数量远小于特征维度:
- 特征维度为2048,但每个批次仅10个样本
- 当样本数小于特征维度时,协方差矩阵
sigma1和sigma2是奇异矩阵(秩不足),矩阵乘积的平方根sqrtm计算会出现数值不稳定,最终导致FID得分出现异常的极大/极小值
解决方案
有两种可行的解决方式:
1. 增加样本数量
确保样本数大于等于特征维度,比如将样本数调整为2048或更多:
# 修改样本生成部分 act1 = random(2048*2048) act1 = act1.reshape((2048,2048)) act2 = random(2048*2048) act2 = act2.reshape((2048,2048))
此时计算同一向量的FID得分会趋近于0。
2. 添加正则化项
对协方差矩阵添加极小的单位矩阵正则化,避免奇异矩阵问题:
def calculate_fid(act1, act2): mu1, sigma1 = act1.mean(axis=0), cov(act1, rowvar=False) mu2, sigma2 = act2.mean(axis=0), cov(act2, rowvar=False) # 添加正则化项 sigma1 += numpy.eye(sigma1.shape[0]) * 1e-6 sigma2 += numpy.eye(sigma2.shape[0]) * 1e-6 ssdiff = numpy.sum((mu1 - mu2)**2.0) covmean = sqrtm(sigma1.dot(sigma2)) if iscomplexobj(covmean): covmean = covmean.real fid = ssdiff + trace(sigma1 + sigma2 - 2.0 * covmean) return fid
修改后,即使样本数较少,也能得到合理的FID得分(同一向量的得分会接近0)。
内容的提问来源于stack exchange,提问作者Prince Patrick
相关产品推荐
相关产品推荐

