加载本地训练的AraBERT分类模型遇RuntimeError问题求助
解决Transformer Pipeline加载本地模型的RuntimeError问题
问题根源
你传给pipeline的是已经实例化的BertForSequenceClassification对象,而pipeline默认会尝试通过字符串型的模型ID去自动推断任务类型,没法从已加载的模型实例里自动识别任务,因此触发报错。
两种解决方法
方法一:明确指定task参数
你的模型用于文本分类(序列分类),直接在创建pipeline时加上task='text-classification'即可:
from transformers import pipeline, AutoModelForSequenceClassification, AutoTokenizer model_name = 'aubmindlab/bert-base-arabertv02' arabert_model = AutoModelForSequenceClassification.from_pretrained('/gdrive/MyDrive/LabelModel') tokenizer = AutoTokenizer.from_pretrained(model_name) text = "أين وقعت غزوة بدر؟" #{'كيان': 0, 'تقريري': 1, 'حدث': 2, 'رقم': 3, 'عاقل': 4, 'موقع': 5, 'وصف': 6} # 明确指定任务类型 pipe = pipeline(task='text-classification', model=arabert_model, tokenizer=tokenizer) pipe(text)
方法二:直接用本地模型路径创建pipeline(更简洁)
无需手动加载模型和tokenizer,直接把本地模型路径传给pipeline,同时指定tokenizer的模型名,pipeline会自动完成加载流程:
from transformers import pipeline text = "أين وقعت غزوة بدر؟" #{'كيان': 0, 'تقريري': 1, 'حدث': 2, 'رقم': 3, 'عاقل': 4, 'موقع': 5, 'وصف': 6} pipe = pipeline( task='text-classification', model='/gdrive/MyDrive/LabelModel', tokenizer='aubmindlab/bert-base-arabertv02' ) pipe(text)
额外提示
如果希望pipeline返回结果直接显示你的自定义阿拉伯语标签(而非默认的LABEL_0这类),建议在训练完成保存模型时,将标签映射(id2label和label2id)写入模型目录的config.json文件中,示例格式:
{ "id2label": { "0": "كيان", "1": "تقريري", "2": "حدث", "3": "رقم", "4": "عاقل", "5": "موقع", "6": "وصف" }, "label2id": { "كيان": 0, "تقريري": 1, "حدث": 2, "رقم": 3, "عاقل": 4, "موقع": 5, "وصف": 6 } }
内容的提问来源于stack exchange,提问作者RJ94
相关产品推荐
相关产品推荐

