Files
Video-Tree-TRM5/app/harness/pools.py
T
iomgaa ddb9a44f75 feat(pools): 三池切分以 unit 为原子,孪生对同池不被拆散
build_pools/_sample_excluding 与 PerCategoryPoolStrategy._split_one_category/
build_incremental 两条切分路径均改为以 QuestionUnit 为采样原子:progressive
exclusion 互斥集合与 train/val 分层划分都按 unit_id 计数(pair 计 1 个 unit),
命中单元整体展开,AR 孪生对两题永不落入不同池/split。

复用 app.harness.question_units 的 build_units/flatten_units,不重写分组逻辑。
single-only 输入下 unit 与 question 一一对应、rng 消耗量不变,采样与划分结果
与逐题口径完全一致;抽出 _unit_correct/_assert_correctness_complete 两个 helper
将 _split_one_category 复杂度压回基线以下。

新增 tests/unit/test_pools_pair_atomic.py 覆盖两条路径的 pair 原子性回归。
2026-07-15 06:09:14 -04:00

912 lines
34 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.
"""三池:held-out test + 验证 + 诊断,分层采样 + 冻结持久化。
三池切分对应训练循环中的 DataLoader 阶段——从题目全集中按
test -> validation -> diagnosis 的顺序 progressive exclusion
以 unit 为原子保证 unit_id 互斥(AR 孪生对两题永不被劈到不同池)。
test 池用自然分布(correct_ratio=None),验证池/诊断池按对错比例分层采样。
"""
from __future__ import annotations
import json
import math
import random
from collections import defaultdict
from dataclasses import dataclass, field
from typing import TYPE_CHECKING
from loguru import logger
from app.harness.question_units import build_units, flatten_units
from app.question_gen import stratified_sample
from core.types import GeneratedQuestion, PoolConfig
if TYPE_CHECKING:
from pathlib import Path
from app.harness.config import RunConfig
from app.ports import PoolStrategy
from core.types import QuestionUnit
@dataclass
class Pools:
"""冻结的三池及其基线指标。
字段:
diagnosis: 诊断池(用于错误归因,对应 loss.backward)。
validation: 验证池(按类局部验证,每题型有保底样本)。
test: held-out 测试池(自然分布,用于最终无偏评估)。
baseline_run_id: 基线 run 标识。
baseline_val_accuracy: 基线在验证池上的准确率。
correctness: 三池所有题的 question_id -> 基线是否答对。
"""
diagnosis: list[GeneratedQuestion]
validation: list[GeneratedQuestion]
test: list[GeneratedQuestion]
baseline_run_id: str
baseline_val_accuracy: float
correctness: dict[str, bool] = field(default_factory=dict)
def build_pools(
questions: list[GeneratedQuestion],
correctness: dict[str, bool],
diag_cfg: dict,
val_cfg: dict,
test_cfg: dict,
baseline_run_id: str,
) -> Pools:
"""先抽 held-out test,再抽验证集,最后抽诊断池,三池互斥。
参数:
questions: 题目全集。
correctness: question_id -> 基线是否答对。
diag_cfg: 诊断池采样配置(size/correct_ratio/task_types[/seed])。
val_cfg: 验证池采样配置,可含 min_per_class 做按类保底。
test_cfg: 测试池配置(size[/seed]);走自然分布,不强制对错比与题型。
baseline_run_id: 基线 run 标识。
返回:
冻结的三池 Pools。
关键实现细节:
切分顺序 test -> validation -> diagnosis;后两步从剩余单元中采样以保证
unit_id 互斥。test 池用 correct_ratio=None 的自然分布采样。以 unit 为采样
原子(pair 计 1 个 unit),孪生对两题永不被劈到不同池;size/correct_ratio
按 unit 计数,single-only 输入下 unit 与 question 一一对应,行为完全不变。
"""
units = build_units(questions)
test = _sample_excluding(
units,
set(),
correctness,
size=test_cfg["size"],
correct_ratio=None,
task_types=None,
seed=test_cfg.get("seed", 0),
min_per_class=None,
)
selected_units = {q.unit_id for q in test}
validation = _sample_excluding(units, selected_units, correctness, **val_cfg)
selected_units |= {q.unit_id for q in validation}
diagnosis = _sample_excluding(units, selected_units, correctness, **diag_cfg)
val_correct = sum(1 for q in validation if correctness.get(q.question_id))
baseline_val_accuracy = val_correct / len(validation) if validation else 0.0
return Pools(
diagnosis=diagnosis,
validation=validation,
test=test,
baseline_run_id=baseline_run_id,
baseline_val_accuracy=baseline_val_accuracy,
correctness={
q.question_id: correctness.get(q.question_id, False)
for q in test + validation + diagnosis
},
)
class GlobalPoolStrategy:
"""全局三分策略:test -> val -> diag progressive exclusion。
封装现有 build_pools 逻辑为 PoolStrategy 接口。
"""
def build(
self,
questions: list[GeneratedQuestion],
correctness: dict[str, bool],
config: PoolConfig,
*,
db_path: Path | None = None,
) -> Pools:
"""委托给现有 build_pools 函数。
参数:
questions: 题目全集。
correctness: question_id -> 基线是否答对。
config: 池构建统一配置。
返回:
冻结的三池 Pools。
"""
return build_pools(
questions,
correctness,
diag_cfg={
"size": config.diag_size,
"correct_ratio": config.diag_correct_ratio,
"task_types": list(config.task_types) if config.task_types else None,
"seed": config.seed,
"min_per_class": None,
},
val_cfg={
"size": config.val_size,
"correct_ratio": config.val_correct_ratio,
"task_types": list(config.task_types) if config.task_types else None,
"seed": config.seed,
"min_per_class": config.eval_min_per_class,
},
test_cfg={"size": config.test_size, "seed": config.seed},
baseline_run_id=config.baseline_run_id,
)
def build_incremental(
self,
new_task_types: list[str],
questions: list[GeneratedQuestion],
correctness: dict[str, bool],
config: PoolConfig,
) -> dict[str, dict[str, list[str]]]:
"""全局策略不支持增量。
异常:
NotImplementedError: 始终抛出。
"""
raise NotImplementedError(
"GlobalPoolStrategy 不支持增量构建,请使用 PerCategoryPoolStrategy。"
)
def _unit_correct(unit: QuestionUnit, correctness: dict[str, bool]) -> bool:
"""单元级正确性:成员全部答对才算对(缺失按 False,宽松口径)。
参数:
unit: 目标单元(single 1 题,pair 2 题)。
correctness: question_id -> 基线是否答对。
返回:
pair 走双向 AND、single 即单题正确性;任一成员缺失或答错即 False。
"""
return all(correctness.get(q.question_id, False) for q in unit.questions)
def _assert_correctness_complete(
units: list[QuestionUnit],
correctness: dict[str, bool],
) -> None:
"""校验 correctness 覆盖所有单元成员题(含 pair 两题),缺失即 fail-fast。
参数:
units: 待校验单元列表。
correctness: question_id -> 基线是否答对。
异常:
ValueError: correctness 中缺少某些 question_id。
"""
missing = [
q.question_id for u in units for q in u.questions if q.question_id not in correctness
]
if missing:
raise ValueError(f"correctness 缺失 {len(missing)} 题: {missing[:5]}")
def _sample_excluding(
units: list[QuestionUnit],
exclude_unit_ids: set[str],
correctness: dict[str, bool],
**cfg: object,
) -> list[GeneratedQuestion]:
"""排除已选 unit 后,以 unit 为原子按 cfg 分层采样,返回展开后的逐题列表。
每个单元以其首题作为分层采样的代表参与 stratified_samplecorrect_ratio /
size 因此按 unit 计数(pair 计 1 个 unit);命中的单元整体展开,孪生对两题
永远同进同出。single-only 输入下 unit 与 question 一一对应、顺序不变,采样
结果与逐题采样完全一致。
参数:
units: 单元全集(single 单封、pair 成对聚合)。
exclude_unit_ids: 已被其他池选走的 unit_id,从候选中剔除以保证三池互斥。
correctness: question_id -> 基线是否答对;单元级正确性取成员的 AND
(缺失按 False,与 stratified_sample 的宽松口径一致)。
cfg: 透传给 stratified_sample 的采样配置
size/correct_ratio/task_types[/seed/min_per_class])。
返回:
采样命中单元展开后的题目列表。
"""
candidates = [u for u in units if u.unit_id not in exclude_unit_ids]
rep_to_unit = {u.questions[0].question_id: u for u in candidates}
reps = [u.questions[0] for u in candidates]
unit_correct = {u.questions[0].question_id: _unit_correct(u, correctness) for u in candidates}
sampled_reps = stratified_sample(reps, unit_correct, **cfg)
sampled_units = [rep_to_unit[rep.question_id] for rep in sampled_reps]
return flatten_units(sampled_units)
def _q_to_dict(q: GeneratedQuestion) -> dict:
"""将 GeneratedQuestion 转为可序列化字典。
参数:
q: 题目对象。
返回:
包含全部字段的字典(options/source_nodes 从 tuple 转为 list)。
"""
return {
"question_id": q.question_id,
"video_id": q.video_id,
"task_type": q.task_type,
"question": q.question,
"options": list(q.options),
"answer": q.answer,
"source_nodes": list(q.source_nodes),
"difficulty": q.difficulty,
"family": q.family,
"skill_target": q.skill_target,
"difficulty_steps": q.difficulty_steps,
}
def _dict_to_q(d: dict) -> GeneratedQuestion:
"""从字典恢复 GeneratedQuestion。
参数:
d: 由 _q_to_dict 产出的字典。
返回:
恢复的 GeneratedQuestion 实例(options/source_nodes 恢复为 tuple)。
"""
return GeneratedQuestion(
question_id=d["question_id"],
video_id=d["video_id"],
task_type=d["task_type"],
question=d["question"],
options=tuple(d["options"]),
answer=d["answer"],
source_nodes=tuple(d.get("source_nodes", ())),
difficulty=d.get("difficulty", "medium"),
family=d.get("family"),
skill_target=d.get("skill_target"),
difficulty_steps=d.get("difficulty_steps"),
)
def save_pools(
pools: Pools,
path: Path,
*,
split_mode: str = "global",
config: PoolConfig | None = None,
) -> None:
"""将三池及基线指标冻结为 JSON。
参数:
pools: 待冻结的三池。
path: 目标 JSON 文件路径。
split_mode: 池划分策略标记("global" / "per_category"),写入 JSON 用于
加载时识别格式。
config: 池构建配置。per_category 模式下必须提供,用于写入 categories
元数据(seed, train_ratio, test_source)以支持增量追加和一致性校验。
异常:
ValueError: split_mode 为 "per_category" 但未提供 config。
"""
if split_mode == "per_category" and config is None:
raise ValueError("per_category 模式下 save_pools 必须提供 config 参数以写入元数据。")
data: dict = {
"split_mode": split_mode,
"baseline_run_id": pools.baseline_run_id,
"baseline_val_accuracy": pools.baseline_val_accuracy,
"correctness": pools.correctness,
"diagnosis": [_q_to_dict(q) for q in pools.diagnosis],
"validation": [_q_to_dict(q) for q in pools.validation],
"test": [_q_to_dict(q) for q in pools.test],
}
if split_mode == "per_category" and config is not None:
# 按 task_type 记录 train/val 的 qid 列表,用于增量追加和一致性校验
categories: dict[str, dict[str, list[str]]] = {}
diag_by_type: dict[str, list[str]] = defaultdict(list)
val_by_type: dict[str, list[str]] = defaultdict(list)
for q in pools.diagnosis:
diag_by_type[q.task_type].append(q.question_id)
for q in pools.validation:
val_by_type[q.task_type].append(q.question_id)
for task_type in sorted(set(diag_by_type) | set(val_by_type)):
categories[task_type] = {
"train": diag_by_type.get(task_type, []),
"val": val_by_type.get(task_type, []),
}
data["categories"] = categories
data["seed"] = config.seed
data["train_ratio"] = config.train_ratio
data["test_source"] = str(config.test_questions_dir) if config.test_questions_dir else None
path.write_text(
json.dumps(data, ensure_ascii=False, indent=2),
encoding="utf-8",
)
def load_pools(path: Path) -> Pools:
"""从 JSON 恢复冻结的三池。
兼容新旧格式:有无 split_mode 字段都能加载。新格式(含 split_mode /
categories)的额外元数据在加载时忽略——Pools 对象只关心三池列表和标量。
参数:
path: 冻结的 pools.json 路径。
返回:
恢复的三池 Pools。
异常:
ValueError: 旧格式 pools.json(无 test 池)。
关键实现细节:
旧格式 pools.json(无 test 池)会以清晰的 ValueError 中止——本项目不做
向后兼容,也不为缺失字段填默认值。删除旧文件后 build_pools 会重新采样切分,
无需重新推理。
"""
d = json.loads(path.read_text(encoding="utf-8"))
if "test" not in d:
raise ValueError(
f"{path} 为旧格式 pools.json(缺 test 池),"
"请删除后重新切分(build_pools 会重新采样,无需重新推理)。"
)
return Pools(
diagnosis=[_dict_to_q(x) for x in d["diagnosis"]],
validation=[_dict_to_q(x) for x in d["validation"]],
test=[_dict_to_q(x) for x in d["test"]],
baseline_run_id=d["baseline_run_id"],
baseline_val_accuracy=d["baseline_val_accuracy"],
correctness=d["correctness"],
)
def _to_pool_config(config: RunConfig, baseline_run_id: str) -> PoolConfig:
"""从 RunConfig + 外部 baseline_run_id 提取 PoolConfig。
baseline_run_id 必须由调用方从 workspace manifest / seed.json 读取,
绝不能用 config.run_id(那是训练 run ID)。
参数:
config: 运行配置。
baseline_run_id: 基线 run 标识(来自 workspace manifest 或 seed.json)。
返回:
PoolConfig 实例。
"""
test_questions_dir: Path | None = None
if config.test_questions:
from app.harness.workspace import resolve_paths
paths = resolve_paths(config.workspace_dir)
test_questions_dir = paths.store_dir / "questions" / config.test_questions
return PoolConfig(
task_types=config.task_types,
seed=0,
baseline_run_id=baseline_run_id,
diag_size=config.diag_size,
diag_correct_ratio=config.diag_correct_ratio,
val_size=config.val_size,
val_correct_ratio=config.val_correct_ratio,
test_size=config.test_size,
eval_min_per_class=config.eval_min_per_class,
train_ratio=config.train_ratio,
test_questions_dir=test_questions_dir,
batch_correct_ratio=config.batch_correct_ratio,
)
def _read_baseline_run_id(config: RunConfig) -> str:
"""从 workspace 的 seed.json 读取 baseline_run_id。
workspace 由 init_workspace_from_seed 从种子创建,seed.json 保存在
store/seeds/<name>/seed.json 中。manifest.json 中 history 首条或 seed
配置字段指向对应种子。
参数:
config: 运行配置(提供 workspace_dir, store_dir, seed)。
返回:
baseline_run_id 字符串。
"""
from app.harness.store import read_seed
meta = read_seed(config.store_dir, config.seed)
return meta["baseline_run_id"]
def _validate_per_category_consistency(
frozen_data: dict,
pool_config: PoolConfig,
baseline_run_id: str,
) -> None:
"""校验已冻结的 per_category pools.json 与当前配置的一致性。
参数:
frozen_data: pools.json 解析后的原始字典。
pool_config: 当前构建配置。
baseline_run_id: 当前基线 run 标识。
异常:
ValueError: 任一关键参数与冻结值不一致。
"""
mismatches: list[str] = []
if frozen_data.get("seed") != pool_config.seed:
mismatches.append(f"seed: 冻结={frozen_data.get('seed')}, 当前={pool_config.seed}")
if frozen_data.get("train_ratio") != pool_config.train_ratio:
mismatches.append(
f"train_ratio: 冻结={frozen_data.get('train_ratio')}, 当前={pool_config.train_ratio}"
)
if frozen_data.get("baseline_run_id") != baseline_run_id:
mismatches.append(
f"baseline_run_id: 冻结={frozen_data.get('baseline_run_id')}, 当前={baseline_run_id}"
)
if frozen_data.get("split_mode") != "per_category":
mismatches.append(f"split_mode: 冻结={frozen_data.get('split_mode')}, 当前=per_category")
if mismatches:
raise ValueError(
"per_category pools.json 与当前配置不一致:\n"
+ "\n".join(f" - {m}" for m in mismatches)
)
def build_or_load_pools(
config: RunConfig,
strategy: PoolStrategy,
db_path: Path,
) -> Pools:
"""train 模式的三池获取入口:pools.json 已存在则加载,否则从基线 db 切分并冻结。
把 main.py train 分支「pools.json 存在则 load_pools 否则 build_pools 再 save_pools」
那段抽成纯函数,使 main 与集成测试共用同一切分逻辑、避免重复。pools.json 是
一次 fresh 训练的冻结切分,resume/重跑同一 workspace 时直接复用以保证三池一致。
参数:
config: 运行配置,提供 workspace_dir 与三池采样旋钮(diag/val/test 各项)。
strategy: 池构建策略(GlobalPoolStrategy / PerCategoryPoolStrategy)。
db_path: harness.db 路径,用于读取基线推理对错。
返回:
冻结的三池 Pools。
关键实现:
baseline_run_id 从 seed.json 读取(非 config.run_id)。per_category 模式
加载时做一致性校验,并支持新类别的增量追加。切分前从基线 db 的 predictions
表读该 run_id 的逐题对错,作为分层采样依据。pools.json 落在
config.workspace_dir 下,存在即视为已冻结。
"""
from app.harness.log import HarnessLog
from app.harness.workspace import resolve_paths
from app.question_gen import load_benchmark
baseline_run_id = _read_baseline_run_id(config)
pool_config = _to_pool_config(config, baseline_run_id)
pools_path = config.workspace_dir / "pools.json"
if pools_path.exists():
# ── 加载已冻结的 pools ──
raw = json.loads(pools_path.read_text(encoding="utf-8"))
frozen_split_mode = raw.get("split_mode", "global")
if frozen_split_mode == "per_category":
_validate_per_category_consistency(raw, pool_config, baseline_run_id)
# 检查是否有新类别需要增量追加
frozen_categories = raw.get("categories", {})
if pool_config.task_types is not None:
requested_types = set(pool_config.task_types)
existing_types = set(frozen_categories.keys())
new_types = requested_types - existing_types
if new_types:
# 增量构建新类别
paths = resolve_paths(config.workspace_dir)
questions = load_benchmark(paths.questions_dir)
with HarnessLog(str(db_path), baseline_run_id) as hlog:
rows = hlog.query(
"SELECT question_id, prediction, answer "
"FROM predictions WHERE run_id=?",
(baseline_run_id,),
)
correctness = {r["question_id"]: r["prediction"] == r["answer"] for r in rows}
new_cats = strategy.build_incremental(
sorted(new_types),
questions,
correctness,
pool_config,
)
# 合并新类别到 categories
frozen_categories.update(new_cats)
raw["categories"] = frozen_categories
# 从 categories 重建 diagnosis/validation 列表
qid_map = {q.question_id: q for q in questions}
new_diag: list[dict] = []
new_val: list[dict] = []
for tt in sorted(frozen_categories.keys()):
cat = frozen_categories[tt]
for qid in cat["train"]:
if qid in qid_map:
new_diag.append(_q_to_dict(qid_map[qid]))
for qid in cat["val"]:
if qid in qid_map:
new_val.append(_q_to_dict(qid_map[qid]))
raw["diagnosis"] = new_diag
raw["validation"] = new_val
raw["correctness"] = {
**raw.get("correctness", {}),
**{
qid: correctness.get(qid, False)
for cat in new_cats.values()
for qid in cat["train"] + cat["val"]
},
}
# 重新冻结
pools_path.write_text(
json.dumps(raw, ensure_ascii=False, indent=2),
encoding="utf-8",
)
logger.info(
"per_category 增量追加 {} 个新类别: {}",
len(new_types),
sorted(new_types),
)
return load_pools(pools_path)
# ── 全新构建 ──
paths = resolve_paths(config.workspace_dir)
questions = load_benchmark(paths.questions_dir)
with HarnessLog(str(db_path), baseline_run_id) as hlog:
rows = hlog.query(
"SELECT question_id, prediction, answer FROM predictions WHERE run_id=?",
(baseline_run_id,),
)
correctness = {r["question_id"]: r["prediction"] == r["answer"] for r in rows}
pools = strategy.build(questions, correctness, pool_config, db_path=db_path)
save_pools(
pools,
pools_path,
split_mode=config.pool_split_mode,
config=pool_config,
)
return pools
class PerCategoryPoolStrategy:
"""Per-category 分层池构建策略。
按题型分组,每个题型内部按 correctness 分层,以 train_ratio 比例
划分 train(映射到 diagnosis 池)和 val(映射到 validation 池)。
与 GlobalPoolStrategy 的全局 progressive exclusion 不同,本策略
保证每个类别内部的 train/val 比例精确对齐。
"""
def build(
self,
questions: list[GeneratedQuestion],
correctness: dict[str, bool],
config: PoolConfig,
*,
db_path: Path | None = None,
) -> Pools:
"""按题型分组后,每组做 correctness 分层的 train/val 划分。
参数:
questions: 题目全集。
correctness: question_id -> 基线是否答对。
config: 池构建配置(使用 train_ratio, task_types, seed,
baseline_run_id, test_questions_dir, batch_correct_ratio)。
db_path: harness.db 路径,用于查询 benchmark 历史推理记录
maintenance 补入时需要)。
返回:
冻结的 Poolsdiagnosis=train, validation=val,
test 从 test_questions_dir 加载或为空列表)。
"""
# Phase 1: 按 task_types 过滤
if config.task_types is not None:
allowed = set(config.task_types)
filtered = [q for q in questions if q.task_type in allowed]
else:
filtered = list(questions)
# Phase 2: 按 task_type 分组
groups: dict[str, list[GeneratedQuestion]] = defaultdict(list)
for q in filtered:
groups[q.task_type].append(q)
# Phase 2.5: 正确率检查 + maintenance 补入
if config.batch_correct_ratio is not None:
self._check_and_supplement_maintenance(
groups,
correctness,
config,
db_path,
)
# Phase 3: 每组分层划分
all_train: list[GeneratedQuestion] = []
all_val: list[GeneratedQuestion] = []
rng = random.Random(config.seed)
for task_type in sorted(groups.keys()):
train_units, val_units = self._split_one_category(
build_units(groups[task_type]),
correctness,
config.train_ratio,
rng,
)
all_train.extend(flatten_units(train_units))
all_val.extend(flatten_units(val_units))
# Phase 4: test 池(从外部目录加载,无则空;按 task_types 过滤)
test: list[GeneratedQuestion] = []
if config.test_questions_dir is not None:
from app.question_gen import load_benchmark
test = load_benchmark(config.test_questions_dir)
if config.task_types is not None:
allowed = set(config.task_types)
test = [q for q in test if q.task_type in allowed]
# Phase 5: 计算 baseline_val_accuracy
val_correct = sum(1 for q in all_val if correctness.get(q.question_id, False))
baseline_val_accuracy = val_correct / len(all_val) if all_val else 0.0
return Pools(
diagnosis=all_train,
validation=all_val,
test=test,
baseline_run_id=config.baseline_run_id,
baseline_val_accuracy=baseline_val_accuracy,
correctness={
q.question_id: correctness.get(q.question_id, False)
for q in all_train + all_val + test
},
)
def _check_and_supplement_maintenance(
self,
groups: dict[str, list[GeneratedQuestion]],
correctness: dict[str, bool],
config: PoolConfig,
db_path: Path | None,
) -> None:
"""按 task_type 检查正确率,过高警告,过低则从 benchmark 补入正确题。
修改 groups 和 correctness(原地更新)。
参数:
groups: task_type -> 题目列表映射(原地追加补入题)。
correctness: question_id -> 是否正确映射(原地追加补入题标记)。
config: 含 batch_correct_ratio 和 test_questions_dir。
db_path: harness.db 路径,用于查询 benchmark 历史推理记录。
"""
import sqlite3
r = config.batch_correct_ratio
assert r is not None # 调用方已保证
# Phase 2.5a: 正确率检查(不依赖 test_questions_dir
for task_type, group in groups.items():
c = sum(1 for q in group if correctness.get(q.question_id, False))
n = len(group)
ratio = c / n if n > 0 else 0.0
if ratio > 1 - r:
logger.warning(
"类别 {} 正确率 {:.1%} 过高(阈值 {:.1%}),出题可能太简单",
task_type,
ratio,
1 - r,
)
# Phase 2.5b: maintenance 补入(需要 test_questions_dir
if config.test_questions_dir is None:
return
from app.question_gen import load_benchmark
bench_questions = load_benchmark(config.test_questions_dir)
# 查询 DB 中 benchmark 题的历史正确性
bench_correctness: dict[str, bool] = {}
if db_path is not None and db_path.exists():
conn = sqlite3.connect(str(db_path))
bench_qids = [q.question_id for q in bench_questions]
if bench_qids:
placeholders = ",".join("?" for _ in bench_qids)
rows = conn.execute(
f"SELECT question_id, prediction, answer FROM predictions " # noqa: S608
f"WHERE question_id IN ({placeholders}) "
f"ORDER BY timestamp DESC",
bench_qids,
).fetchall()
for qid, pred, ans in rows:
if qid not in bench_correctness:
bench_correctness[qid] = pred == ans
conn.close()
# 按 task_type 索引 benchmark 题
bench_by_type: dict[str, list[GeneratedQuestion]] = defaultdict(list)
for q in bench_questions:
bench_by_type[q.task_type].append(q)
for task_type, group in groups.items():
c = sum(1 for q in group if correctness.get(q.question_id, False))
w = len(group) - c
n = len(group)
ratio = c / n if n > 0 else 0.0
if ratio >= r:
continue
k = math.ceil((r * w - (1 - r) * c) / (1 - r))
existing_ids = {q.question_id for q in group}
candidates = [
q
for q in bench_by_type.get(task_type, [])
if bench_correctness.get(q.question_id, False) and q.question_id not in existing_ids
]
if not candidates:
logger.warning(
"类别 {} 需补入 {} 道正确题,但 benchmark 中无可用候选",
task_type,
k,
)
continue
actual = min(k, len(candidates))
for q in candidates[:actual]:
supplemented = GeneratedQuestion(
question_id=q.question_id,
video_id=q.video_id,
task_type=q.task_type,
question=q.question,
options=q.options,
answer=q.answer,
source_nodes=q.source_nodes,
difficulty=q.difficulty,
family="VME_MAINTENANCE",
skill_target=q.skill_target,
difficulty_steps=q.difficulty_steps,
)
group.append(supplemented)
correctness[supplemented.question_id] = True
logger.info(
"类别 {} 正确率 {:.1%} < {:.1%},从 benchmark 补入 {} 道 maintenance 正确题",
task_type,
ratio,
r,
actual,
)
def _split_one_category(
self,
units: list[QuestionUnit],
correctness: dict[str, bool],
train_ratio: float,
rng: random.Random,
) -> tuple[list[QuestionUnit], list[QuestionUnit]]:
"""单类别 correctness 分层划分,以 unit 为原子(pair 计 1 个 unit)。
孪生对两题作为一个整体落入 train 或 val,绝不被拆散;single-only 输入下
unit 与 question 一一对应、rng 消耗量不变,划分结果与逐题划分完全一致。
参数:
units: 单类别全部单元。
correctness: question_id -> 基线是否答对;单元级正确性取成员的 AND。
train_ratio: train 占单元总量的比例。
rng: 随机数生成器(保证跨类别可复现)。
返回:
(train_units, val_units) 单元列表元组,两侧互斥且总量 == len(units)。
异常:
ValueError: correctness 中缺少某些 question_id。
"""
n_total = len(units)
n_train = round(n_total * train_ratio)
n_val = n_total - n_train
_assert_correctness_complete(units, correctness)
correct_units = [u for u in units if _unit_correct(u, correctness)]
wrong_units = [u for u in units if not _unit_correct(u, correctness)]
n_correct = len(correct_units)
# 全 correct 或全 wrong -> 退化为非分层随机划分
if n_correct == 0 or n_correct == n_total:
label = "全部正确" if n_correct == n_total else "全部错误"
logger.warning(
"类别 {} {} ({} 单元),退化为非分层随机划分",
units[0].task_type,
label,
n_total,
)
shuffled = list(units)
rng.shuffle(shuffled)
return shuffled[:n_train], shuffled[n_train:]
# 分层: 按 correctness 比例分配到 train
train_correct = math.floor(n_correct * n_train / n_total)
train_wrong = n_train - train_correct
rng.shuffle(correct_units)
rng.shuffle(wrong_units)
train = correct_units[:train_correct] + wrong_units[:train_wrong]
val = correct_units[train_correct:] + wrong_units[train_wrong:]
assert len(train) == n_train, f"train 数量不匹配: {len(train)} != {n_train}"
assert len(val) == n_val, f"val 数量不匹配: {len(val)} != {n_val}"
return train, val
def build_incremental(
self,
new_task_types: list[str],
questions: list[GeneratedQuestion],
correctness: dict[str, bool],
config: PoolConfig,
) -> dict[str, dict[str, list[str]]]:
"""增量划分:仅处理 new_task_types 中的类别。
参数:
new_task_types: 需要增量划分的类别列表。
questions: 题目全集(从中筛选指定类别)。
correctness: question_id -> 基线是否答对。
config: 池构建配置(使用 train_ratio, seed)。
返回:
{task_type: {"train": [qid, ...], "val": [qid, ...]}}。
"""
target_types = set(new_task_types)
groups: dict[str, list[GeneratedQuestion]] = defaultdict(list)
for q in questions:
if q.task_type in target_types:
groups[q.task_type].append(q)
result: dict[str, dict[str, list[str]]] = {}
rng = random.Random(config.seed)
for task_type in sorted(groups.keys()):
train_units, val_units = self._split_one_category(
build_units(groups[task_type]),
correctness,
config.train_ratio,
rng,
)
result[task_type] = {
"train": [q.question_id for q in flatten_units(train_units)],
"val": [q.question_id for q in flatten_units(val_units)],
}
return result