Files
Video-Tree-TRM5/app/question_gen/run_store.py
T

472 lines
15 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.
"""SQLite 出题管线日志记录器 — 记录批次运行与逐题门判定。
写入频率低(每批次 ~240 题),使用同步 sqlite3 即可。
schema 定义参照 research-wiki/schemas/question-gen-runs.md 和 question-gen-items.md。
"""
from __future__ import annotations
import sqlite3
from dataclasses import dataclass
from datetime import UTC, datetime
from typing import TYPE_CHECKING, Protocol, runtime_checkable
from loguru import logger
if TYPE_CHECKING:
from pathlib import Path
# ---------------------------------------------------------------------------
# GateReport ProtocolTask 6 尚未实现,此处声明 duck-type 接口)
# ---------------------------------------------------------------------------
@runtime_checkable
class _GateResultLike(Protocol):
"""单门判定结果的最小接口。"""
@property
def verdict(self) -> object:
"""PASS / FAIL / SKIP 枚举值,.value 为小写字符串。"""
...
@property
def reason(self) -> str:
"""判定理由。"""
...
@runtime_checkable
class GateReportLike(Protocol):
"""GateReport 的最小 duck-type 接口。"""
@property
def key_verify(self) -> _GateResultLike: ...
@property
def blind_answer(self) -> _GateResultLike: ...
@property
def multi_true(self) -> _GateResultLike: ...
@property
def leak_test(self) -> _GateResultLike: ...
@property
def passed(self) -> bool: ...
@property
def reject_reason(self) -> str | None: ...
# ---------------------------------------------------------------------------
# 数据类
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class RunStats:
"""批次运行的统计摘要。
Attributes
----------
total_slots : int
目标题数(slot 总数)。
accepted : int
最终通过门判定的题数。
rejected : int
最终被拒绝的题数。
heavy_sampled : int
进行重量抽检的题数。
"""
total_slots: int
accepted: int
rejected: int
heavy_sampled: int
# ---------------------------------------------------------------------------
# DDL
# ---------------------------------------------------------------------------
_DDL_RUNS = """
CREATE TABLE IF NOT EXISTS question_gen_runs (
run_id TEXT PRIMARY KEY,
git_sha TEXT NOT NULL,
config_snapshot TEXT,
started_at TEXT NOT NULL,
ended_at TEXT,
status TEXT NOT NULL DEFAULT 'running',
total_slots INTEGER,
accepted INTEGER,
rejected INTEGER,
heavy_sampled INTEGER
);
"""
_DDL_ITEMS = """
CREATE TABLE IF NOT EXISTS question_gen_items (
item_id TEXT PRIMARY KEY,
run_id TEXT NOT NULL REFERENCES question_gen_runs(run_id),
slot_id TEXT NOT NULL,
video_id TEXT NOT NULL,
family TEXT NOT NULL,
task_type TEXT NOT NULL,
skill_target TEXT NOT NULL,
attempt INTEGER NOT NULL,
question_text TEXT NOT NULL,
sub_pattern TEXT,
gate_key_verify TEXT,
gate_blind_answer TEXT,
gate_multi_true TEXT,
gate_leak_test TEXT,
gate_reject_reason TEXT,
final_status TEXT NOT NULL DEFAULT 'pending',
difficulty_steps INTEGER,
selector_scores TEXT,
created_at TEXT NOT NULL DEFAULT (datetime('now'))
);
"""
_DDL_INDEXES = [
"CREATE INDEX IF NOT EXISTS idx_qgr_status ON question_gen_runs(status);",
"CREATE INDEX IF NOT EXISTS idx_qgi_run ON question_gen_items(run_id);",
"CREATE INDEX IF NOT EXISTS idx_qgi_slot ON question_gen_items(slot_id);",
"CREATE INDEX IF NOT EXISTS idx_qgi_status ON question_gen_items(final_status);",
"CREATE INDEX IF NOT EXISTS idx_qgi_family ON question_gen_items(family);",
"CREATE INDEX IF NOT EXISTS idx_qgi_task_type ON question_gen_items(task_type);",
]
# ---------------------------------------------------------------------------
# Store 实现
# ---------------------------------------------------------------------------
class QuestionGenStore:
"""SQLite 出题管线日志记录器。
记录每次出题批次的元数据(run)和每题的门判定结果(item)。
使用同步 sqlite3,写入频率低无需异步。
Parameters
----------
db_path : Path
SQLite 数据库文件路径。父目录必须存在。
"""
def __init__(self, db_path: Path) -> None:
self._db_path = db_path
self._conn = sqlite3.connect(
str(db_path),
check_same_thread=False,
timeout=10.0,
)
self._conn.execute("PRAGMA journal_mode=WAL")
self._conn.execute("PRAGMA busy_timeout=5000")
self._conn.execute("PRAGMA foreign_keys=ON")
self._init_schema()
logger.debug("QuestionGenStore 已初始化: {}", db_path)
def _init_schema(self) -> None:
"""幂等创建表和索引。多次调用安全。"""
self._conn.execute(_DDL_RUNS)
self._conn.execute(_DDL_ITEMS)
for idx_sql in _DDL_INDEXES:
self._conn.execute(idx_sql)
self._conn.commit()
# 幂等迁移:为已有表加 sub_pattern / selector_scores 列
cols = {r[1] for r in self._conn.execute("PRAGMA table_info(question_gen_items)")}
if "sub_pattern" not in cols:
self._conn.execute("ALTER TABLE question_gen_items ADD COLUMN sub_pattern TEXT")
self._conn.commit()
if "selector_scores" not in cols:
self._conn.execute("ALTER TABLE question_gen_items ADD COLUMN selector_scores TEXT")
self._conn.commit()
def record_run_start(self, run_id: str, git_sha: str, config_snapshot: str) -> None:
"""记录批次开始。
Parameters
----------
run_id : str
批次唯一标识(UUID)。
git_sha : str
当前代码版本 HEAD commit。
config_snapshot : str
科研配置快照(JSON 序列化字符串)。
"""
now = datetime.now(tz=UTC).isoformat(timespec="seconds")
self._conn.execute(
"""
INSERT INTO question_gen_runs (run_id, git_sha, config_snapshot, started_at, status)
VALUES (?, ?, ?, ?, 'running')
""",
(run_id, git_sha, config_snapshot, now),
)
self._conn.commit()
logger.info("出题批次已开始: run_id={}", run_id)
def record_run_end(self, run_id: str, status: str, stats: RunStats) -> None:
"""记录批次结束,更新状态与统计。
Parameters
----------
run_id : str
批次唯一标识。
status : str
最终状态(completed / failed)。
stats : RunStats
批次统计摘要。
"""
now = datetime.now(tz=UTC).isoformat(timespec="seconds")
cursor = self._conn.execute(
"""
UPDATE question_gen_runs
SET ended_at=?, status=?, total_slots=?, accepted=?, rejected=?, heavy_sampled=?
WHERE run_id=?
""",
(
now,
status,
stats.total_slots,
stats.accepted,
stats.rejected,
stats.heavy_sampled,
run_id,
),
)
self._conn.commit()
if cursor.rowcount == 0:
raise ValueError(f"run_id 不存在: {run_id}")
logger.info(
"出题批次已结束: run_id={}, status={}, accepted={}/{}",
run_id,
status,
stats.accepted,
stats.total_slots,
)
def record_item(
self,
item_id: str,
run_id: str,
slot_id: str,
video_id: str,
family: str,
task_type: str,
skill_target: str,
attempt: int,
question_text: str,
sub_pattern: str | None = None,
) -> None:
"""记录一道新生成的题目(初始状态 pending)。
Parameters
----------
item_id : str
题目唯一 ID(每轮独立)。
run_id : str
关联的批次 ID。
slot_id : str
逻辑 slot 标识(同 slot 多次重出共享)。
video_id : str
视频 ID。
family : str
题族(retrieval/reasoning/enumeration/visual/spatial)。
task_type : str
Video-MME 12 类主标签。
skill_target : str
M1-M5 + 题族子标签。
attempt : int
当前重出轮次(1-based)。
question_text : str
题目文本。
sub_pattern : str | None
子模式标识(如有)。
"""
now = datetime.now(tz=UTC).isoformat(timespec="seconds")
self._conn.execute(
"""
INSERT INTO question_gen_items
(item_id, run_id, slot_id, video_id, family, task_type,
skill_target, attempt, question_text, sub_pattern, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
item_id,
run_id,
slot_id,
video_id,
family,
task_type,
skill_target,
attempt,
question_text,
sub_pattern,
now,
),
)
self._conn.commit()
def update_gates(self, item_id: str, report: GateReportLike) -> None:
"""更新门判定结果及 final_status。
Parameters
----------
item_id : str
题目唯一 ID。
report : GateReportLike
门判定报告(duck-type,需有 key_verify/blind_answer/multi_true/leak_test
属性,每个属性具有 .verdict.value 和 .reason;以及 passed/reject_reason 属性)。
"""
final_status = "accepted" if report.passed else "rejected"
cursor = self._conn.execute(
"""
UPDATE question_gen_items
SET gate_key_verify=?, gate_blind_answer=?, gate_multi_true=?,
gate_leak_test=?, gate_reject_reason=?, final_status=?
WHERE item_id=?
""",
(
report.key_verify.verdict.value,
report.blind_answer.verdict.value,
report.multi_true.verdict.value,
report.leak_test.verdict.value,
report.reject_reason,
final_status,
item_id,
),
)
self._conn.commit()
if cursor.rowcount == 0:
raise ValueError(f"item_id 不存在: {item_id}")
def mark_item_rejected(self, item_id: str, reason: str) -> None:
"""将已记录的 item 标记为 rejected(用于门控外的拒绝场景,如去重)。
Parameters
----------
item_id : str
题目唯一 ID。
reason : str
拒绝原因描述。
"""
cursor = self._conn.execute(
"UPDATE question_gen_items SET final_status='rejected', gate_reject_reason=? "
"WHERE item_id=?",
(reason, item_id),
)
self._conn.commit()
if cursor.rowcount == 0:
raise ValueError(f"item_id 不存在: {item_id}")
def update_difficulty(self, item_id: str, difficulty_steps: int) -> None:
"""更新重量抽检产出的 Agent 步数。
Parameters
----------
item_id : str
题目唯一 ID。
difficulty_steps : int
Agent 完成该题所需步数。
"""
cursor = self._conn.execute(
"UPDATE question_gen_items SET difficulty_steps=? WHERE item_id=?",
(difficulty_steps, item_id),
)
self._conn.commit()
if cursor.rowcount == 0:
raise ValueError(f"item_id 不存在: {item_id}")
def update_selector_scores(self, item_id: str, selector_scores_json: str) -> None:
"""写入 grounded selector 打分观测(JSON 字符串)。
Parameters
----------
item_id : str
题目唯一 ID。
selector_scores_json : str
观测 JSONcorrect_score / chosen / pool_size / anneal_rounds / hard_fail。
Raises
------
ValueError
item_id 不存在时抛出。
"""
cursor = self._conn.execute(
"UPDATE question_gen_items SET selector_scores=? WHERE item_id=?",
(selector_scores_json, item_id),
)
self._conn.commit()
if cursor.rowcount == 0:
raise ValueError(f"item_id 不存在: {item_id}")
def get_run_stats(self, run_id: str) -> RunStats:
"""查询批次统计摘要。
Parameters
----------
run_id : str
批次唯一标识。
Returns
-------
RunStats
该批次的统计数据。
Raises
------
ValueError
run_id 不存在时抛出。
"""
row = self._conn.execute(
"SELECT total_slots, accepted, rejected, heavy_sampled "
"FROM question_gen_runs WHERE run_id=?",
(run_id,),
).fetchone()
if row is None:
raise ValueError(f"run_id 不存在: {run_id}")
return RunStats(
total_slots=row[0] or 0,
accepted=row[1] or 0,
rejected=row[2] or 0,
heavy_sampled=row[3] or 0,
)
def load_progress(self) -> dict[str, str]:
"""加载已接受 slot 的进度映射(用于断点续跑)。
从最近一次 running 状态的批次中,只读取 accepted 的 slot。
rejected 的 slot 不纳入 progress,以便重跑时重新尝试。
Returns
-------
dict[str, str]
{slot_id: "accepted"} 映射。无进度时返回空 dict。
"""
row = self._conn.execute(
"SELECT run_id FROM question_gen_runs WHERE status='running' "
"ORDER BY started_at DESC LIMIT 1",
).fetchone()
if row is None:
return {}
run_id = row[0]
rows = self._conn.execute(
"SELECT DISTINCT slot_id FROM question_gen_items "
"WHERE run_id=? AND final_status='accepted'",
(run_id,),
).fetchall()
return {row[0]: "accepted" for row in rows}
def close(self) -> None:
"""关闭数据库连接。"""
self._conn.close()
logger.debug("QuestionGenStore 已关闭")