运行timm示例代码时遇TypeError:__new__()缺少必填参数'task'
解决torchmetrics.Accuracy初始化时的TypeError: missing 'task'参数问题
Hey there, let's break down and fix this error you're hitting with the TimmMixupTrainer class!
问题原因
这个TypeError出现的核心原因是:torchmetrics在v1.0.0及以上的版本中,强制要求初始化Accuracy指标时必须传入task参数。你使用的gist代码是基于旧版torchmetrics编写的,当时可以直接用Accuracy()初始化,但新版torchmetrics需要明确指定任务类型(比如分类的具体类别形式)才能正常工作。
解决方案
你有两种简单可行的修复方式:
方案1:适配新版torchmetrics,修改Accuracy初始化代码
从你用timm模型做训练的场景来看,这应该是图像分类任务,所以需要给Accuracy指定task和num_classes参数。
找到TimmMixupTrainer类中初始化self.train_acc和self.val_acc的代码段,替换成如下内容:
# 请把YOUR_CLASS_COUNT替换成你本地数据集的实际类别数(比如10代表10分类任务) self.train_acc = Accuracy(task="multiclass", num_classes=YOUR_CLASS_COUNT) self.val_acc = Accuracy(task="multiclass", num_classes=YOUR_CLASS_COUNT)
- 如果是二分类任务,把
task改成"binary";如果是多标签分类任务,改成"multilabel"即可。
方案2:降级torchmetrics到兼容的旧版本
如果你不想修改代码,可以安装一个不需要task参数的旧版torchmetrics,在终端运行:
pip install torchmetrics==0.9.3
注意:这只是临时快速修复,长期来看更推荐方案1,因为旧版本可能存在bug或缺少新特性更新。
验证修改
完成上述任一修改后,重新运行你的训练脚本,这个TypeError应该就能被解决,训练器可以正常初始化精度指标了。
内容的提问来源于stack exchange,提问作者Walter
相关产品推荐
相关产品推荐

