使用自定义特征提取器计算FID报错:类型不匹配问题排查
自定义特征提取器计算FID报错修复
问题代码
import torch _ = torch.manual_seed(123) from torchmetrics.image.fid import FrechetInceptionDistance from torchvision.models import inception_v3 net = inception_v3() checkpoint = torch.load('checkpoint.pt') net.load_state_dict(checkpoint['state_dict']) net.eval() fid = FrechetInceptionDistance(feature=net) # 生成两个略有重叠的图像强度分布 imgs_dist1 = torch.randint(0, 200, (100, 3, 299, 299), dtype=torch.uint8) imgs_dist2 = torch.randint(100, 255, (100, 3, 299, 299), dtype=torch.uint8) fid.update(imgs_dist1, real=True) fid.update(imgs_dist2, real=False) result = fid.compute() print(result)
报错信息
Traceback (most recent call last): File "foo.py", line 12, in <module> fid = FrechetInceptionDistance(feature=net) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Lib/site-packages/torchmetrics/image/fid.py", line 304, in __init__ num_features = self.inception(dummy_image).shape[-1] ^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Lib/site-packages/torch/nn/modules/module.py", line 1518, in _wrapped_call_impl return self._call_impl(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Lib/site-packages/torch/nn/modules/module.py", line 1527, in _call_impl return forward_call(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Lib/site-packages/torchvision/models/inception.py", line 166, in forward x, aux = self._forward(x) ^^^^^^^^^^^^^^^^ File "/Lib/site-packages/torchvision/models/inception.py", line 105, in _forward x = self.Conv2d_1a_3x3(x) ^^^^^^^^^^^^^^^^^^^^^ File "/Lib/site-packages/torch/nn/modules/module.py", line 1518, in _wrapped_call_impl return self._call_impl(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Lib/site-packages/torch/nn/modules/module.py", line 1527, in _call_impl return forward_call(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Lib/site-packages/torchvision/models/inception.py", line 405, in forward x = self.conv(x) ^^^^^^^^^^^^ File "/Lib/site-packages/torch/nn/modules/module.py", line 1518, in _wrapped_call_impl return self._call_impl(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Lib/site-packages/torch/nn/modules/module.py", line 1527, in _call_impl return forward_call(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Lib/site-packages/torch/nn/modules/conv.py", line 460, in forward return self._conv_forward(input, self.weight, self.bias) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Lib/site-packages/torch/nn/modules/conv.py", line 456, in _conv_forward return F.conv2d(input, weight, bias, self.stride, ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ RuntimeError: expected scalar type Byte but found Float
问题原因
- 数据类型不兼容:
FrechetInceptionDistance初始化时会生成Float类型的测试图像,但加载后的模型权重为Byte类型,导致卷积层输入与权重类型不匹配。 - 特征提取错误:直接使用
inception_v3作为特征提取器时,模型默认返回分类logits,而非FID计算所需的中间层特征(如pool3输出)。
修复方案
1. 修正模型数据类型
加载 checkpoint 后,将模型转为 Float32 类型,确保与输入数据类型一致:
net = net.float()
2. 包装模型以提取正确特征
创建包装类,让inception_v3返回FID标准的中间层特征,并处理输入归一化:
import torch.nn as nn import torchvision.transforms.functional as TF class InceptionFeatureExtractor(nn.Module): def __init__(self, inception_model): super().__init__() # 截取到pool3层,提取特征而非logits self.feature_extractor = nn.Sequential(*list(inception_model.children())[:-1]) def forward(self, x): # 将uint8图像转为0-1的Float,并应用inception标准归一化 if x.dtype == torch.uint8: x = x.float() / 255.0 x = TF.normalize(x, mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) features = self.feature_extractor(x) return features.view(features.size(0), -1) # 展平为特征向量
3. 修改初始化与调用逻辑
关闭FID的默认归一化,使用自定义特征提取器:
import torch _ = torch.manual_seed(123) from torchmetrics.image.fid import FrechetInceptionDistance from torchvision.models import inception_v3 net = inception_v3() checkpoint = torch.load('checkpoint.pt') net.load_state_dict(checkpoint['state_dict']) net.eval() net = net.float() # 修正数据类型 # 创建自定义特征提取器 feature_extractor = InceptionFeatureExtractor(net) # 初始化FID,关闭默认归一化 fid = FrechetInceptionDistance(feature=feature_extractor, normalize=False) # 测试数据 imgs_dist1 = torch.randint(0, 200, (100, 3, 299, 299), dtype=torch.uint8) imgs_dist2 = torch.randint(100, 255, (100, 3, 299, 299), dtype=torch.uint8) fid.update(imgs_dist1, real=True) fid.update(imgs_dist2, real=False) result = fid.compute() print(result)
内容的提问来源于stack exchange,提问作者A.A
相关产品推荐
相关产品推荐

