feat: add question_units helper as pair contract entry point

build_units/flatten_units/validate_units/unit_correctness——pair 契约唯一入口。
build_units 对孤儿/超员/角色缺失重复 fail-fast raise ValueError(防 next 静默
StopIteration);unit_correctness 走 per_q[qid] KeyError 防静默兜底。
This commit is contained in:
2026-07-15 05:52:30 -04:00
parent 7ef9b99217
commit bef46636fe
2 changed files with 170 additions and 0 deletions
+97
View File
@@ -0,0 +1,97 @@
"""QuestionUnit 组装/展开/校验/单元正确性——pair 契约的唯一入口。
pool 构建、批处理、推理、评测(Task 3+)均通过本模块聚合/展开孪生对,
保证 AR pair 的"两题作为整体调度"契约只在一处实现、fail-fast 暴露非法配对。
"""
from collections import defaultdict
from core.types import GeneratedQuestion, QuestionUnit
def build_units(questions: list[GeneratedQuestion]) -> list[QuestionUnit]:
"""将扁平题目列表聚合为单元列表:single 单封、pair 按 pair_id 成对聚合。
参数:
questions: 待聚合的题目列表,可混含 single 与孪生对成员。
返回:
单元列表,先 single 后 pair,顺序稳定(single 保留输入顺序,
pair 按首次出现的 pair_id 顺序)。
关键实现:
- pair_id 为空 → single 单元;非空 → 归入对应 pair 桶。
- 每个 pair 桶必须恰好 2 条,否则视为孤儿/超员,raise ValueError。
- 显式检查 pair_original / pair_mirror 角色齐备且唯一,缺失或重复
直接 raise ValueError(防 next(...) 静默 StopIteration),保持
fail-fast。合法孪生对交由 QuestionUnit.from_pair 做一致性断言。
"""
by_pair: dict[str, list[GeneratedQuestion]] = defaultdict(list)
singles: list[QuestionUnit] = []
for q in questions:
if q.pair_id:
by_pair[q.pair_id].append(q)
else:
singles.append(QuestionUnit.from_single(q))
pairs: list[QuestionUnit] = []
for pid, qs in by_pair.items():
if len(qs) != 2:
raise ValueError(f"pair {pid} 数量={len(qs)}≠2(孤儿或超员)")
originals = [q for q in qs if q.question_role == "pair_original"]
mirrors = [q for q in qs if q.question_role == "pair_mirror"]
if len(originals) != 1 or len(mirrors) != 1:
raise ValueError(
f"pair {pid} 角色非法:original={len(originals)} mirror={len(mirrors)}"
"需各恰好 1 条"
)
pairs.append(QuestionUnit.from_pair(originals[0], mirrors[0]))
return singles + pairs
def flatten_units(units: list[QuestionUnit]) -> list[GeneratedQuestion]:
"""将单元列表无损展开回扁平题目列表。
参数:
units: 单元列表。
返回:
展开后的题目列表,保持单元顺序及单元内题目顺序。
"""
return [q for u in units for q in u.questions]
def validate_units(units: list[QuestionUnit]) -> list[QuestionUnit]:
"""校验单元列表结构合法性,通过则原样返回(便于链式调用)。
参数:
units: 待校验单元列表。
返回:
校验通过的原单元列表。
关键实现:
pair 单元必须恰好含 2 题,否则 raise ValueErrorsingle 单元无需额外
校验(构造时即为 1 题)。用于消费方在使用前做一道防御闸门。
"""
for u in units:
if u.kind == "pair" and u.size != 2:
raise ValueError(f"unit {u.unit_id} pair 不成对(size={u.size}")
return units
def unit_correctness(unit: QuestionUnit, per_q: dict[str, bool]) -> bool:
"""计算单元级正确性:AR pair 走双向 AND,single 即单题正确性。
参数:
unit: 目标单元。
per_q: 题目 question_id → 该题是否作答正确的映射。
返回:
单元内所有题目均正确时为 True,否则 False。
关键实现:
直接以 per_q[q.question_id] 取值,缺任一题触发 KeyError(防静默兜底),
强制上游先补齐全部单题结果再计单元正确性。
"""
return all(per_q[q.question_id] for q in unit.questions)