search_r1_like_qa_em.py — 这个文件实现了 Search-R1 风格的 QA 精确匹配(Exact Match)评分¶
模块路径: verl.utils.reward_score.search_r1_like_qa_em
文件概述¶
这个文件实现了 Search-R1 风格的 QA 精确匹配(Exact Match)评分,适用于开放域问答数据集(如 NQ、TriviaQA、HotpotQA 等)。模型需要在 <answer>...</answer> 标签中给出答案。
改编自 Search-R1 项目。
关键代码讲解¶
1. 答案标准化 normalize_answer¶
def normalize_answer(s):
def remove_articles(text):
return re.sub(r"\b(a|an|the)\b", " ", text) # 去冠词
def white_space_fix(text):
return " ".join(text.split()) # 统一空格
def remove_punc(text):
return "".join(ch for ch in text if ch not in string.punctuation) # 去标点
def lower(text):
return text.lower() # 转小写
return white_space_fix(remove_articles(remove_punc(lower(s))))
标准化流程:小写 → 去标点 → 去冠词 → 统一空格。这样 "The United States" 和 "united states" 被视为相同。
2. 精确匹配检查¶
def em_check(prediction, golden_answers):
"""完全匹配:标准化后的预测必须与某个标准答案完全相同"""
normalized_prediction = normalize_answer(prediction)
for golden_answer in golden_answers:
if normalize_answer(golden_answer) == normalized_prediction:
return 1
return 0
def subem_check(prediction, golden_answers):
"""子串匹配:标准答案是预测的子串即可"""
normalized_prediction = normalize_answer(prediction)
for golden_answer in golden_answers:
if normalize_answer(golden_answer) in normalized_prediction:
return 1
return 0
两种匹配模式:EM(完全匹配)和 SubEM(子串匹配)。
3. 答案提取¶
def extract_solution(solution_str):
answer_pattern = r"<answer>(.*?)</answer>"
match = re.finditer(answer_pattern, solution_str, re.DOTALL)
matches = list(match)
if len(matches) < 1:
return None
return matches[-1].group(1).strip() # 取最后一个 <answer> 标签
4. 评分函数¶
def compute_score(solution_str, ground_truth, method="strict", format_score=0.0, score=1.0):
answer = extract_solution(solution_str)
open_count, close_count = count_answer_tags(solution_str)
if answer is None:
return 0
else:
if em_check(answer, ground_truth["target"]):
if open_count > 10 or close_count > 10: # 防止模型输出大量 </answer> 刷分
score = score / 4
return score
else:
return format_score
注意这里有一个防作弊机制:如果模型输出了超过10个 <answer> 标签,分数降为 1/4,防止模型学会通过大量重复来"猜"答案。
核心类/函数列表¶
| 函数名 | 作用 |
|---|---|
normalize_answer |
标准化答案(去冠词、标点、统一大小写) |
em_check |
精确匹配检查 |
subem_check |
子串匹配检查 |
extract_solution |
从 <answer> 标签中提取答案 |
compute_score |
主评分函数 |
compute_score_subem |
使用子串匹配的评分函数 |
与其他模块的关系¶
- 被
__init__.py调用,处理searchR1_nq、searchR1_triviaqa等数据集 - 改编自 Search-R1 项目的 QA 评分逻辑
小结¶
这个模块针对 QA 任务的评分特点设计:答案是自然语言文本(非数学表达式),需要做文本标准化后比较。防作弊机制是 RL 训练中的重要设计 -- 防止模型利用漏洞获取高奖励。