跳转至

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 训练中的重要设计 -- 防止模型利用漏洞获取高奖励。