初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD Source: Original Platform
This commit is contained in:
83
slime/rollout/rm_hub/__init__.py
Normal file
83
slime/rollout/rm_hub/__init__.py
Normal file
@@ -0,0 +1,83 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import asyncio
|
||||
import random
|
||||
|
||||
import aiohttp
|
||||
|
||||
from slime.utils.misc import load_function
|
||||
from slime.utils.types import Sample
|
||||
|
||||
from .deepscaler import get_deepscaler_rule_based_reward
|
||||
from .f1 import f1_score
|
||||
from .gpqa import compute_gpqa_reward
|
||||
from .math_dapo_utils import compute_score as compute_score_dapo
|
||||
from .math_utils import extract_answer as extract_boxed_answer
|
||||
from .math_utils import grade_answer_verl
|
||||
|
||||
|
||||
async def remote_rm(args, sample: Sample):
|
||||
payload = {
|
||||
"prompt": sample.prompt,
|
||||
"response": sample.response,
|
||||
"label": sample.label,
|
||||
}
|
||||
session_kwargs = {}
|
||||
async with aiohttp.ClientSession(**session_kwargs) as session:
|
||||
async with session.post(args.rm_url, json=payload) as resp:
|
||||
resp.raise_for_status()
|
||||
return await resp.json()
|
||||
|
||||
|
||||
async def async_rm(args, sample: Sample, **kwargs):
|
||||
if args.custom_rm_path is not None:
|
||||
rm_function = load_function(args.custom_rm_path)
|
||||
return await rm_function(args, sample, **kwargs)
|
||||
|
||||
metadata = sample.metadata if isinstance(sample.metadata, dict) else {}
|
||||
rm_type = (metadata.get("rm_type") or args.rm_type or "").strip()
|
||||
response = sample.response
|
||||
label = sample.label
|
||||
if rm_type.startswith("boxed_"):
|
||||
response = extract_boxed_answer(response) or ""
|
||||
rm_type = rm_type[len("boxed_") :]
|
||||
|
||||
# This function is intended for remote or time-consuming reward model evaluation.
|
||||
# Implement the actual logic as needed.
|
||||
if rm_type == "remote_rm":
|
||||
return await remote_rm(args, sample)
|
||||
elif rm_type == "deepscaler":
|
||||
return get_deepscaler_rule_based_reward(response, label)
|
||||
elif rm_type == "dapo":
|
||||
return compute_score_dapo(response, label)
|
||||
elif rm_type == "math":
|
||||
return 1 if grade_answer_verl(response, label) else 0
|
||||
elif rm_type == "f1":
|
||||
return f1_score(response, label)[0]
|
||||
elif rm_type == "gpqa":
|
||||
return compute_gpqa_reward(response, label, metadata=metadata)
|
||||
elif rm_type == "ifbench":
|
||||
from .ifbench import compute_ifbench_reward
|
||||
|
||||
return compute_ifbench_reward(response, label, metadata=metadata)
|
||||
elif rm_type == "random":
|
||||
return random.randint(0, 1)
|
||||
elif rm_type:
|
||||
raise NotImplementedError(f"Rule-based RM for {rm_type} is not implemented.")
|
||||
else:
|
||||
raise NotImplementedError("Rule-based RM type is not specified.")
|
||||
|
||||
|
||||
async def batched_async_rm(
|
||||
args,
|
||||
samples: list[Sample],
|
||||
**kwargs,
|
||||
) -> list[int | float]:
|
||||
if args.custom_rm_path is not None:
|
||||
# Ensure the custom reward function is implemented in batch mode
|
||||
rm_function = load_function(args.custom_rm_path)
|
||||
return await rm_function(args, samples, **kwargs)
|
||||
tasks = [async_rm(args, sample, **kwargs) for sample in samples]
|
||||
rewards = await asyncio.gather(*tasks)
|
||||
return rewards
|
||||
BIN
slime/rollout/rm_hub/__pycache__/__init__.cpython-312.pyc
Normal file
BIN
slime/rollout/rm_hub/__pycache__/__init__.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/rollout/rm_hub/__pycache__/deepscaler.cpython-312.pyc
Normal file
BIN
slime/rollout/rm_hub/__pycache__/deepscaler.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/rollout/rm_hub/__pycache__/f1.cpython-312.pyc
Normal file
BIN
slime/rollout/rm_hub/__pycache__/f1.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/rollout/rm_hub/__pycache__/gpqa.cpython-312.pyc
Normal file
BIN
slime/rollout/rm_hub/__pycache__/gpqa.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/rollout/rm_hub/__pycache__/math_dapo_utils.cpython-312.pyc
Normal file
BIN
slime/rollout/rm_hub/__pycache__/math_dapo_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/rollout/rm_hub/__pycache__/math_utils.cpython-312.pyc
Normal file
BIN
slime/rollout/rm_hub/__pycache__/math_utils.cpython-312.pyc
Normal file
Binary file not shown.
45
slime/rollout/rm_hub/deepscaler.py
Normal file
45
slime/rollout/rm_hub/deepscaler.py
Normal file
@@ -0,0 +1,45 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from .math_utils import extract_answer, grade_answer_mathd, grade_answer_sympy
|
||||
|
||||
|
||||
def get_deepscaler_rule_based_reward(response, label):
|
||||
if "</think>" in response:
|
||||
model_solution = response.split("</think>")[-1]
|
||||
elif "###Response" in response:
|
||||
model_solution = response.split("###Response")[1]
|
||||
else:
|
||||
return 0
|
||||
|
||||
model_answer = extract_answer(model_solution)
|
||||
if model_answer is None:
|
||||
return 0
|
||||
if label == "":
|
||||
return 0
|
||||
|
||||
# Convert single answer to list for uniform processing
|
||||
assert isinstance(label, (str, float, int))
|
||||
ground_truths = [label]
|
||||
|
||||
# Process each ground truth
|
||||
processed_ground_truths = []
|
||||
for truth in ground_truths:
|
||||
truth = str(truth)
|
||||
if "\\boxed" in truth:
|
||||
processed_truth = extract_answer(truth)
|
||||
if processed_truth is not None:
|
||||
processed_ground_truths.append(processed_truth)
|
||||
else:
|
||||
processed_ground_truths.append(truth)
|
||||
|
||||
if not processed_ground_truths:
|
||||
return 0
|
||||
|
||||
# Check against all possible correct answers
|
||||
for ground_truth in processed_ground_truths:
|
||||
is_correct = grade_answer_mathd(model_answer, ground_truth) or grade_answer_sympy(model_answer, ground_truth)
|
||||
if is_correct:
|
||||
return 1
|
||||
|
||||
return 0
|
||||
50
slime/rollout/rm_hub/f1.py
Normal file
50
slime/rollout/rm_hub/f1.py
Normal file
@@ -0,0 +1,50 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import re
|
||||
import string
|
||||
from collections import Counter
|
||||
|
||||
|
||||
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):
|
||||
exclude = set(string.punctuation)
|
||||
return "".join(ch for ch in text if ch not in exclude)
|
||||
|
||||
def lower(text):
|
||||
return text.lower()
|
||||
|
||||
return white_space_fix(remove_articles(remove_punc(lower(s))))
|
||||
|
||||
|
||||
def f1_score(prediction, ground_truth):
|
||||
ZERO_METRIC = (0, 0, 0)
|
||||
|
||||
if prediction is None:
|
||||
return ZERO_METRIC
|
||||
|
||||
normalized_prediction = normalize_answer(prediction)
|
||||
normalized_ground_truth = normalize_answer(ground_truth)
|
||||
|
||||
if normalized_prediction in ["yes", "no", "noanswer"] and normalized_prediction != normalized_ground_truth:
|
||||
return ZERO_METRIC
|
||||
if normalized_ground_truth in ["yes", "no", "noanswer"] and normalized_prediction != normalized_ground_truth:
|
||||
return ZERO_METRIC
|
||||
|
||||
prediction_tokens = normalized_prediction.split()
|
||||
ground_truth_tokens = normalized_ground_truth.split()
|
||||
common = Counter(prediction_tokens) & Counter(ground_truth_tokens)
|
||||
num_same = sum(common.values())
|
||||
if num_same == 0:
|
||||
return ZERO_METRIC
|
||||
precision = 1.0 * num_same / len(prediction_tokens)
|
||||
recall = 1.0 * num_same / len(ground_truth_tokens)
|
||||
f1 = (2 * precision * recall) / (precision + recall)
|
||||
return f1, precision, recall
|
||||
132
slime/rollout/rm_hub/gpqa.py
Normal file
132
slime/rollout/rm_hub/gpqa.py
Normal file
@@ -0,0 +1,132 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import re
|
||||
import string
|
||||
from collections.abc import Iterable
|
||||
|
||||
DEFAULT_VALID_LETTERS = list(string.ascii_uppercase[:8])
|
||||
|
||||
|
||||
def _strip_chain_of_thought(text: str) -> str:
|
||||
if not text:
|
||||
return ""
|
||||
|
||||
if "</think>" in text:
|
||||
return text.rsplit("</think>", 1)[-1]
|
||||
|
||||
return text
|
||||
|
||||
|
||||
def _normalize_text(text: str) -> str:
|
||||
return re.sub(r"[^a-z0-9]+", " ", text.lower()).strip()
|
||||
|
||||
|
||||
def _extract_letter_from_response(response: str, valid_letters: Iterable[str]) -> str | None:
|
||||
"""
|
||||
Best-effort extraction of the selected option letter from the model response.
|
||||
"""
|
||||
if not response:
|
||||
return None
|
||||
|
||||
text = _strip_chain_of_thought(response)
|
||||
patterns = [
|
||||
r"(?:answer|option|choice)\s*(?:is|:)?\s*([A-Z])",
|
||||
r"([A-Z])\s*(?:is\s*(?:the)?\s*correct)",
|
||||
r"final\s*(?:answer|option)\s*(?:is|:)?\s*([A-Z])",
|
||||
]
|
||||
|
||||
valid_letters = {letter.upper() for letter in valid_letters}
|
||||
for pattern in patterns:
|
||||
match = re.search(pattern, text, flags=re.IGNORECASE)
|
||||
if match:
|
||||
letter = match.group(1).upper()
|
||||
if letter in valid_letters:
|
||||
return letter
|
||||
|
||||
# Fallback: last standalone capital letter that is valid.
|
||||
candidates = re.findall(r"\b([A-Z])\b", text)
|
||||
for letter in reversed(candidates):
|
||||
letter = letter.upper()
|
||||
if letter in valid_letters:
|
||||
return letter
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def compute_gpqa_reward(response: str, label, metadata: dict | None = None) -> float:
|
||||
"""Rule-based scorer for GPQA-style multiple-choice evaluation."""
|
||||
if response is None:
|
||||
return 0.0
|
||||
|
||||
metadata = metadata or {}
|
||||
|
||||
choices = metadata.get("choices")
|
||||
if isinstance(choices, dict):
|
||||
choices = list(choices.values())
|
||||
elif choices is not None:
|
||||
choices = list(choices)
|
||||
|
||||
valid_letters = metadata.get("valid_letters")
|
||||
if valid_letters:
|
||||
valid_letters = [str(letter).upper() for letter in valid_letters]
|
||||
elif choices:
|
||||
valid_letters = list(string.ascii_uppercase[: len(choices)])
|
||||
else:
|
||||
valid_letters = DEFAULT_VALID_LETTERS
|
||||
|
||||
correct_letter = metadata.get("correct_letter")
|
||||
if isinstance(correct_letter, str):
|
||||
correct_letter = correct_letter.strip().upper()
|
||||
else:
|
||||
correct_letter = None
|
||||
|
||||
label_text = None
|
||||
if isinstance(label, str):
|
||||
label_text = label.strip()
|
||||
if len(label_text) == 1 and label_text.upper() in valid_letters and not correct_letter:
|
||||
correct_letter = label_text.upper()
|
||||
elif isinstance(label, (int, float)):
|
||||
idx = int(label)
|
||||
if 0 <= idx < len(valid_letters):
|
||||
correct_letter = valid_letters[idx]
|
||||
|
||||
if not correct_letter and choices and label_text:
|
||||
normalized_label = _normalize_text(label_text)
|
||||
for idx, choice in enumerate(choices):
|
||||
if _normalize_text(str(choice)) == normalized_label:
|
||||
correct_letter = valid_letters[idx]
|
||||
metadata.setdefault("correct_answer", choice)
|
||||
break
|
||||
|
||||
extracted_letter = _extract_letter_from_response(response, valid_letters)
|
||||
if extracted_letter and correct_letter:
|
||||
return 1.0 if extracted_letter == correct_letter else 0.0
|
||||
|
||||
candidate_answers = []
|
||||
if correct_letter and choices:
|
||||
try:
|
||||
idx = valid_letters.index(correct_letter)
|
||||
except ValueError:
|
||||
idx = None
|
||||
if idx is not None and idx < len(choices):
|
||||
candidate_answers.append(str(choices[idx]))
|
||||
|
||||
for key in ("correct_answer", "answer_text"):
|
||||
value = metadata.get(key)
|
||||
if value:
|
||||
candidate_answers.append(str(value))
|
||||
|
||||
if label_text:
|
||||
candidate_answers.append(label_text)
|
||||
|
||||
normalized_targets = {_normalize_text(text) for text in candidate_answers if text}
|
||||
normalized_response = _normalize_text(_strip_chain_of_thought(response))
|
||||
for target in normalized_targets:
|
||||
if target and target in normalized_response:
|
||||
return 1.0
|
||||
|
||||
if extracted_letter and not correct_letter and label_text:
|
||||
return 1.0 if extracted_letter == label_text.strip().upper() else 0.0
|
||||
|
||||
return 0.0
|
||||
173
slime/rollout/rm_hub/ifbench.py
Normal file
173
slime/rollout/rm_hub/ifbench.py
Normal file
@@ -0,0 +1,173 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import logging
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from collections.abc import Sequence
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_WORKSPACE_ROOT = Path(__file__).resolve().parents[3]
|
||||
_WORKSPACE_PARENT = _WORKSPACE_ROOT.parent
|
||||
_LOCAL_IFBENCH_REQUIREMENTS = _WORKSPACE_ROOT / "examples" / "eval_multi_task" / "requirements_ifbench.txt"
|
||||
|
||||
|
||||
def _ensure_ifbench_repo() -> Path:
|
||||
"""Clone IFBench repo if needed and ensure it is available on sys.path."""
|
||||
|
||||
repo_path = _WORKSPACE_PARENT / "IFBench"
|
||||
|
||||
if not repo_path.exists():
|
||||
clone_cmd = ["git", "clone", "https://github.com/allenai/IFBench.git", str(repo_path)]
|
||||
try:
|
||||
subprocess.run(clone_cmd, check=True, capture_output=True)
|
||||
except Exception as exc:
|
||||
raise ImportError(
|
||||
"Unable to automatically clone IFBench. Please clone "
|
||||
"https://github.com/allenai/IFBench.git into the repo root."
|
||||
) from exc
|
||||
|
||||
repo_str = str(repo_path)
|
||||
if repo_str not in sys.path:
|
||||
sys.path.insert(0, repo_str)
|
||||
|
||||
current_pythonpath = os.environ.get("PYTHONPATH")
|
||||
if current_pythonpath is None:
|
||||
os.environ["PYTHONPATH"] = repo_str
|
||||
elif repo_str not in current_pythonpath.split(os.pathsep):
|
||||
os.environ["PYTHONPATH"] = os.pathsep.join([repo_str, current_pythonpath])
|
||||
|
||||
return repo_path
|
||||
|
||||
|
||||
def _ensure_ifbench_dependencies(repo_path: Path) -> None:
|
||||
"""Install IFBench requirements the first time the module is imported."""
|
||||
|
||||
requirements_file = _LOCAL_IFBENCH_REQUIREMENTS
|
||||
|
||||
if not requirements_file.exists():
|
||||
logger.debug("Local IFBench requirements file not found at %s; skipping install.", requirements_file)
|
||||
return
|
||||
|
||||
sentinel = repo_path / ".deps_installed"
|
||||
if sentinel.exists():
|
||||
return
|
||||
|
||||
install_cmd = [sys.executable, "-m", "pip", "install", "-r", str(requirements_file)]
|
||||
try:
|
||||
subprocess.run(install_cmd, check=True)
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to install IFBench dependencies automatically: %s", exc)
|
||||
else:
|
||||
sentinel.write_text("installed\n")
|
||||
|
||||
|
||||
def _load_evaluation_lib():
|
||||
repo_path = _ensure_ifbench_repo()
|
||||
try:
|
||||
return importlib.import_module("evaluation_lib")
|
||||
except ImportError:
|
||||
_ensure_ifbench_dependencies(repo_path)
|
||||
return importlib.import_module("evaluation_lib")
|
||||
|
||||
|
||||
evaluation_lib = _load_evaluation_lib()
|
||||
InputExample = evaluation_lib.InputExample
|
||||
|
||||
|
||||
JsonDict = dict[str, Any]
|
||||
KwargsDict = dict[str, str | int | float | None]
|
||||
|
||||
|
||||
def _normalize_instruction_ids(raw_ids: Sequence[Any]) -> list[str]:
|
||||
"""Ensure instruction identifiers are clean strings."""
|
||||
|
||||
normalized: list[str] = []
|
||||
for entry in raw_ids or []:
|
||||
if entry is None:
|
||||
continue
|
||||
text = str(entry).strip()
|
||||
if not text:
|
||||
continue
|
||||
normalized.append(text)
|
||||
return normalized
|
||||
|
||||
|
||||
def _coerce_kwargs_list(
|
||||
raw_kwargs: Any,
|
||||
num_instructions: int,
|
||||
) -> list[KwargsDict]:
|
||||
"""Convert stored kwargs into the list structure expected by IFBench."""
|
||||
|
||||
if isinstance(raw_kwargs, list):
|
||||
processed: list[KwargsDict] = []
|
||||
for entry in raw_kwargs:
|
||||
if isinstance(entry, dict):
|
||||
processed.append(dict(entry))
|
||||
else:
|
||||
processed.append({})
|
||||
elif isinstance(raw_kwargs, dict):
|
||||
processed = [dict(raw_kwargs) for _ in range(num_instructions)]
|
||||
else:
|
||||
processed = [{} for _ in range(num_instructions)]
|
||||
|
||||
if len(processed) < num_instructions:
|
||||
tail = processed[-1] if processed else {}
|
||||
processed.extend([dict(tail) for _ in range(num_instructions - len(processed))])
|
||||
elif len(processed) > num_instructions:
|
||||
processed = processed[:num_instructions]
|
||||
|
||||
# Remove explicit None values to match official preprocessing.
|
||||
sanitized: list[KwargsDict] = []
|
||||
for entry in processed:
|
||||
sanitized.append({k: v for k, v in entry.items() if v is not None})
|
||||
return sanitized
|
||||
|
||||
|
||||
def _build_input_example(metadata: JsonDict) -> InputExample | None:
|
||||
instruction_ids = _normalize_instruction_ids(metadata.get("instruction_id_list") or [])
|
||||
if not instruction_ids:
|
||||
logger.debug("Missing instruction identifiers in metadata: %s", metadata)
|
||||
return None
|
||||
|
||||
prompt_text = metadata.get("prompt_text")
|
||||
if prompt_text is None:
|
||||
prompt_text = ""
|
||||
else:
|
||||
prompt_text = str(prompt_text)
|
||||
|
||||
raw_kwargs = metadata.get("kwargs")
|
||||
kwargs_list = _coerce_kwargs_list(raw_kwargs, len(instruction_ids))
|
||||
|
||||
return InputExample(
|
||||
key=int(metadata.get("record_id") or 0),
|
||||
instruction_id_list=instruction_ids,
|
||||
prompt=prompt_text,
|
||||
kwargs=kwargs_list,
|
||||
)
|
||||
|
||||
|
||||
def compute_ifbench_reward(response: str, label: Any, metadata: JsonDict | None = None) -> float:
|
||||
"""Score a model response using the official IFBench rules."""
|
||||
|
||||
if metadata is None:
|
||||
logger.debug("No metadata provided for IFBench scoring.")
|
||||
return 0.0
|
||||
|
||||
if response is None:
|
||||
return 0.0
|
||||
|
||||
inp = _build_input_example(metadata)
|
||||
if inp is None:
|
||||
return 0.0
|
||||
|
||||
prompt_to_response = {inp.prompt: str(response or "")}
|
||||
output = evaluation_lib.test_instruction_following_strict(inp, prompt_to_response)
|
||||
return 1.0 if output.follow_all_instructions else 0.0
|
||||
292
slime/rollout/rm_hub/math_dapo_utils.py
Normal file
292
slime/rollout/rm_hub/math_dapo_utils.py
Normal file
@@ -0,0 +1,292 @@
|
||||
# Copyright 2024 Bytedance Ltd. and/or its affiliates
|
||||
# Copyright 2022 EleutherAI and the HuggingFace Inc. team. All rights reserved.
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# Adapted from https://github.com/EleutherAI/lm-evaluation-harness/blob/main/lm_eval/tasks/hendrycks_math/utils.py
|
||||
|
||||
import re
|
||||
import signal
|
||||
|
||||
|
||||
def last_boxed_only_string(string: str) -> str | None:
|
||||
"""Extract the last LaTeX boxed expression from a string.
|
||||
|
||||
Args:
|
||||
string: Input string containing LaTeX code
|
||||
|
||||
Returns:
|
||||
The last boxed expression or None if not found
|
||||
"""
|
||||
idx = string.rfind("\\boxed{")
|
||||
if idx < 0:
|
||||
return None
|
||||
|
||||
i = idx
|
||||
right_brace_idx = None
|
||||
num_left_braces_open = 0
|
||||
|
||||
while i < len(string):
|
||||
if string[i] == "{":
|
||||
num_left_braces_open += 1
|
||||
if string[i] == "}":
|
||||
num_left_braces_open -= 1
|
||||
if num_left_braces_open == 0:
|
||||
right_brace_idx = i
|
||||
break
|
||||
i += 1
|
||||
|
||||
return string[idx : right_brace_idx + 1] if right_brace_idx is not None else None
|
||||
|
||||
|
||||
def remove_boxed(s: str) -> str:
|
||||
"""Remove the LaTeX boxed command from a string.
|
||||
|
||||
Args:
|
||||
s: String with format "\\boxed{content}"
|
||||
|
||||
Returns:
|
||||
The content inside the boxed command
|
||||
"""
|
||||
left = "\\boxed{"
|
||||
assert s[: len(left)] == left, f"box error: {s}"
|
||||
assert s[-1] == "}", f"box error: {s}"
|
||||
return s[len(left) : -1]
|
||||
|
||||
|
||||
class timeout:
|
||||
|
||||
def __init__(self, seconds=1, error_message="Timeout"):
|
||||
self.seconds = seconds
|
||||
self.error_message = error_message
|
||||
|
||||
def handle_timeout(self, signum, frame):
|
||||
raise TimeoutError(self.error_message)
|
||||
|
||||
def __enter__(self):
|
||||
signal.signal(signal.SIGALRM, self.handle_timeout)
|
||||
signal.alarm(self.seconds)
|
||||
|
||||
def __exit__(self, type, value, traceback):
|
||||
signal.alarm(0)
|
||||
|
||||
|
||||
# Constants for normalization
|
||||
SUBSTITUTIONS = [
|
||||
("an ", ""),
|
||||
("a ", ""),
|
||||
(".$", "$"),
|
||||
("\\$", ""),
|
||||
(r"\ ", ""),
|
||||
(" ", ""),
|
||||
("mbox", "text"),
|
||||
(",\\text{and}", ","),
|
||||
("\\text{and}", ","),
|
||||
("\\text{m}", "\\text{}"),
|
||||
]
|
||||
|
||||
REMOVED_EXPRESSIONS = [
|
||||
"square",
|
||||
"ways",
|
||||
"integers",
|
||||
"dollars",
|
||||
"mph",
|
||||
"inches",
|
||||
"hours",
|
||||
"km",
|
||||
"units",
|
||||
"\\ldots",
|
||||
"sue",
|
||||
"points",
|
||||
"feet",
|
||||
"minutes",
|
||||
"digits",
|
||||
"cents",
|
||||
"degrees",
|
||||
"cm",
|
||||
"gm",
|
||||
"pounds",
|
||||
"meters",
|
||||
"meals",
|
||||
"edges",
|
||||
"students",
|
||||
"childrentickets",
|
||||
"multiples",
|
||||
"\\text{s}",
|
||||
"\\text{.}",
|
||||
"\\text{\ns}",
|
||||
"\\text{}^2",
|
||||
"\\text{}^3",
|
||||
"\\text{\n}",
|
||||
"\\text{}",
|
||||
r"\mathrm{th}",
|
||||
r"^\circ",
|
||||
r"^{\circ}",
|
||||
r"\;",
|
||||
r",\!",
|
||||
"{,}",
|
||||
'"',
|
||||
"\\dots",
|
||||
"<|im_end|>",
|
||||
"<|endoftext|>",
|
||||
]
|
||||
|
||||
|
||||
def normalize_final_answer(final_answer: str) -> str:
|
||||
"""Normalize a final answer to a quantitative reasoning question.
|
||||
|
||||
Args:
|
||||
final_answer: The answer string to normalize
|
||||
|
||||
Returns:
|
||||
Normalized answer string
|
||||
"""
|
||||
final_answer = str(final_answer)
|
||||
final_answer = final_answer.split("=")[-1]
|
||||
|
||||
# Apply substitutions and removals
|
||||
for before, after in SUBSTITUTIONS:
|
||||
final_answer = final_answer.replace(before, after)
|
||||
for expr in REMOVED_EXPRESSIONS:
|
||||
final_answer = final_answer.replace(expr, "")
|
||||
|
||||
# Extract and normalize LaTeX math
|
||||
final_answer = re.sub(r"(.*?)(\$)(.*?)(\$)(.*)", "$\\3$", final_answer)
|
||||
final_answer = re.sub(r"(\\text\{)(.*?)(\})", "\\2", final_answer)
|
||||
final_answer = re.sub(r"(\\textbf\{)(.*?)(\})", "\\2", final_answer)
|
||||
final_answer = re.sub(r"(\\overline\{)(.*?)(\})", "\\2", final_answer)
|
||||
final_answer = re.sub(r"(\\boxed\{)(.*)(\})", "\\2", final_answer)
|
||||
|
||||
# Normalize shorthand TeX:
|
||||
# \fracab -> \frac{a}{b}
|
||||
# \frac{abc}{bef} -> \frac{abc}{bef}
|
||||
# \fracabc -> \frac{a}{b}c
|
||||
# \sqrta -> \sqrt{a}
|
||||
# \sqrtab -> sqrt{a}b
|
||||
final_answer = re.sub(r"(frac)([^{])(.)", "frac{\\2}{\\3}", final_answer)
|
||||
final_answer = re.sub(r"(sqrt)([^{])", "sqrt{\\2}", final_answer)
|
||||
final_answer = final_answer.replace("$", "")
|
||||
|
||||
# Normalize numbers
|
||||
if final_answer.replace(",", "").isdigit():
|
||||
final_answer = final_answer.replace(",", "")
|
||||
|
||||
return final_answer.strip()
|
||||
|
||||
|
||||
def is_correct_minerva(
|
||||
solution_str: str, gt: str, gt_need_extract: bool = False, answer_pattern: str = r"(?i)Answer\s*:\s*([^\n]+)"
|
||||
) -> tuple[bool, str]:
|
||||
"""Check if the solution is correct according to Minerva criteria.
|
||||
|
||||
Args:
|
||||
solution_str: The solution string to check
|
||||
gt: The ground truth answer
|
||||
gt_need_extract: Whether the ground truth needs extraction
|
||||
answer_pattern: Regex pattern to extract the answer
|
||||
|
||||
Returns:
|
||||
Tuple of (is_correct, normalized_prediction)
|
||||
"""
|
||||
# Extract answer from solution
|
||||
match = re.findall(answer_pattern, solution_str)
|
||||
extracted_answer = match[-1] if match else "[INVALID]"
|
||||
pred = normalize_final_answer(extracted_answer)
|
||||
|
||||
# Process ground truth
|
||||
if gt_need_extract:
|
||||
gt = normalize_final_answer(remove_boxed(last_boxed_only_string(gt)))
|
||||
else:
|
||||
gt = normalize_final_answer(gt)
|
||||
|
||||
gt = str(int(float(gt))) # in dapo, all answers are integers
|
||||
|
||||
return (pred == gt), pred
|
||||
|
||||
|
||||
def is_correct_strict_box(pred: str, gt: str, pause_tokens_index: list[int] | None = None) -> tuple[int, str | None]:
|
||||
"""Check if the prediction is correct using strict boxed answer criteria.
|
||||
|
||||
Args:
|
||||
pred: The prediction string
|
||||
gt: The ground truth answer
|
||||
pause_tokens_index: Indices of pause tokens
|
||||
|
||||
Returns:
|
||||
Tuple of (score, extracted_prediction)
|
||||
"""
|
||||
# Extract the relevant part of the prediction
|
||||
if pause_tokens_index is not None:
|
||||
assert len(pause_tokens_index) == 4
|
||||
pred = pred[pause_tokens_index[-1] - 100 :]
|
||||
else:
|
||||
pred = pred[-100:]
|
||||
|
||||
# Extract and check the boxed answer
|
||||
boxed_pred = last_boxed_only_string(pred)
|
||||
extracted_pred = remove_boxed(boxed_pred) if boxed_pred is not None else None
|
||||
|
||||
return 1 if (extracted_pred == gt) else -1, extracted_pred
|
||||
|
||||
|
||||
def verify(
|
||||
solution_str: str, answer: str, strict_box_verify: bool = False, pause_tokens_index: list[int] | None = None
|
||||
) -> bool:
|
||||
"""Verify if the solution is correct.
|
||||
|
||||
Args:
|
||||
solution_str: The solution string to verify
|
||||
answer: The ground truth answer
|
||||
strict_box_verify: Whether to use strict box verification
|
||||
pause_tokens_index: Indices of pause tokens
|
||||
|
||||
Returns:
|
||||
True if the solution is correct, False otherwise
|
||||
"""
|
||||
if strict_box_verify:
|
||||
correct, pred = is_correct_strict_box(solution_str, answer, pause_tokens_index)
|
||||
return correct == 1, pred
|
||||
|
||||
correct, pred = is_correct_minerva(solution_str, answer)
|
||||
return correct, pred
|
||||
|
||||
|
||||
def compute_score(
|
||||
solution_str: str,
|
||||
ground_truth: str,
|
||||
strict_box_verify: bool = False,
|
||||
pause_tokens_index: list[int] | None = None,
|
||||
) -> float:
|
||||
"""Compute the reward score for a solution.
|
||||
|
||||
Args:
|
||||
solution_str: The solution string
|
||||
ground_truth: The ground truth answer
|
||||
config: Configuration object containing reward model settings
|
||||
pause_tokens_index: Indices of pause tokens
|
||||
|
||||
Returns:
|
||||
Reward score (1.0 for correct, -1.0 for incorrect)
|
||||
"""
|
||||
# Limit solution length for efficiency
|
||||
solution_str = solution_str[-300:] # The longest answer in MATH-500 has 159 characters
|
||||
|
||||
# Verify the solution
|
||||
correct, pred = verify(solution_str, ground_truth, strict_box_verify, pause_tokens_index)
|
||||
|
||||
reward = 1.0 if correct else -1.0
|
||||
acc = correct
|
||||
|
||||
return {
|
||||
"score": reward,
|
||||
"acc": acc,
|
||||
"pred": pred,
|
||||
}
|
||||
491
slime/rollout/rm_hub/math_utils.py
Normal file
491
slime/rollout/rm_hub/math_utils.py
Normal file
@@ -0,0 +1,491 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# from https://github.com/agentica-project/deepscaler/blob/e6080ccd974eb64bd3430f0b36108244a6fee330/deepscaler/rewards/math_utils/utils.py
|
||||
"""
|
||||
Answer checker API that uses sympy to simplify expressions and check for equality.
|
||||
|
||||
Call grade_answer(given_answer: str, ground_truth: str).
|
||||
"""
|
||||
import re
|
||||
|
||||
import sympy
|
||||
from pylatexenc import latex2text
|
||||
from sympy.parsing import sympy_parser
|
||||
|
||||
|
||||
# Dan Hendrycks' code
|
||||
def mathd_normalize_answer(answer: str | None) -> str | None:
|
||||
if answer is None:
|
||||
return None
|
||||
answer = answer.strip()
|
||||
try:
|
||||
# Remove enclosing `\text{}`.
|
||||
m = re.search("^\\\\text\{(?P<text>.+?)\}$", answer)
|
||||
if m is not None:
|
||||
answer = m.group("text").strip()
|
||||
return _strip_string(answer)
|
||||
except Exception:
|
||||
return answer
|
||||
|
||||
|
||||
def _strip_string(string):
|
||||
def _fix_fracs(string):
|
||||
substrs = string.split("\\frac")
|
||||
new_str = substrs[0]
|
||||
if len(substrs) > 1:
|
||||
substrs = substrs[1:]
|
||||
for substr in substrs:
|
||||
new_str += "\\frac"
|
||||
if substr[0] == "{":
|
||||
new_str += substr
|
||||
else:
|
||||
try:
|
||||
assert len(substr) >= 2
|
||||
except Exception:
|
||||
return string
|
||||
a = substr[0]
|
||||
b = substr[1]
|
||||
if b != "{":
|
||||
if len(substr) > 2:
|
||||
post_substr = substr[2:]
|
||||
new_str += "{" + a + "}{" + b + "}" + post_substr
|
||||
else:
|
||||
new_str += "{" + a + "}{" + b + "}"
|
||||
else:
|
||||
if len(substr) > 2:
|
||||
post_substr = substr[2:]
|
||||
new_str += "{" + a + "}" + b + post_substr
|
||||
else:
|
||||
new_str += "{" + a + "}" + b
|
||||
string = new_str
|
||||
return string
|
||||
|
||||
def _fix_a_slash_b(string):
|
||||
if len(string.split("/")) != 2:
|
||||
return string
|
||||
a = string.split("/")[0]
|
||||
b = string.split("/")[1]
|
||||
try:
|
||||
a = int(a)
|
||||
b = int(b)
|
||||
assert string == f"{a}/{b}"
|
||||
new_string = "\\frac{" + str(a) + "}{" + str(b) + "}"
|
||||
return new_string
|
||||
except Exception:
|
||||
return string
|
||||
|
||||
def _remove_right_units(string):
|
||||
# "\\text{ " only ever occurs (at least in the val set) when describing units
|
||||
if "\\text{ " in string:
|
||||
splits = string.split("\\text{ ")
|
||||
assert len(splits) == 2
|
||||
return splits[0]
|
||||
else:
|
||||
return string
|
||||
|
||||
def _fix_sqrt(string):
|
||||
if "\\sqrt" not in string:
|
||||
return string
|
||||
splits = string.split("\\sqrt")
|
||||
new_string = splits[0]
|
||||
for split in splits[1:]:
|
||||
if split[0] != "{":
|
||||
a = split[0]
|
||||
new_substr = "\\sqrt{" + a + "}" + split[1:]
|
||||
else:
|
||||
new_substr = "\\sqrt" + split
|
||||
new_string += new_substr
|
||||
return new_string
|
||||
|
||||
# linebreaks
|
||||
string = string.replace("\n", "")
|
||||
|
||||
# remove inverse spaces
|
||||
string = string.replace("\\!", "")
|
||||
|
||||
# replace \\ with \
|
||||
string = string.replace("\\\\", "\\")
|
||||
|
||||
# replace tfrac and dfrac with frac
|
||||
string = string.replace("tfrac", "frac")
|
||||
string = string.replace("dfrac", "frac")
|
||||
|
||||
# remove \left and \right
|
||||
string = string.replace("\\left", "")
|
||||
string = string.replace("\\right", "")
|
||||
|
||||
# Remove circ (degrees)
|
||||
string = string.replace("^{\\circ}", "")
|
||||
string = string.replace("^\\circ", "")
|
||||
|
||||
# remove dollar signs
|
||||
string = string.replace("\\$", "")
|
||||
|
||||
# remove units (on the right)
|
||||
string = _remove_right_units(string)
|
||||
|
||||
# remove percentage
|
||||
string = string.replace("\\%", "")
|
||||
string = string.replace("\%", "")
|
||||
|
||||
# " 0." equivalent to " ." and "{0." equivalent to "{." Alternatively, add "0" if "." is the start of the string
|
||||
string = string.replace(" .", " 0.")
|
||||
string = string.replace("{.", "{0.")
|
||||
# if empty, return empty string
|
||||
if len(string) == 0:
|
||||
return string
|
||||
if string[0] == ".":
|
||||
string = "0" + string
|
||||
|
||||
# to consider: get rid of e.g. "k = " or "q = " at beginning
|
||||
if len(string.split("=")) == 2:
|
||||
if len(string.split("=")[0]) <= 2:
|
||||
string = string.split("=")[1]
|
||||
|
||||
# fix sqrt3 --> sqrt{3}
|
||||
string = _fix_sqrt(string)
|
||||
|
||||
# remove spaces
|
||||
string = string.replace(" ", "")
|
||||
|
||||
# \frac1b or \frac12 --> \frac{1}{b} and \frac{1}{2}, etc. Even works with \frac1{72} (but not \frac{72}1). Also does a/b --> \\frac{a}{b}
|
||||
string = _fix_fracs(string)
|
||||
|
||||
# manually change 0.5 --> \frac{1}{2}
|
||||
if string == "0.5":
|
||||
string = "\\frac{1}{2}"
|
||||
|
||||
# NOTE: X/Y changed to \frac{X}{Y} in dataset, but in simple cases fix in case the model output is X/Y
|
||||
string = _fix_a_slash_b(string)
|
||||
|
||||
return string
|
||||
|
||||
|
||||
# sympy might hang -- we don't care about trying to be lenient in these cases
|
||||
BAD_SUBSTRINGS = ["^{", "^("]
|
||||
BAD_REGEXES = ["\^[0-9]+\^", "\^[0-9][0-9]+"]
|
||||
TUPLE_CHARS = "()[]"
|
||||
|
||||
|
||||
def _sympy_parse(expr: str):
|
||||
"""Parses an expression with sympy."""
|
||||
py_expr = expr.replace("^", "**")
|
||||
return sympy_parser.parse_expr(
|
||||
py_expr,
|
||||
transformations=(sympy_parser.standard_transformations + (sympy_parser.implicit_multiplication_application,)),
|
||||
)
|
||||
|
||||
|
||||
def _parse_latex(expr: str) -> str:
|
||||
"""Attempts to parse latex to an expression sympy can read."""
|
||||
expr = expr.replace("\\tfrac", "\\frac")
|
||||
expr = expr.replace("\\dfrac", "\\frac")
|
||||
expr = expr.replace("\\frac", " \\frac") # Play nice with mixed numbers.
|
||||
expr = latex2text.LatexNodes2Text().latex_to_text(expr)
|
||||
|
||||
# Replace the specific characters that this parser uses.
|
||||
expr = expr.replace("√", "sqrt")
|
||||
expr = expr.replace("π", "pi")
|
||||
expr = expr.replace("∞", "inf")
|
||||
expr = expr.replace("∪", "U")
|
||||
expr = expr.replace("·", "*")
|
||||
expr = expr.replace("×", "*")
|
||||
|
||||
return expr.strip()
|
||||
|
||||
|
||||
def _is_float(num: str) -> bool:
|
||||
try:
|
||||
float(num)
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def _is_int(x: float) -> bool:
|
||||
try:
|
||||
return abs(x - int(round(x))) <= 1e-7
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def _is_frac(expr: str) -> bool:
|
||||
return bool(re.search(r"^-?[0-9]+.?/0*[1-9][0-9]*.?$", expr))
|
||||
|
||||
|
||||
def _str_is_int(x: str) -> bool:
|
||||
try:
|
||||
x = _strip_properly_formatted_commas(x)
|
||||
x = float(x)
|
||||
return abs(x - int(round(x))) <= 1e-7
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def _str_to_int(x: str) -> int:
|
||||
x = x.replace(",", "")
|
||||
x = float(x)
|
||||
return int(x)
|
||||
|
||||
|
||||
def _inject_implicit_mixed_number(step: str):
|
||||
"""
|
||||
Automatically make a mixed number evalable
|
||||
e.g. 7 3/4 => 7+3/4
|
||||
"""
|
||||
p1 = re.compile("([0-9]) +([0-9])")
|
||||
step = p1.sub("\\1+\\2", step) ## implicit mults
|
||||
return step
|
||||
|
||||
|
||||
def _strip_properly_formatted_commas(expr: str):
|
||||
# We want to be careful because we don't want to strip tuple commas
|
||||
p1 = re.compile("(\d)(,)(\d\d\d)($|\D)")
|
||||
while True:
|
||||
next_expr = p1.sub("\\1\\3\\4", expr)
|
||||
if next_expr == expr:
|
||||
break
|
||||
expr = next_expr
|
||||
return next_expr
|
||||
|
||||
|
||||
def _normalize(expr: str) -> str:
|
||||
"""Normalize answer expressions."""
|
||||
if expr is None:
|
||||
return None
|
||||
|
||||
# Remove enclosing `\text{}`.
|
||||
m = re.search("^\\\\text\{(?P<text>.+?)\}$", expr)
|
||||
if m is not None:
|
||||
expr = m.group("text")
|
||||
|
||||
expr = expr.replace("\\%", "%")
|
||||
expr = expr.replace("\\$", "$")
|
||||
expr = expr.replace("$", "")
|
||||
expr = expr.replace("%", "")
|
||||
expr = expr.replace(" or ", " , ")
|
||||
expr = expr.replace(" and ", " , ")
|
||||
|
||||
expr = expr.replace("million", "*10^6")
|
||||
expr = expr.replace("billion", "*10^9")
|
||||
expr = expr.replace("trillion", "*10^12")
|
||||
|
||||
for unit in [
|
||||
"degree",
|
||||
"cm",
|
||||
"centimeter",
|
||||
"meter",
|
||||
"mile",
|
||||
"second",
|
||||
"minute",
|
||||
"hour",
|
||||
"day",
|
||||
"week",
|
||||
"month",
|
||||
"year",
|
||||
"foot",
|
||||
"feet",
|
||||
"inch",
|
||||
"yard",
|
||||
]:
|
||||
expr = re.sub(f"{unit}(es)?(s)? *(\^[0-9]+)?", "", expr)
|
||||
expr = re.sub("\^ *\\\\circ", "", expr)
|
||||
|
||||
if len(expr) > 0 and expr[0] == "{" and expr[-1] == "}":
|
||||
expr = expr[1:-1]
|
||||
|
||||
expr = re.sub(",\\\\! *", "", expr)
|
||||
if _is_float(expr) and _is_int(float(expr)):
|
||||
expr = str(int(round(float(expr))))
|
||||
if "\\" in expr:
|
||||
try:
|
||||
expr = _parse_latex(expr)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# edge case with mixed numbers and negative signs
|
||||
expr = re.sub("- *", "-", expr)
|
||||
|
||||
expr = _inject_implicit_mixed_number(expr)
|
||||
expr = expr.replace(" ", "")
|
||||
|
||||
# if we somehow still have latex braces here, just drop them
|
||||
expr = expr.replace("{", "")
|
||||
expr = expr.replace("}", "")
|
||||
|
||||
# don't be case sensitive for text answers
|
||||
expr = expr.lower()
|
||||
|
||||
if _str_is_int(expr):
|
||||
expr = str(_str_to_int(expr))
|
||||
|
||||
return expr
|
||||
|
||||
|
||||
def count_unknown_letters_in_expr(expr: str):
|
||||
expr = expr.replace("sqrt", "")
|
||||
expr = expr.replace("frac", "")
|
||||
letters_in_expr = set([x for x in expr if x.isalpha()])
|
||||
return len(letters_in_expr)
|
||||
|
||||
|
||||
def should_allow_eval(expr: str):
|
||||
# we don't want to try parsing unknown text or functions of more than two variables
|
||||
if count_unknown_letters_in_expr(expr) > 2:
|
||||
return False
|
||||
|
||||
for bad_string in BAD_SUBSTRINGS:
|
||||
if bad_string in expr:
|
||||
return False
|
||||
|
||||
for bad_regex in BAD_REGEXES:
|
||||
if re.search(bad_regex, expr) is not None:
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def are_equal_under_sympy(ground_truth_normalized: str, given_normalized: str):
|
||||
are_equal = False
|
||||
try:
|
||||
expr = f"({ground_truth_normalized})-({given_normalized})"
|
||||
if should_allow_eval(expr):
|
||||
sympy_diff = _sympy_parse(expr)
|
||||
simplified = sympy.simplify(sympy_diff)
|
||||
if simplified == 0:
|
||||
are_equal = True
|
||||
except Exception:
|
||||
pass
|
||||
return are_equal
|
||||
|
||||
|
||||
def split_tuple(expr: str):
|
||||
"""
|
||||
Split the elements in a tuple/interval, while handling well-formatted commas in large numbers
|
||||
"""
|
||||
expr = _strip_properly_formatted_commas(expr)
|
||||
if len(expr) == 0:
|
||||
return []
|
||||
if (
|
||||
len(expr) > 2
|
||||
and expr[0] in TUPLE_CHARS
|
||||
and expr[-1] in TUPLE_CHARS
|
||||
and all([ch not in expr[1:-1] for ch in TUPLE_CHARS])
|
||||
):
|
||||
elems = [elem.strip() for elem in expr[1:-1].split(",")]
|
||||
else:
|
||||
elems = [expr]
|
||||
return elems
|
||||
|
||||
|
||||
def last_boxed_only_string(string):
|
||||
idx = string.rfind("\\boxed")
|
||||
if idx < 0:
|
||||
idx = string.rfind("\\fbox")
|
||||
if idx < 0:
|
||||
return None
|
||||
|
||||
i = idx
|
||||
right_brace_idx = None
|
||||
num_left_braces_open = 0
|
||||
while i < len(string):
|
||||
if string[i] == "{":
|
||||
num_left_braces_open += 1
|
||||
if string[i] == "}":
|
||||
num_left_braces_open -= 1
|
||||
if num_left_braces_open == 0:
|
||||
right_brace_idx = i
|
||||
break
|
||||
i += 1
|
||||
|
||||
if right_brace_idx is None:
|
||||
retval = None
|
||||
else:
|
||||
retval = string[idx : right_brace_idx + 1]
|
||||
|
||||
return retval
|
||||
|
||||
|
||||
def remove_boxed(s):
|
||||
left = "\\boxed{"
|
||||
try:
|
||||
assert s[: len(left)] == left
|
||||
assert s[-1] == "}"
|
||||
return s[len(left) : -1]
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def extract_boxed_answer(solution: str) -> str:
|
||||
"""Extract the answer from inside a LaTeX \\boxed{} command"""
|
||||
solution = last_boxed_only_string(solution)
|
||||
solution = remove_boxed(solution)
|
||||
return solution
|
||||
|
||||
|
||||
def grade_answer_sympy(given_answer: str, ground_truth: str) -> bool:
|
||||
ground_truth_normalized = _normalize(ground_truth)
|
||||
given_normalized = _normalize(given_answer)
|
||||
|
||||
if ground_truth_normalized is None:
|
||||
return False
|
||||
|
||||
if ground_truth_normalized == given_normalized:
|
||||
return True
|
||||
|
||||
if len(given_normalized) == 0:
|
||||
return False
|
||||
|
||||
ground_truth_elems = split_tuple(ground_truth_normalized)
|
||||
given_elems = split_tuple(given_normalized)
|
||||
|
||||
if len(ground_truth_elems) > 1 and (
|
||||
ground_truth_normalized[0] != given_normalized[0] or ground_truth_normalized[-1] != given_normalized[-1]
|
||||
):
|
||||
is_correct = False
|
||||
elif len(ground_truth_elems) != len(given_elems):
|
||||
is_correct = False
|
||||
else:
|
||||
for ground_truth_elem, given_elem in zip(ground_truth_elems, given_elems, strict=False):
|
||||
if _is_frac(ground_truth_elem) and _is_frac(given_elem):
|
||||
# if fractions aren't reduced, then shouldn't be marked as correct
|
||||
# so, we don't want to allow sympy.simplify in this case
|
||||
is_correct = ground_truth_elem == given_elem
|
||||
elif _str_is_int(ground_truth_elem) != _str_is_int(given_elem):
|
||||
# if the ground truth answer is an integer, we require the given answer to be a strict match (no sympy.simplify)
|
||||
is_correct = False
|
||||
else:
|
||||
is_correct = are_equal_under_sympy(ground_truth_elem, given_elem)
|
||||
if not is_correct:
|
||||
break
|
||||
|
||||
return is_correct
|
||||
|
||||
|
||||
def grade_answer_mathd(given_answer: str, ground_truth: str) -> bool:
|
||||
ground_truth_normalized_mathd = mathd_normalize_answer(ground_truth)
|
||||
given_answer_normalized_mathd = mathd_normalize_answer(given_answer)
|
||||
|
||||
# be at least as lenient as mathd
|
||||
if ground_truth_normalized_mathd == given_answer_normalized_mathd:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def extract_answer(passage: str) -> str:
|
||||
if "\\boxed" in passage:
|
||||
return extract_boxed_answer(passage)
|
||||
return None
|
||||
|
||||
|
||||
def grade_answer_verl(solution_str, ground_truth):
|
||||
if not ground_truth:
|
||||
return False
|
||||
ground_truth = str(ground_truth)
|
||||
if "\\boxed" in ground_truth:
|
||||
ground_truth = extract_answer(ground_truth)
|
||||
given_answer = extract_answer(solution_str)
|
||||
if given_answer is None:
|
||||
return False
|
||||
return grade_answer_mathd(given_answer, ground_truth) or grade_answer_sympy(given_answer, ground_truth)
|
||||
Reference in New Issue
Block a user