使用Catalyst处理表格分类任务时数据类型冲突报错如何解决
问题解决方法
错误核心原因
你当前用的nn.CrossEntropyLoss要求:当模型输出为N类的原始logits时,传入的标签必须是torch.long类型的类别索引(二分类场景下就是取值为0/1的整数),你当前把标签转换成了float32类型,所以触发类型不匹配报错。
你之前把标签改成long类型仍然报错,是因为模型最后输出层加了ReLU激活,nn.CrossEntropyLoss需要输入未经过激活的原始logits,ReLU会截断负的输出值,引发数值问题甚至报错。
方案1:最小改动适配CrossEntropyLoss(推荐)
步骤1:修正标签数据类型
把创建标签tensor的代码替换为:
# 把布尔型target转成int64,对应pytorch的long类型 tt = torch.tensor(train['target'].values.astype('int64')) vt = torch.tensor(val['target'].values.astype('int64'))
步骤2:移除模型最后一层的ReLU激活
把模型的out层修改为:
self.out = nn.Linear(in_features= 5, out_features= 2)
不需要套Sequential加ReLU,CrossEntropyLoss内部会自动处理归一化逻辑。
步骤3:可选优化回调参数
二分类任务不需要计算top2准确率,你可以把AccuracyCallback的topk_args调整为(1,),避免无意义的计算:
dl.AccuracyCallback(input_key="logits", target_key="targets", topk_args=(1,)),
调整完以上内容即可正常训练。
方案2:切换为二分类专用损失函数
如果你不想修改标签类型,也可以换用二分类场景更常用的BCEWithLogitsLoss,改动如下:
步骤1:修改模型最后一层输出维度
self.out = nn.Linear(in_features= 5, out_features= 1)
步骤2:更换损失函数
criterion = nn.BCEWithLogitsLoss()
步骤3:调整回调适配单输出
callbacks=[ dl.AccuracyCallback(input_key="logits", target_key="targets", threshold=0.5), dl.PrecisionRecallF1SupportCallback( input_key="logits", target_key="targets", num_classes=2, threshold=0.5 ), ]
标签保持你原来的float32类型即可正常训练。
内容的提问来源于stack exchange,提问作者user9102437
相关产品推荐
相关产品推荐

