FastAI v2中如何使用已微调文本多分类模型对test.csv执行预测(替代v1的DatasetType.Test)
FastAI v2 测试集批量预测正确方法
我来帮你解决FastAI版本兼容的问题,同时搞定你循环预测导致内核崩溃的麻烦:
为什么原v1代码报错?
FastAI v2已经移除了DatasetType枚举类,所以直接调用DatasetType.Test会触发NameError。v2改用显式创建**测试数据加载器(Test DataLoader)**的方式指定预测数据集,这比v1的设计更灵活直观。
正确的批量预测步骤
你之前逐行循环预测的方式会重复创建单个样本的DataLoader,既浪费内存又效率极低,这就是内核崩溃的根源。FastAI v2推荐用批量预测,步骤如下:
创建测试数据加载器
假设你已经有训练时用的DataLoaders对象(比如命名为dls),用它生成测试集的DataLoader,同时指定ordered=True保证预测结果和你的test_df顺序完全一致:test_dl = dls.test_dl(test_df, ordered=True)批量获取预测结果
把创建好的test_dl传给learn.get_preds()的dl参数,因为测试集没有标签,第二个返回值可以用下划线_忽略:preds, _ = learn.get_preds(dl=test_dl)转换为numpy数组(和v1代码输出一致)
最后把预测张量转成numpy数组,和你原来v1代码的labels结果完全相同:labels = preds.numpy()
完整示例代码
如果你的训练流程基于TextDataLoaders,完整的预测代码大概是这样:
# 假设你已完成模型训练,learn是训练好的模型,test_df是测试数据集 test_dl = dls.test_dl(test_df, ordered=True) preds, _ = learn.get_preds(dl=test_dl) labels = preds.numpy() # 可选:如果需要得到每个样本的类别标签(而非概率) predicted_classes = preds.argmax(dim=1).numpy()
补充说明
ordered=True非常关键,它确保预测结果的顺序和test_df的行顺序一一对应,避免出现样本与预测结果错位的问题。- 批量预测是FastAI的设计初衷,不仅效率高,还能避免内存溢出问题,完全不需要逐行遍历。
内容的提问来源于stack exchange,提问作者lihas sharma
相关产品推荐
相关产品推荐

