使用RoBERTa-base QA模型返回上下文而非答案的问题排查
解决roberta-base-squad2模型返回上下文而非目标技能的问题
问题根源
代码出现返回上下文/问题+上下文的情况,核心原因包括:
- 重复初始化模型:每次调用
generate_skills都重新加载tokenizer和模型,既低效又可能引发状态异常 - 未处理无答案场景:SQuAD2.0模型支持识别无答案情况,但代码直接取
argmax,当模型判定上下文无匹配答案时,会选中整个上下文的起止区间 - 问题与上下文匹配度不足:如果
top_words是零散关键词而非结构化技能描述,模型无法定位明确答案区间,只能返回完整上下文 - 答案提取逻辑粗糙:直接取概率最高的起止位置,未过滤覆盖整个上下文的无效长答案
修正后的代码
import pandas as pd import torch from transformers import AutoTokenizer, AutoModelForQuestionAnswering # 全局加载模型和tokenizer,避免重复初始化 tokenizer = AutoTokenizer.from_pretrained("deepset/roberta-base-squad2") model = AutoModelForQuestionAnswering.from_pretrained("deepset/roberta-base-squad2") # 设备配置,优先使用GPU加速 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model.to(device) def generate_skills(question, context): # 处理空上下文或缺失值 if not context or pd.isna(context): return "" inputs = tokenizer(question, context, return_tensors='pt', truncation=True, max_length=512).to(device) with torch.no_grad(): # 关闭梯度计算,节省资源 outputs = model(**inputs) start_scores = outputs.start_logits[0] end_scores = outputs.end_logits[0] # 判断无答案场景:SQuAD2.0中<s> token分数最高表示无匹配答案 cls_score = start_scores[0] max_start_score = torch.max(start_scores[1:]) if cls_score > max_start_score: return "" # 提取答案区间,过滤过长的无效答案 start_index = torch.argmax(start_scores) end_index = torch.argmax(end_scores) + 1 context_token_len = len(tokenizer(context, return_tensors='pt')['input_ids'][0]) if (end_index - start_index) > context_token_len * 0.8: return "" tokens = inputs['input_ids'][0][start_index:end_index] answer = tokenizer.decode(tokens, skip_special_tokens=True).strip() # 过滤与问题/上下文完全重复的无效答案 if answer == question or answer == context: return "" return answer def generate_skills_for_row(row): context = row['top_words'] # 调整问题表述,更贴合工作活动类上下文 question = "从以下工作活动中提取该岗位所需的必要技能:" return generate_skills(question, context) # 示例数据(替换为你的实际数据集) df = pd.DataFrame({'top_words': [ "负责数据清洗、特征工程,使用Python和SQL处理大规模数据集", "参与机器学习模型构建,熟练掌握Scikit-learn和TensorFlow", "协助业务部门进行数据可视化,使用Tableau制作报表" ]}) df['skills'] = df.apply(generate_skills_for_row, axis=1) print(df)
额外优化建议
- 预处理上下文:如果
top_words是零散关键词,先拼接成连贯句子(如用逗号连接),帮助模型更好理解内容 - 使用Pipeline简化代码:Hugging Face的
pipeline已封装无答案判断和后处理逻辑,代码更简洁:from transformers import pipeline qa_pipeline = pipeline("question-answering", model="deepset/roberta-base-squad2", device=device) def generate_skills(question, context): result = qa_pipeline(question=question, context=context) return result['answer'] if result['score'] > 0.3 else "" # 设置分数阈值过滤低置信度答案 - 调整问题表述:根据
top_words内容细化问题,比如针对数据科学家岗位可提问:"该数据科学家岗位需要哪些必要技能?" - 设置置信度阈值:通过判断
start_scores和end_scores的最大值,过滤低置信度答案,避免无效输出
内容的提问来源于stack exchange,提问作者Moe_blg
相关产品推荐
相关产品推荐

