如何解决RuntimeError:输入与权重张量类型不匹配问题(Fastai多分类)
解决Fastai多类别图像分类中Input与Weight张量类型不匹配的错误
错误原因
在Kaggle GPU环境中,dls.train.one_batch()返回的张量默认加载到GPU(torch.cuda.FloatTensor),但cnn_learner初始化的模型默认留在CPU(torch.FloatTensor),二者设备不匹配导致触发RuntimeError。
代码笔误修正
先修正原始代码中的两处基础错误:
- 导入包拼写错误:
from fastasi.vision.all import *→from fastai.vision.all import * untar_data调用缺少闭合括号:path = untar_data(URLs.PASCAL_2007→path = untar_data(URLs.PASCAL_2007)
解决方案
以下三种方式均可解决问题,优先推荐前两种以充分利用GPU加速:
方式一:将模型移至GPU
在调用learn.model(x)前执行:
learn.model = learn.model.cuda()
方式二:自动匹配设备(推荐)
利用Fastai内置的设备属性,将输入张量移至模型所在设备:
activs = learn.model(x.to(learn.dls.device))
learn.dls.device会自动检测当前可用设备(GPU/CPU),无需手动指定。
方式三:将输入移至CPU(不推荐)
仅临时测试时可用,会浪费GPU资源:
x = x.cpu() activs = learn.model(x)
修正后的完整代码
from fastai.vision.all import * import numpy as np import pandas as pd path = untar_data(URLs.PASCAL_2007) df = pd.read_csv(path/'train.csv') def get_x(r): return r['fname'] def get_y(r): return r['labels'] def splitter(df): train = df.index[~df['is_valid']].tolist() valid = df.index[df['is_valid']].tolist() return train, valid dblock = DataBlock(blocks = (ImageBlock, MultiCategoryBlock), splitter = splitter, get_x = get_x, get_y = get_y, item_tfms = RandomResizedCrop(128, min_scale = 0.35)) dls = dblock.dataloaders(df) learn = cnn_learner(dls, resnet18) x,y = dls.train.one_batch() # 使用推荐的自动匹配设备方式 activs = learn.model(x.to(learn.dls.device)) activs.shape
内容的提问来源于stack exchange,提问作者Ibrat Usmonov
相关产品推荐
相关产品推荐

