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

使用自定义特征提取器计算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

问题原因

  1. 数据类型不兼容:FrechetInceptionDistance初始化时会生成Float类型的测试图像,但加载后的模型权重为Byte类型,导致卷积层输入与权重类型不匹配。
  2. 特征提取错误:直接使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 16:07:04