You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

FastAI v2中如何使用已微调文本多分类模型对test.csv执行预测(替代v1的DatasetType.Test)

FastAI v2 测试集批量预测正确方法

我来帮你解决FastAI版本兼容的问题,同时搞定你循环预测导致内核崩溃的麻烦:

为什么原v1代码报错?

FastAI v2已经移除了DatasetType枚举类,所以直接调用DatasetType.Test会触发NameError。v2改用显式创建**测试数据加载器(Test DataLoader)**的方式指定预测数据集,这比v1的设计更灵活直观。

正确的批量预测步骤

你之前逐行循环预测的方式会重复创建单个样本的DataLoader,既浪费内存又效率极低,这就是内核崩溃的根源。FastAI v2推荐用批量预测,步骤如下:

  1. 创建测试数据加载器
    假设你已经有训练时用的DataLoaders对象(比如命名为dls),用它生成测试集的DataLoader,同时指定ordered=True保证预测结果和你的test_df顺序完全一致:

    test_dl = dls.test_dl(test_df, ordered=True)
    
  2. 批量获取预测结果
    把创建好的test_dl传给learn.get_preds()的dl参数,因为测试集没有标签,第二个返回值可以用下划线_忽略:

    preds, _ = learn.get_preds(dl=test_dl)
    
  3. 转换为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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.30 21:22:46