Torchmetrics FID自定义特征提取器时的异常问题求助
解决TorchMetrics FID自定义特征提取器的虚拟输入不匹配问题
问题根源
TorchMetrics的FrechetInceptionDistance默认会生成**[1, 3, 299, 299]的uint8格式虚拟图片**,这是为了适配原生InceptionV3的输入要求,但你的自定义特征提取器是针对MNIST(单通道、通常为float32格式)设计的,因此会出现类型和通道不匹配的报错。
两种解决方案
方案1:在特征提取器中适配输入格式
直接在自定义特征提取器的forward方法开头,将输入转换为符合MNIST要求的格式:
import torch as th import torch.nn as nn import torch.nn.functional as F from torchmetrics.image.fid import FrechetInceptionDistance N = 32 * 27 * 27 # 对应conv输出的flatten维度:(28-2+1)=27,32通道 class SimpleConvFeatureExtractor(nn.Module): def __init__(self, embed_dim): super().__init__() self.conv = nn.Conv2d(in_channels=1, out_channels=32, kernel_size=2) self.out = nn.Sequential(nn.Linear(N, embed_dim)) def forward(self, x): # 适配默认dummy输入:将3通道转单通道,uint8转float32 if x.shape[1] == 3: x = x.mean(dim=1, keepdim=True) # 3通道转单通道(取均值) x = x.to(th.float32) # 原forward逻辑 x = F.silu(self.conv(x)) x = self.out(x.view(x.shape[0], -1)) return x # 初始化FID fid = FrechetInceptionDistance(feature=SimpleConvFeatureExtractor(128))
方案2:重写FID类的虚拟输入生成逻辑
如果不想修改特征提取器,可以子类化FrechetInceptionDistance,重写_dummy_input方法,生成符合MNIST规格的虚拟输入:
import torch as th import torch.nn as nn import torch.nn.functional as F from torchmetrics.image.fid import FrechetInceptionDistance N = 32 * 27 * 27 class SimpleConvFeatureExtractor(nn.Module): def __init__(self, embed_dim): super().__init__() self.conv = nn.Conv2d(in_channels=1, out_channels=32, kernel_size=2) self.out = nn.Sequential(nn.Linear(N, embed_dim)) def forward(self, x): x = F.silu(self.conv(x)) x = self.out(x.view(x.shape[0], -1)) return x # 自定义FID类,重写虚拟输入 class MNISTFrechetInceptionDistance(FrechetInceptionDistance): def _dummy_input(self): # 生成MNIST规格的虚拟输入:[batch_size=1, channel=1, height=28, width=28],float32格式 return th.randn(1, 1, 28, 28, dtype=th.float32) # 初始化自定义FID fid = MNISTFrechetInceptionDistance(feature=SimpleConvFeatureExtractor(128))
说明
- 方案1更灵活,即使后续输入有其他格式变化,特征提取器也能自动适配;
- 方案2更彻底,直接从根源修改虚拟输入的规格,完全匹配你的模型预期。
内容的提问来源于stack exchange,提问作者smartstix
相关产品推荐
相关产品推荐

