ai-agent-book 精选快照(<2MB 代码与文档,来自 github.com/bojieli/ai-agent-book)
Build latest book artifacts / build (push) Canceled after 0s
dependency resolution / resolve (3.11) (push) Canceled after 0s
dependency resolution / resolve (3.13) (push) Canceled after 0s
deploy-pages / build (push) Canceled after 0s
deploy-pages / deploy (push) Canceled after 0s
i18n consistency check / check (push) Canceled after 0s
provider adoption tests / test (chapter2/context-compression) (push) Canceled after 0s
provider adoption tests / test (chapter2/prompt-injection) (push) Canceled after 0s
provider adoption tests / test (chapter2/system-hint) (push) Canceled after 0s
provider adoption tests / test (chapter3/log-sanitization) (push) Canceled after 0s
web-search-agent tests / test (push) Canceled after 0s
web-search-agent tests / agentbook (push) Canceled after 0s

This commit is contained in:
2026-08-20 13:12:50 +00:00
commit b119135836
10275 changed files with 3284984 additions and 0 deletions
File diff suppressed because it is too large Load Diff
+328
View File
@@ -0,0 +1,328 @@
#!/usr/bin/env python3
"""
从 lighteval 缓存的 parquet 文件中提取答案并计算 GSM8K 准确率
支持 \\boxed{} 和 #### 两种答案格式
"""
import re
import pandas as pd
import argparse
from pathlib import Path
from typing import Optional
def extract_answer_from_boxed(text: str) -> Optional[str]:
"""\\boxed{} 格式中提取答案(同时支持 \\(\\boxed{}\\) 形式)"""
if not text:
return None
# 如果是 bytes,转换为字符串
if isinstance(text, bytes):
text = text.decode('utf-8', errors='ignore')
# 确保是字符串
text = str(text)
# Balanced braces so nested LaTeX like \boxed{\frac{1}{2}} is not truncated.
marker = "\\boxed{"
start = text.find(marker)
if start < 0:
return None
i = start + len(marker)
depth = 1
while i < len(text) and depth:
ch = text[i]
if ch == "{":
depth += 1
elif ch == "}":
depth -= 1
i += 1
if depth != 0:
return None
return text[start + len(marker) : i - 1].strip()
def extract_answer_from_gsm8k_format(text: str) -> Optional[str]:
"""从 #### number 格式中提取答案"""
if not text:
return None
# 如果是 bytes,转换为字符串
if isinstance(text, bytes):
text = text.decode('utf-8', errors='ignore')
# 确保是字符串
text = str(text)
if "####" in text:
parts = text.split("####")
if len(parts) > 1:
return parts[-1].strip()
return None
def _format_normalized_number(num: float) -> str:
if num.is_integer():
return str(int(num))
return str(num)
def normalize_number(text: str) -> Optional[str]:
"""标准化数字格式:去除逗号、空格、LaTeX 符号等"""
if not text:
return None
# 如果是 bytes,转换为字符串
if isinstance(text, bytes):
text = text.decode('utf-8', errors='ignore')
# 确保是字符串
text = str(text)
# Unwrap LaTeX formatting before parsing; the wrapped content may itself be numeric.
cleaned = re.sub(r'\\(?:text|mathrm|mathbf)\s*\{([^}]*)\}', r'\1', text)
cleaned = cleaned.replace("\\$", "").replace("$", "").replace("\\,", "").replace("\\text", "")
cleaned = cleaned.replace(",", "")
# Evaluate \frac{a}{b} before brace stripping (else "\frac{6}{2}" becomes "frac62").
frac = re.search(r'(-)?\s*\\(?:d)?frac\s*\{([^{}]+)\}\s*\{([^{}]+)\}', cleaned)
if frac:
try:
sign = -1.0 if frac.group(1) else 1.0
num_match = re.match(r'\s*(-?\s*\d+(?:\.\d+)?)', frac.group(2))
den_match = re.match(r'\s*(-?\s*\d+(?:\.\d+)?)', frac.group(3))
if not num_match or not den_match:
raise ValueError("fraction component does not start with a number")
num = float(num_match.group(1).replace(" ", ""))
den = float(den_match.group(1).replace(" ", ""))
if den != 0:
return _format_normalized_number(sign * (num / den))
except ValueError:
pass
# Plain a/b before taking the first digit run alone; allow spaces and units.
slash = re.search(r'(-?\s*\d+(?:\.\d+)?)\s*/\s*(-?\s*\d+(?:\.\d+)?)', cleaned)
if slash:
try:
num = float(slash.group(1).replace(" ", ""))
den = float(slash.group(2).replace(" ", ""))
if den != 0:
return _format_normalized_number(num / den)
except ValueError:
pass
# 去除 LaTeX 及货币符号
text = text.replace("\\$", "")
text = text.replace("$", "")
text = text.replace("\\,", "")
text = text.replace("\\text", "")
text = text.replace("{", "").replace("}", "")
# 去除逗号和空格
text = text.replace(",", "").replace(" ", "")
# 提取数字(包括小数和负数)
match = re.search(r'-?\d+\.?\d*', text)
if match:
num_str = match.group(0)
try:
return _format_normalized_number(float(num_str))
except ValueError:
return None
return None
def extract_and_normalize_answer(text: str) -> Optional[str]:
"""从模型输出中提取并标准化答案"""
if not text:
return None
# 如果是 bytes,转换为字符串
if isinstance(text, bytes):
text = text.decode('utf-8', errors='ignore')
# 确保是字符串
text = str(text)
# 先尝试提取 boxed 格式
answer = extract_answer_from_boxed(text)
# 如果没找到,尝试 GSM8K 格式
if not answer:
answer = extract_answer_from_gsm8k_format(text)
# 如果还是没找到,尝试从最后一句话提取数字
if not answer:
# 取最后 200 个字符,避免提取到过程中的数字
last_part = text[-200:] if len(text) > 200 else text
answer = last_part
# 标准化数字格式
return normalize_number(answer)
def load_gsm8k_answers(split: str = "test") -> dict:
"""加载 GSM8K 数据集的金标答案
返回一个字典,键是数据集中的原始索引(0-1318),值是标准化后的答案
"""
try:
from datasets import load_dataset
dataset = load_dataset("gsm8k", "main", split=split)
answers = {}
# 注意:这里的索引是数据集中的顺序索引,不是 sample_id
for idx in range(len(dataset)):
item = dataset[idx]
# GSM8K 答案格式:计算过程\n#### 答案
gold_answer = item["answer"]
# 提取 #### 后面的数字
normalized = extract_answer_from_gsm8k_format(gold_answer)
if normalized:
normalized = normalize_number(normalized)
answers[idx] = normalized
print(f"✅ 加载了 {len(answers)} 个金标答案")
return answers
except ImportError:
print("❌ 错误:需要安装 datasets 库")
print("运行:pip install datasets")
return {}
except Exception as e:
print(f"❌ 加载金标答案时出错: {e}")
return {}
def evaluate_from_parquet(parquet_path: str, verbose: bool = False):
"""从 parquet 文件评测"""
print(f"📂 读取预测结果: {parquet_path}")
df = pd.read_parquet(parquet_path)
print(f"📊 总样本数: {len(df)}")
# 加载金标答案
print("📥 加载 GSM8K 金标答案...")
gold_answers = load_gsm8k_answers()
if not gold_answers:
print("❌ 无法加载金标答案,退出")
return
# 评测
correct = 0
total = 0
errors = []
# 调试:显示前几个 sample_id
if verbose:
print(f"\n前 5 个 sample_id: {df['sample_id'].head().tolist()}")
print(f"金标答案的键范围: {min(gold_answers.keys()) if gold_answers else 'N/A'} - {max(gold_answers.keys()) if gold_answers else 'N/A'}")
for idx, row in df.iterrows():
sample_id = row['sample_id']
sample_data = row['sample']
# 转换 sample_id 为原生 intparquet 的数值列返回 np.int64
# 直接放进结果里会让最后的 json.dump 抛
# "Object of type int64 is not JSON serializable",把 -o 输出截断。
try:
sample_id = int(sample_id)
except (TypeError, ValueError):
if verbose:
print(f"⚠️ 样本 {sample_id}: 无法转换为整数")
continue
# 提取模型输出
text_field = sample_data.get('text', [''])
if isinstance(text_field, list):
model_output = text_field[0] if text_field else ''
else:
model_output = text_field if text_field is not None else ''
# 确保 model_output 是字符串
if isinstance(model_output, bytes):
model_output = model_output.decode('utf-8', errors='ignore')
model_output = str(model_output) if model_output else ''
# 提取并标准化答案
pred_answer = extract_and_normalize_answer(model_output)
gold_answer = gold_answers.get(sample_id)
if gold_answer is None:
if verbose and idx < 5:
print(f"⚠️ 样本 {sample_id}: 找不到金标答案")
continue
total += 1
is_correct = pred_answer == gold_answer
if is_correct:
correct += 1
else:
errors.append({
'sample_id': sample_id,
'predicted': pred_answer,
'gold': gold_answer,
'output': model_output[:200] + "..." if len(model_output) > 200 else model_output
})
if verbose and idx < 5:
print(f"\n样本 {sample_id}:")
print(f" 预测: {pred_answer}")
print(f" 金标: {gold_answer}")
print(f" 正确: {'' if is_correct else ''}")
# 计算准确率
accuracy = correct / total * 100 if total > 0 else 0
print("\n" + "="*80)
print("📈 评测结果")
print("="*80)
print(f"总样本数: {total}")
print(f"正确数量: {correct}")
print(f"错误数量: {total - correct}")
print(f"准确率: {accuracy:.2f}%")
print("="*80)
# 显示部分错误样本
if errors and verbose:
print("\n❌ 前 10 个错误样本:")
for i, error in enumerate(errors[:10], 1):
print(f"\n{i}. 样本 {error['sample_id']}:")
print(f" 预测: {error['predicted']}")
print(f" 金标: {error['gold']}")
print(f" 输出: {error['output']}")
return {
'total': total,
'correct': correct,
'accuracy': accuracy,
'errors': errors
}
def main():
parser = argparse.ArgumentParser(description='从 lighteval 缓存评测 GSM8K 结果')
parser.add_argument('parquet_file', type=str, help='Parquet 文件路径')
parser.add_argument('-v', '--verbose', action='store_true', help='显示详细信息和错误样本')
parser.add_argument('-o', '--output', type=str, help='保存结果到 JSON 文件')
args = parser.parse_args()
if not Path(args.parquet_file).exists():
print(f"❌ 错误:文件不存在: {args.parquet_file}")
return
results = evaluate_from_parquet(args.parquet_file, verbose=args.verbose)
if args.output and results:
import json
with open(args.output, 'w') as f:
json.dump(results, f, indent=2, ensure_ascii=False)
print(f"\n💾 结果已保存到: {args.output}")
if __name__ == "__main__":
main()
@@ -0,0 +1,31 @@
"""Regression: negative currency formats and negative fractions must preserve negative sign."""
from evaluate_from_cache import extract_answer_from_gsm8k_format
from evaluate_from_cache import extract_and_normalize_answer, normalize_number
def test_negative_dollar_amount():
assert normalize_number("-$42") == "-42"
def test_negative_latex_dollar_amount():
assert normalize_number(r"-\$42") == "-42"
def test_boxed_negative_dollar_amount():
assert extract_and_normalize_answer(r"\boxed{-\$42}") == "-42"
def test_negative_latex_frac():
assert normalize_number(r"-\frac{6}{2}") == "-3"
assert normalize_number(r"-\dfrac{6}{2}") == "-3"
def test_extract_gsm8k_multiple_hash_markers():
assert extract_answer_from_gsm8k_format("#### step 1 #### 42") == "42"
assert extract_and_normalize_answer("#### step 1 #### 42") == "42"
def test_negative_fraction_with_space():
assert normalize_number("- 1/2") == "-0.5"
assert extract_and_normalize_answer(r"\boxed{- 1/2}") == "-0.5"
+9
View File
@@ -0,0 +1,9 @@
"""Test import bootstrap for the Intuitor experiment."""
from pathlib import Path
import sys
EXPERIMENT_ROOT = Path(__file__).resolve().parents[1]
if str(EXPERIMENT_ROOT) not in sys.path:
sys.path.insert(0, str(EXPERIMENT_ROOT))
@@ -0,0 +1,18 @@
"""Regression: nested braces inside \\boxed{} must not truncate."""
from evaluate_from_cache import extract_answer_from_boxed
def test_boxed_nested_frac():
text = r"The answer is \boxed{\frac{1}{2}}"
assert extract_answer_from_boxed(text) == r"\frac{1}{2}"
def test_boxed_simple_integer():
text = r"Final answer: \boxed{42}"
assert extract_answer_from_boxed(text) == "42"
def test_boxed_deeper_nesting():
text = r"\boxed{\frac{a}{b+c}}"
assert extract_answer_from_boxed(text) == r"\frac{a}{b+c}"
@@ -0,0 +1,22 @@
"""Regression: empty sample text list must not IndexError."""
def _extract_model_output(sample_data):
text_field = sample_data.get('text', [''])
if isinstance(text_field, list):
return text_field[0] if text_field else ''
return text_field if text_field is not None else ''
def test_empty_text_list():
assert _extract_model_output({"text": []}) == ""
def test_nonempty_text_list():
assert _extract_model_output({"text": ["hello"]}) == "hello"
def test_source_guards_empty_list():
from pathlib import Path
src = (Path(__file__).resolve().parents[1] / "evaluate_from_cache.py").read_text()
assert "text_field[0] if text_field else ''" in src
@@ -0,0 +1,20 @@
"""Regression: \\frac{a}{b} and a/b must evaluate, not concatenate digit runs."""
from evaluate_from_cache import extract_and_normalize_answer, normalize_number
def test_frac_six_over_two():
assert extract_and_normalize_answer(r"\boxed{\frac{6}{2}}") == "3"
def test_plain_slash_six_over_two():
assert extract_and_normalize_answer(r"\boxed{6/2}") == "3"
def test_frac_one_half():
assert extract_and_normalize_answer(r"\boxed{\frac{1}{2}}") == "0.5"
def test_plain_integer_unchanged():
assert extract_and_normalize_answer(r"\boxed{42}") == "42"
assert normalize_number("1,234") == "1234"
@@ -0,0 +1,26 @@
"""Regression: LaTeX fractions and division with formatting/units must evaluate correctly."""
from evaluate_from_cache import extract_and_normalize_answer, normalize_number
def test_frac_thin_space_evaluates():
assert normalize_number(r"\frac{1\,000}{2}") == "500"
assert normalize_number(r"\dfrac{1\,500}{3}") == "500"
assert normalize_number(r"-\frac{1\,000}{2}") == "-500"
def test_frac_text_units_evaluates():
assert normalize_number(r"\frac{100\text{ kg}}{2}") == "50"
assert normalize_number(r"\frac{100}{2\text{ kg}}") == "50"
assert normalize_number(r"\frac{\text{100 kg}}{2}") == "50"
def test_frac_numeric_format_wrappers_evaluate():
assert normalize_number(r"\frac{\mathrm{1,000}}{2}") == "500"
assert normalize_number(r"\frac{\mathbf{6}}{2}") == "3"
def test_slash_division_with_units_and_formatting():
assert normalize_number("6/2 kg") == "3"
assert normalize_number("$6/2$") == "3"
assert normalize_number("1,000 / 2") == "500"