Torchvision v2变换API报错求助:参数缺失及类不存在问题
解决方案
核心原因
你使用的torchvision 0.15.x版本并不支持教程里的v2新API(ToDtype带scale参数、ToPureTensor、tv_tensors),官方文档"main"分支对应最新版本内容,和你本地的旧版本不匹配。
方案一:升级到兼容版本(推荐)
Windows CPU环境下,在conda环境中执行以下命令,自动安装PyTorch和torchvision的兼容稳定版本(会覆盖现有版本):
conda install pytorch torchvision torchaudio cpuonly -c pytorch
升级完成后,运行以下代码验证:
import torchvision print(torchvision.__version__) # 显示0.16.0及以上即为成功 from torchvision.transforms import v2 as T from torchvision import tv_tensors # 能正常导入就说明没问题
之后再运行教程里的代码就不会报错了。
方案二:改用旧版本兼容代码(不升级)
如果不想升级,把出错的get_transform函数替换成torchvision v1的写法,功能完全等价:
from torchvision import transforms as T def get_transform(train): transforms = [] if train: transforms.append(T.RandomHorizontalFlip(0.5)) transforms.append(T.ToTensor()) # 替代原来的ToDtype+ToPureTensor,自动完成转float和0-1缩放 return T.Compose(transforms)
内容的提问来源于stack exchange,提问作者blundered_bishop
相关产品推荐
相关产品推荐

