初始化项目,由ModelHub XC社区提供模型

Model: ayh015/myLightningOPD
Source: Original Platform
This commit is contained in:
ModelHub XC
2026-08-27 23:50:14 +08:00
commit d4e0a1af66
368 changed files with 559583 additions and 0 deletions

View 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

Binary file not shown.

Binary file not shown.

View 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

View 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

View 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

View 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

View 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,
}

View 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)