修改OpenPrompt代码加载本地阿拉伯语二分类数据集解决train报错
OpenPrompt加载阿拉伯语自定义二分类数据集修改方案
原有代码报错核心原因
原有代码是针对SuperGLUE的CB数据集编写的,存在3处硬编码问题,会直接导致训练环节报错:
- 写死了公开数据集的本地存储路径,和个人云盘里的自定义数据集路径不匹配
- 绑定了CB数据集专属字段名
premise/hypothesis/idx,自定义阿拉伯语数据集不存在这些字段 - 缺少阿拉伯语编码兼容、标签格式校验逻辑,非规范输入会在tokenizer或者训练计算loss阶段触发异常
分步修改操作
- 确认云盘数据集正确挂载到运行环境,验证路径可访问:普通csv/json格式的数据集无需提前转成HuggingFace Disk格式,直接用
load_dataset读取对应格式即可;如果是通过save_to_disk存储的数据集,确认路径下存在dataset_dict.json等必要文件后再加载。 - 替换字段映射逻辑:将代码中取
premise/hypothesis/idx的部分,替换为自定义数据集的对应列名——单句二分类任务(如阿语文本情感分类、违规内容识别)无需传text_b参数;无全局唯一id字段时,直接用遍历序号作为guid即可。 - 提前做标签校验:二分类标签必须转为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
相关产品推荐
相关产品推荐

