Files
rag/exam_pkg/grader.py
lacerate551 a31ee4bba0 feat(exam_pkg): 优化出题系统稳定性与推理模型适配
- core/llm_utils: MiMo模型自动注入thinking=disabled参数,全局生效
- config: 新增LLM_DISABLE_THINKING配置项(默认true)
- generator: 推理模型自适应max_tokens(1.5x)、429限流重试+指数退避、v2管线补题机制
- generator: analyze_document_for_exam新增max_total参数控制AI出题上限
- generator: validate_questions_schema兼容type和question_type字段
- grader: 主观题max_tokens从1000提升至2000、fuzzy_match增加编辑距离容错
- manager: results变量初始化防NameError、透传max_total参数
- api: /exam/generate-smart支持max_total请求参数
2026-06-22 18:51:11 +08:00

535 lines
17 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
批题器 - 本地批阅逻辑(后续可迁移到 Dify 工作流)
核心功能:
1. 本地批阅选择题/判断题
2. 填空题模糊匹配
3. 主观题 LLM 评分
4. 并发批阅 + 限流 + 顺序保持
使用方式:
from exam_pkg.grader import AnswerGrader
grader = AnswerGrader()
results = grader.grade_answers(answers)
"""
import json
import logging
import threading
import time
from concurrent.futures import ThreadPoolExecutor, as_completed
from functools import wraps
from typing import List, Dict, Any, Optional
# 导入 LLM 工具函数
from core.llm_utils import call_llm, extract_json_object
logger = logging.getLogger(__name__)
# 导入 LLM 配置
try:
from config import API_KEY, BASE_URL, MODEL
LLM_AVAILABLE = True
except ImportError:
API_KEY = None
BASE_URL = None
MODEL = None
LLM_AVAILABLE = False
# 推理模型识别(与 generator.py 共享同一套关键词)
_REASONING_MODEL_KEYWORDS = ('mimo', 'qwq', 'deepseek-r1', 'deepseek-reasoner', 'o1', 'o3')
def _is_reasoning_model(model_name: str) -> bool:
if not model_name:
return False
name_lower = model_name.lower()
return any(kw in name_lower for kw in _REASONING_MODEL_KEYWORDS)
def _get_effective_max_tokens(base_max: int, model_name: str) -> int:
"""推理模型需要 1.5x token 预算给思考链"""
if _is_reasoning_model(model_name):
return max(base_max, int(base_max * 1.5))
return base_max
# ==================== 装饰器 ====================
def retry(times: int = 2, delay: float = 1.0):
"""
🔥 P1 改进:重试装饰器
Args:
times: 重试次数
delay: 重试间隔(秒)
"""
def decorator(func):
@wraps(func)
def wrapper(*args, **kwargs):
last_error = None
for i in range(times):
try:
return func(*args, **kwargs)
except Exception as e:
last_error = e
if i < times - 1:
time.sleep(delay)
raise last_error
return wrapper
return decorator
# ==================== 限流 ====================
# 🔥 P2 改进:限流信号量
MAX_CONCURRENT_GRADING = 3
grading_semaphore = threading.Semaphore(MAX_CONCURRENT_GRADING)
# ==================== 本地批阅函数 ====================
def grade_objective(answer: Dict) -> Dict:
"""
批阅客观题(选择/判断)
🔥 本地直接判断,无 LLM 调用
"""
q_type = answer['question_type']
question_content = answer.get('content', {})
correct_answer = question_content.get('answer')
student_answer = answer.get('student_answer')
max_score = answer.get('max_score', 2.0)
# 判断正确性
if q_type == 'single_choice':
correct = student_answer == correct_answer
elif q_type == 'multiple_choice':
# 多选题:答案顺序无关
correct = set(student_answer) == set(correct_answer) if isinstance(student_answer, list) else False
elif q_type == 'true_false':
correct = student_answer == correct_answer
else:
correct = False
return {
"question_id": answer.get('question_id'),
"score": max_score if correct else 0,
"max_score": max_score,
"grading_status": "success",
"details": {
"correct": correct,
"student_answer": student_answer,
"correct_answer": correct_answer,
"feedback": f"正确答案: {correct_answer}" if not correct else "正确!"
}
}
def grade_fill_blank(answer: Dict) -> Dict:
"""
批阅填空题 - 支持同义词匹配
填空题答案格式:[["答案1", "同义词1", ...], ["答案2", ...], ...]
学生答案格式:["学生答案1", "学生答案2", ...]
"""
question_content = answer.get('content', {})
correct_answers = question_content.get('answer', []) # [[答案1, 同义词...], ...]
student_answers = answer.get('student_answer', [])
max_score = answer.get('max_score', 4.0)
if not correct_answers or not student_answers:
return {
"question_id": answer.get('question_id'),
"score": 0,
"max_score": max_score,
"grading_status": "failed",
"details": {"error": "答案格式错误correct_answers 或 student_answer 为空"}
}
# 计算每空分数
score_per_blank = max_score / len(correct_answers)
blank_scores = []
total_score = 0
for i, correct_list in enumerate(correct_answers):
if i >= len(student_answers):
blank_scores.append(0)
continue
student_ans = student_answers[i]
# 检查是否匹配任一正确答案
matched = any(
fuzzy_match(student_ans, correct)
for correct in correct_list
)
blank_score = score_per_blank if matched else 0
blank_scores.append(blank_score)
total_score += blank_score
return {
"question_id": answer.get('question_id'),
"score": round(total_score, 1),
"max_score": max_score,
"grading_status": "success",
"details": {
"blank_scores": blank_scores,
"total_blanks": len(correct_answers),
"correct_blanks": sum(1 for s in blank_scores if s > 0)
}
}
def fuzzy_match(student_answer: str, correct_answer: str) -> bool:
"""
模糊匹配(支持同义词和小编辑距离容错)
策略:
1. 精确匹配(去空格、转小写、统一标点)
2. 编辑距离容错≥4字答案允许≤2字符差异
"""
if not student_answer or not correct_answer:
return False
# 标准化:去空格、转小写、统一标点
def _normalize(text: str) -> str:
t = text.strip().lower()
# 统一常见中文标点变体
t = t.replace('', '(').replace('', ')').replace('', ',')
t = t.replace('', ';').replace('', ':').replace('"', '"').replace('"', '"')
return t
s = _normalize(student_answer)
c = _normalize(correct_answer)
if s == c:
return True
# 编辑距离容错答案≥4字时允许≤2字符差异
if len(s) >= 4 and len(c) >= 4:
dist = _edit_distance(s, c)
if dist <= 2:
return True
return False
def _edit_distance(s1: str, s2: str) -> int:
"""计算两个字符串的编辑距离Levenshtein"""
if len(s1) < len(s2):
return _edit_distance(s2, s1)
if len(s2) == 0:
return len(s1)
prev_row = list(range(len(s2) + 1))
for i, c1 in enumerate(s1):
curr_row = [i + 1]
for j, c2 in enumerate(s2):
# 插入、删除、替换
insertions = prev_row[j + 1] + 1
deletions = curr_row[j] + 1
substitutions = prev_row[j] + (c1 != c2)
curr_row.append(min(insertions, deletions, substitutions))
prev_row = curr_row
return prev_row[-1]
# ==================== AnswerGrader 类 ====================
class AnswerGrader:
"""本地批题器 - 使用本地 OpenAI 客户端"""
def __init__(self):
self.client = None
if LLM_AVAILABLE and API_KEY:
try:
from openai import OpenAI
self.client = OpenAI(api_key=API_KEY, base_url=BASE_URL)
except ImportError:
pass
self.model = MODEL
def grade_answers(self, answers: List[Dict]) -> List[Dict]:
"""
批阅答案列表
🔥 P1 改进:
- 本地批阅选择题/判断题
- 填空题本地模糊匹配
- 主观题调用 LLM 评分
- 结果顺序保持
"""
results_map = {}
# 分离题型
local_questions = [] # 选择题、判断题
fill_blank_questions = [] # 填空题
llm_questions = [] # 主观题
for ans in answers:
q_type = ans.get('question_type')
if q_type in ['single_choice', 'multiple_choice', 'true_false']:
local_questions.append(ans)
elif q_type == 'fill_blank':
fill_blank_questions.append(ans)
else:
llm_questions.append(ans)
# 本地批阅选择题/判断题
for ans in local_questions:
result = grade_objective(ans)
results_map[ans.get('question_id')] = result
# 本地批阅填空题
for ans in fill_blank_questions:
result = grade_fill_blank(ans)
results_map[ans.get('question_id')] = result
# 🔥 P1 改进:并发调用 LLM 批阅主观题
if llm_questions:
try:
self._grade_subjective_concurrently(llm_questions, results_map)
except Exception as e:
logger.error(f"主观题并发批阅整体异常: {e}")
# 兜底:为所有未完成的主观题设置失败状态
for ans in llm_questions:
qid = ans.get('question_id')
if qid not in results_map:
results_map[qid] = {
"question_id": qid,
"score": 0,
"max_score": ans.get('max_score', 10),
"grading_status": "failed",
"details": {"error": f"批阅系统异常: {str(e)}"}
}
# 🔥 P1 改进:按原始顺序重组结果
results = [results_map.get(ans.get('question_id')) for ans in answers]
return results
def _grade_subjective_concurrently(self, questions: List[Dict], results_map: Dict):
"""并发批阅主观题"""
with ThreadPoolExecutor(max_workers=MAX_CONCURRENT_GRADING) as executor:
# 建立映射关系
future_to_qid = {
executor.submit(self._grade_subjective, ans): ans.get('question_id')
for ans in questions
}
# 收集结果
try:
for future in as_completed(future_to_qid, timeout=60):
qid = future_to_qid[future]
try:
result = future.result(timeout=15)
results_map[qid] = result
except Exception as e:
# 失败时返回默认结果
results_map[qid] = {
"question_id": qid,
"score": 0,
"max_score": next(
(a.get('max_score', 10) for a in questions if a.get('question_id') == qid),
10
),
"grading_status": "failed",
"details": {"error": str(e)}
}
except TimeoutError:
# 超时:为未完成的任务设置失败状态
for future, qid in future_to_qid.items():
if qid not in results_map:
results_map[qid] = {
"question_id": qid,
"score": 0,
"max_score": next(
(a.get('max_score', 10) for a in questions if a.get('question_id') == qid),
10
),
"grading_status": "failed",
"details": {"error": "批阅超时"}
}
@retry(times=3, delay=1)
def _grade_subjective(self, answer: Dict) -> Dict:
"""
批阅主观题 - 调用 LLM 评分
🔥 改进3次重试 + 解析失败自动重试
"""
with grading_semaphore: # 限流
prompt = self._build_grading_prompt(answer)
response = self._call_llm(prompt)
try:
return self._parse_grading_result(response, answer)
except (json.JSONDecodeError, ValueError, TypeError) as e:
# 解析失败抛异常 → 触发 @retry 重试
logger.warning(f"[主观题评分] 解析失败将重试: {e}, 原始响应前200字: {str(response)[:200]}")
raise
def _build_grading_prompt(self, answer: Dict) -> str:
"""构造评分 Prompt"""
question_content = answer.get('content', {})
scoring_points = question_content.get('data', {}).get('scoring_points', [])
stem = question_content.get('stem', '')
reference_answer = question_content.get('answer', '')
# 如果缺少评分标准,在 prompt 中补充提示
scoring_section = ""
if scoring_points:
scoring_section = json.dumps(scoring_points, ensure_ascii=False, indent=2)
else:
scoring_section = "(未提供评分标准,请根据参考答案自行判断要点)"
return f"""请批阅以下简答题。
## 题目
{stem if stem else '(未提供题目)'}
## 参考答案
{reference_answer if reference_answer else '(未提供参考答案)'}
## 评分标准
{scoring_section}
## 学生答案
{answer.get('student_answer', '')}
## 满分
{answer.get('max_score', 10)}
## 输出约束
1. 必须输出合法 JSON不要包含任何占位符或中文说明
2. score 为数字,不能超过满分
3. achieved 为 0-1 之间的数字
4. 所有字段必须填入实际评分值
## 输出格式示例JSON
{{
"score": 7,
"scoring_breakdown": [
{{"point": "核心概念正确", "weight": 0.5, "achieved": 0.8, "comment": "基本概念描述准确"}}
],
"highlights": ["回答条理清晰"],
"shortcomings": ["缺少具体应用场景"],
"overall_feedback": "整体回答较好,但不够全面"
}}
请直接输出 JSON"""
def _call_llm(self, prompt: str) -> str:
"""调用本地 LLM"""
if not self.client:
raise ValueError("LLM 未配置,无法进行主观题评分")
messages = [
{"role": "system", "content": "你是一个专业的阅卷老师请严格按照JSON格式输出评分结果。"},
{"role": "user", "content": prompt}
]
result = call_llm(
client=self.client,
prompt=prompt,
model=self.model,
temperature=0.3,
max_tokens=_get_effective_max_tokens(2000, self.model),
messages=messages
)
if result is None:
raise Exception("LLM 调用失败")
return result
def _extract_json(self, response: str) -> dict:
"""多策略从 LLM 响应中提取 JSON 对象(使用共享工具)"""
result = extract_json_object(response)
if result is not None:
return result
if not response:
raise ValueError("LLM 返回为空")
raise json.JSONDecodeError(
f"无法从 LLM 响应中提取有效 JSON响应前300字: {response[:300]}",
response, 0
)
def _parse_grading_result(self, response: str, answer: Dict) -> Dict:
"""
解析评分结果
解析失败时抛出异常(由调用方的 @retry 处理重试)
"""
max_score = answer.get('max_score', 10)
# 检查主观题内容完整性
question_content = answer.get('content', {})
warnings = []
if not question_content.get('stem'):
warnings.append("缺少题目(stem)")
if not question_content.get('answer'):
warnings.append("缺少参考答案(answer)")
if not question_content.get('data', {}).get('scoring_points'):
warnings.append("缺少评分标准(scoring_points),评分结果仅供参考")
# 多策略提取 JSON失败抛异常 → 触发重试)
result = self._extract_json(response)
# 校验关键字段
score = result.get('score')
if score is None or not isinstance(score, (int, float)):
raise ValueError(f"score 字段缺失或类型错误: {score}")
score = min(float(score), max_score) # 不能超过满分
details = {
"scoring_breakdown": result.get('scoring_breakdown', []),
"highlights": result.get('highlights', []),
"shortcomings": result.get('shortcomings', []),
"overall_feedback": result.get('overall_feedback', '')
}
if warnings:
details["warnings"] = warnings
return {
"question_id": answer.get('question_id'),
"score": score,
"max_score": max_score,
"grading_status": "success",
"details": details
}
# ==================== 批题入口函数 ====================
def grade_answers(answers: List[Dict], request_id: str = None) -> Dict:
"""
批阅答案入口函数
🔥 P1/P2 改进:
- 加 timeout + retry
- 结果顺序保持
- 限流控制
Args:
answers: 答案列表
request_id: 请求 ID幂等性支持
Returns:
批阅结果
"""
grader = AnswerGrader()
results = grader.grade_answers(answers)
# 计算总分
total_score = sum(r.get('score', 0) for r in results if r)
total_max = sum(r.get('max_score', 0) for r in results if r)
return {
"success": True,
"request_id": request_id,
"results": results,
"total_score": round(total_score, 1),
"total_max_score": total_max,
"score_rate": round(total_score / total_max * 100, 1) if total_max > 0 else 0
}