#!/usr/bin/env python3 # -*- coding: utf-8 -*- import json import argparse from pathlib import Path def calculate_accuracy(jsonl_file_path: str) -> dict: """ 计算 JSONL 文件中 answer 和 agent_answer 的准确率 Args: jsonl_file_path: JSONL 文件路径 Returns: 包含准确率统计的字典 """ correct = 0 total = 0 errors = [] # 打开并逐行读取 JSONL 文件 with open(jsonl_file_path, 'r', encoding='utf-8') as f: for line_num, line in enumerate(f, 1): line = line.strip() if not line: # 跳过空行 continue try: data = json.loads(line) # 检查必要字段是否存在 if 'answer' not in data or 'agent_answer' not in data: print(f"⚠️ 第 {line_num} 行缺少 'answer' 或 'agent_answer' 字段,跳过") continue # 获取真实答案和模型答案 true_answer = str(data['answer']).strip() pred_answer = str(data['agent_answer']).strip() # 比较是否一致 if true_answer == pred_answer: correct += 1 else: errors.append({ 'line': line_num, 'true': true_answer, 'pred': pred_answer, 'question': data.get('question', 'N/A')[:100] + '...' # 截取前100字符 }) total += 1 except json.JSONDecodeError: print(f"⚠️ 第 {line_num} 行 JSON 解析失败,跳过") continue # 计算准确率 accuracy = correct / total if total > 0 else 0 return { 'total_samples': total, 'correct_samples': correct, 'incorrect_samples': total - correct, 'accuracy': accuracy, 'accuracy_percentage': round(accuracy * 100, 2), 'errors': errors[:10] # 只显示前10个错误,避免输出过多 } def main(): parser = argparse.ArgumentParser(description='计算 JSONL 问答数据的准确率') parser.add_argument('file', type=str, help='JSONL 文件路径') parser.add_argument('--show-errors', action='store_true', help='显示错误样本详情') args = parser.parse_args() # 检查文件是否存在 if not Path(args.file).exists(): print(f"❌ 文件不存在: {args.file}") return # 计算准确率 print(f"📊 正在分析文件: {args.file}") result = calculate_accuracy(args.file) # 打印结果 print("\n" + "="*50) print(f"📈 准确率统计结果") print("="*50) print(f"总样本数: {result['total_samples']}") print(f"正确样本数: {result['correct_samples']}") print(f"错误样本数: {result['incorrect_samples']}") print(f"准确率: {result['accuracy_percentage']}%") print("="*50) # 显示错误详情(如果有) if args.show_errors and result['errors']: print("\n🔍 错误样本示例 (前10个):") print("-"*50) for err in result['errors']: print(f"第 {err['line']} 行:") print(f" 真实答案: {err['true']}") print(f" 模型答案: {err['pred']}") print(f" 问题: {err['question']}") print("-"*30) if __name__ == "__main__": main()