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

修改OpenPrompt代码加载本地阿拉伯语二分类数据集解决train报错

OpenPrompt加载阿拉伯语自定义二分类数据集修改方案

原有代码报错核心原因

原有代码是针对SuperGLUE的CB数据集编写的,存在3处硬编码问题,会直接导致训练环节报错:

  • 写死了公开数据集的本地存储路径,和个人云盘里的自定义数据集路径不匹配
  • 绑定了CB数据集专属字段名premise/hypothesis/idx,自定义阿拉伯语数据集不存在这些字段
  • 缺少阿拉伯语编码兼容、标签格式校验逻辑,非规范输入会在tokenizer或者训练计算loss阶段触发异常

分步修改操作

  1. 确认云盘数据集正确挂载到运行环境,验证路径可访问:普通csv/json格式的数据集无需提前转成HuggingFace Disk格式,直接用load_dataset读取对应格式即可;如果是通过save_to_disk存储的数据集,确认路径下存在dataset_dict.json等必要文件后再加载。
  2. 替换字段映射逻辑:将代码中取premise/hypothesis/idx的部分,替换为自定义数据集的对应列名——单句二分类任务(如阿语文本情感分类、违规内容识别)无需传text_b参数;无全局唯一id字段时,直接用遍历序号作为guid即可。
  3. 提前做标签校验:二分类标签必须转为0/1整数格式,禁止传入字符串类型标签值;阿拉伯语文本需确保以utf-8编码读取,避免乱码导致tokenizer解析失败。

修改后可直接运行的代码

!pip install openprompt
!git clone https://github.com/thunlp/OpenPrompt.git
%cd OpenPrompt

from datasets import load_dataset, load_from_disk
from openprompt.data_utils import InputExample

# -------------------------- 需自行修改的配置部分开始 --------------------------
# 数据集加载方式二选一:
# 方式1:加载云盘里的csv/json等普通格式数据集,以csv格式为例
# raw_dataset = load_dataset("csv", data_files={
#     "train": "/个人云盘挂载路径/train.csv",
#     "validation": "/个人云盘挂载路径/val.csv",
#     "test": "/个人云盘挂载路径/test.csv"
# }, encoding="utf-8")
# 方式2:加载之前通过save_to_disk存储的数据集
raw_dataset = load_from_disk("/个人云盘挂载路径/阿拉伯语二分类数据集存储路径")

# 替换为自定义数据集的实际列名,比如阿语单句分类的文本列名为"arabic_content",标签列名为"tag"
TEXT_COL_NAME = "arabic_content"
# 句子对任务填第二个文本列名,单句任务设为None即可
TEXT_B_COL_NAME = None
LABEL_COL_NAME = "tag"
# 如果标签是字符串格式(比如"positive"/"negative"),在这里映射成0/1整数
LABEL2ID = {"negative":0, "positive":1}
# -------------------------- 配置部分结束 --------------------------

dataset = {}
for split in ['train', 'validation', 'test']:
    dataset[split] = []
    for idx, data in enumerate(raw_dataset[split]):
        # 统一标签为整数格式
        raw_label = data[LABEL_COL_NAME]
        if isinstance(raw_label, str):
            label = LABEL2ID[raw_label]
        else:
            label = int(raw_label)
        
        # 构造输入样本
        input_example_params = {
            "text_a": data[TEXT_COL_NAME],
            "label": label,
            "guid": idx
        }
        if TEXT_B_COL_NAME is not None:
            input_example_params["text_b"] = data[TEXT_B_COL_NAME]
        
        input_example = InputExample(**input_example_params)
        dataset[split].append(input_example)

# 打印第一条样本验证加载结果
print(dataset['train'][0])

注意:如果训练阶段仍报编码相关错误,加载数据集时可显式指定encoding="utf-8-sig",适配部分Windows环境导出的阿拉伯语文本文件。

内容的提问来源于stack exchange,提问作者avery

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 22:18:42