Files
Video-Tree-TRM5/app/question_gen/generator_v2.py
iomgaa afa77173e3 refactor(question_gen): adapt generator/gates/store signatures for strategy
- generator_v2: _load_prompt_template takes template_name str instead of
  QuestionFamilySpec; _build_v2_prompt takes prompt_template + strategy_name
  + sub_pattern_instruction; generate_one_v2 takes discrete params
  (prompt_template, strategy_name, skill_target, sub_pattern_instruction)
- gates: _gate_leak_test and run_gates take leak_probe_template str
  instead of QuestionFamilySpec
- run_store: add sub_pattern column to DDL + idempotent migration;
  record_item accepts optional sub_pattern param
- Remove QuestionFamilySpec imports from generator_v2 and gates modules
- Update test call sites accordingly

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-07-14 05:50:47 -04:00

462 lines
14 KiB
Python
Raw Permalink 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.
"""v2 生成器 — 基于策略特化 prompt 模板的单题 VLM 出题模块。
使用 VLMProvider 接口调用视觉语言模型,结合 per-strategy prompt 模板
和 MaterialContext 素材上下文,生成一道四选一候选题。
典型调用路径::
candidate = await generate_one_v2(
vlm=vlm_client,
tree=tree_index,
material=material_ctx,
task_type="Action Reasoning",
seq=1,
video_id="vid_001",
prompt_template="retrieval.md",
strategy_name="RETRIEVAL",
skill_target="M1",
session_id="sess_001",
)
"""
from __future__ import annotations
import json
from dataclasses import dataclass, field
from pathlib import Path
from typing import TYPE_CHECKING
from json_repair import repair_json
from loguru import logger
if TYPE_CHECKING:
from app.question_gen.sampler_v2 import MaterialContext
from app.tree.index import TreeIndex
from core.protocols import VLMProvider
# ---------------------------------------------------------------------------
# 常量
# ---------------------------------------------------------------------------
_PROMPTS_DIR = Path(__file__).resolve().parent.parent.parent / "store" / "prompts" / "question_gen"
_VALID_ANSWERS = frozenset({"A", "B", "C", "D"})
_VALID_DIFFICULTIES = frozenset({"easy", "medium", "hard"})
# ---------------------------------------------------------------------------
# CandidateQuestion(规范定义位置)
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class CandidateQuestion:
"""候选题目 — 出题管线生成、待门控审核的题目数据。
属性:
question_id: 题目唯一标识(格式 "{video_id}_{task_type}_{seq:04d}")。
video_id: 所属视频标识。
task_type: 题型(如 "Action Reasoning")。
skill_target: 目标失败机制编号(M1-M5)。
question: 题目文本。
options: 选项元组(如 ("A. ...", "B. ...", "C. ...", "D. ..."))。
answer: 正确答案字母(如 "A")。
source_nodes: 来源节点 ID 元组。
difficulty: 难度等级(easy/medium/hard)。
subtitle_sentences: 验证材料 — 字幕句子元组。
frame_paths: 验证材料 — 帧图片路径元组。
"""
question_id: str
video_id: str
task_type: str
skill_target: str
question: str
options: tuple[str, ...]
answer: str
source_nodes: tuple[str, ...]
difficulty: str
subtitle_sentences: tuple[str, ...] = field(default_factory=tuple)
frame_paths: tuple[str, ...] = field(default_factory=tuple)
# ---------------------------------------------------------------------------
# Prompt 模板加载
# ---------------------------------------------------------------------------
def _load_prompt_template(template_name: str) -> str:
"""加载 prompt 模板文件。
参数:
template_name: store/prompts/question_gen/ 下的模板文件名。
返回:
模板内容字符串。
异常:
FileNotFoundError: 模板文件不存在。
"""
path = _PROMPTS_DIR / template_name
if not path.exists():
msg = f"Prompt 模板文件不存在: {path}"
raise FileNotFoundError(msg)
return path.read_text(encoding="utf-8")
# ---------------------------------------------------------------------------
# Prompt 构建
# ---------------------------------------------------------------------------
def _build_v2_prompt(
prompt_template: str,
strategy_name: str,
material: MaterialContext,
task_type: str,
seq: int,
*,
reject_reason: str | None = None,
sub_pattern_instruction: str | None = None,
) -> tuple[list[dict[str, str]], list[str]]:
"""构建 VLM 出题调用的 messages 和帧路径列表。
参数:
prompt_template: prompt 模板文件名(store/prompts/question_gen/ 下)。
strategy_name: 策略名称(如 "RETRIEVAL")。
material: 采样素材上下文。
task_type: 任务类型字符串。
seq: 当前序号。
reject_reason: 上一次被门控拒绝的原因(用于引导 VLM 避免相同错误)。
sub_pattern_instruction: 子模式特殊指令(如有)。
返回:
二元组:
- messages: 适配 VLMProvider 的 message 列表(system + user)。
- frame_paths: 需发送给 VLM 的帧路径列表。
"""
# Phase 1: 加载策略模板作为 system prompt
template_content = _load_prompt_template(prompt_template)
system_message = template_content
# Phase 2: 构建 user prompt — 聚合素材信息
user_parts: list[str] = []
user_parts.append(f"## Task Type: {task_type}")
user_parts.append(f"## Question Family: {strategy_name}")
user_parts.append(f"## Sequence: #{seq}")
# 字幕素材
if material.subtitle_sentences:
user_parts.append("\n## Subtitle Content:")
for i, sent in enumerate(material.subtitle_sentences, 1):
user_parts.append(f" {i}. {sent}")
# 跨 L2 上下文
if material.cross_l2_texts:
user_parts.append("\n## Cross-Segment Context:")
for text in material.cross_l2_texts:
user_parts.append(f" - {text}")
# 帧路径提示(VLM 会接收实际图像,此处仅作文本参考)
if material.frame_paths:
user_parts.append(f"\n## Visual Frames: {len(material.frame_paths)} frames attached.")
# 拒绝原因注入
if reject_reason is not None:
user_parts.append(
f"\n## IMPORTANT - Previous Attempt Rejected:\n"
f"Your previous question was rejected for the following reason:\n"
f'"{reject_reason}"\n'
f"Please generate a NEW question that avoids this issue."
)
# 子模式特殊指令注入
if sub_pattern_instruction is not None:
user_parts.append(f"\n## Special Focus:\n{sub_pattern_instruction}")
# 输出格式指令
user_parts.append(
"\n## Output Format:\n"
"Respond with ONLY a JSON object in this exact format:\n"
"```json\n"
"{\n"
' "question": "Your question text here",\n'
' "options": ["A. ...", "B. ...", "C. ...", "D. ..."],\n'
' "answer": "A",\n'
' "difficulty": "easy|medium|hard"\n'
"}\n"
"```"
)
user_content = "\n".join(user_parts)
messages = [
{"role": "system", "content": system_message},
{"role": "user", "content": user_content},
]
# Phase 3: 帧路径
frame_paths = list(material.frame_paths)
return messages, frame_paths
# ---------------------------------------------------------------------------
# 响应解析
# ---------------------------------------------------------------------------
def _extract_json_from_text(raw: str) -> str:
"""从可能被 markdown 代码块包裹的文本中提取 JSON 部分。
参数:
raw: VLM 原始返回文本。
返回:
清理后的 JSON 字符串。
"""
content = raw.strip()
if "```" in content:
parts = content.split("```")
for part in parts:
stripped = part.strip()
if stripped.startswith("json"):
stripped = stripped[4:].strip()
if stripped.startswith("{"):
return stripped
return content
@dataclass(frozen=True)
class _ValidatedFields:
"""字段校验通过后的中间结构。"""
question: str
options: tuple[str, ...]
answer: str
difficulty: str
def _validate_parsed_fields(data: dict) -> _ValidatedFields:
"""校验 VLM 响应 JSON 的必填字段并规范化。
参数:
data: 已解析的 JSON 字典。
返回:
_ValidatedFields 实例(字段已规范化)。
异常:
ValueError: 缺少必填字段或字段值非法。
"""
missing = [f for f in ("question", "options", "answer", "difficulty") if f not in data]
if missing:
msg = f"VLM 响应缺少必填字段: {', '.join(missing)}"
raise ValueError(msg)
question_text = str(data["question"])
options_raw = data["options"]
answer = str(data["answer"]).strip().upper()
difficulty = str(data["difficulty"]).strip().lower()
# 校验 optionsanswer 约束为 A-D,因此必须恰好 4 个选项)
if not isinstance(options_raw, list) or len(options_raw) != 4:
msg = f"options 字段必须恰好包含 4 个选项,实际数量: {len(options_raw) if isinstance(options_raw, list) else type(options_raw).__name__}"
raise ValueError(msg)
# 校验 answer
if answer not in _VALID_ANSWERS:
msg = f"answer 字段值 '{answer}' 非法,必须为 A/B/C/D 之一"
raise ValueError(msg)
# 校验 difficulty(宽容处理:非法值回退为 medium)
if difficulty not in _VALID_DIFFICULTIES:
logger.warning(
"difficulty '{}' 不在预设范围 {},回退为 'medium'",
difficulty,
_VALID_DIFFICULTIES,
)
difficulty = "medium"
return _ValidatedFields(
question=question_text,
options=tuple(str(o) for o in options_raw),
answer=answer,
difficulty=difficulty,
)
def _parse_v2_response(
raw: str,
video_id: str,
task_type: str,
skill_target: str,
seq: int,
source_nodes: tuple[str, ...],
) -> CandidateQuestion:
"""解析 VLM 返回的 JSON 响应,构造 CandidateQuestion。
流程:
1. 提取 JSON(处理 markdown 包裹)。
2. json_repair 修复常见格式错误。
3. 校验必填字段(委托 _validate_parsed_fields)。
4. 构造 CandidateQuestion 实例。
参数:
raw: VLM 原始返回文本。
video_id: 视频标识。
task_type: 任务类型。
skill_target: 目标技能编号。
seq: 当前序号。
source_nodes: 来源节点 ID 元组。
返回:
CandidateQuestion 实例。
异常:
ValueError: JSON 解析失败或缺少必填字段或字段值非法。
"""
# Phase 1: 提取 + 修复 JSON
json_text = _extract_json_from_text(raw)
repaired = repair_json(json_text, return_objects=False)
# Phase 2: 解析为字典
try:
data = json.loads(repaired)
except json.JSONDecodeError as e:
msg = f"VLM 响应 JSON 解析失败: {e}. 原始文本: {raw[:200]}"
raise ValueError(msg) from e
if not isinstance(data, dict):
msg = f"VLM 响应顶层不是 JSON 对象: type={type(data).__name__}"
raise ValueError(msg)
# Phase 3: 字段校验
fields = _validate_parsed_fields(data)
# Phase 4: 构造 CandidateQuestion
question_id = f"{video_id}_{task_type}_{seq:04d}"
return CandidateQuestion(
question_id=question_id,
video_id=video_id,
task_type=task_type,
skill_target=skill_target,
question=fields.question,
options=fields.options,
answer=fields.answer,
source_nodes=source_nodes,
difficulty=fields.difficulty,
)
# ---------------------------------------------------------------------------
# 主入口
# ---------------------------------------------------------------------------
async def generate_one_v2(
vlm: VLMProvider,
tree: TreeIndex,
material: MaterialContext,
task_type: str,
seq: int,
*,
video_id: str,
prompt_template: str,
strategy_name: str,
skill_target: str,
reject_reason: str | None = None,
sub_pattern_instruction: str | None = None,
session_id: str,
) -> CandidateQuestion:
"""调用 VLM 生成一道候选题目。
流程:
1. 构建 per-strategy prompt + 帧路径。
2. 调用 VLMProvider.chat_with_images。
3. 解析响应为 CandidateQuestion。
4. 附加素材验证信息(subtitle_sentences、frame_paths)。
参数:
vlm: VLM 调用端口。
tree: 视频树索引(当前未直接使用,预留后续扩展)。
material: 采样素材上下文。
task_type: 任务类型字符串。
seq: 当前序号。
video_id: 视频标识。
prompt_template: prompt 模板文件名。
strategy_name: 策略名称(如 "RETRIEVAL")。
skill_target: 目标失败机制编号(M1-M5)。
reject_reason: 上一次被门控拒绝的原因。
sub_pattern_instruction: 子模式特殊指令(如有)。
session_id: 会话 ID(遥测关联)。
返回:
CandidateQuestion 实例(包含验证材料)。
异常:
ValueError: VLM 响应解析失败。
FileNotFoundError: prompt 模板不存在。
"""
# Phase 1: 构建 prompt
messages, frame_paths = _build_v2_prompt(
prompt_template=prompt_template,
strategy_name=strategy_name,
material=material,
task_type=task_type,
seq=seq,
reject_reason=reject_reason,
sub_pattern_instruction=sub_pattern_instruction,
)
# Phase 2: 调用 VLM
logger.debug(
"generate_one_v2: strategy={}, task_type={}, seq={}, frames={}",
strategy_name,
task_type,
seq,
len(frame_paths),
)
response = await vlm.chat_with_images(
messages,
frame_paths,
session_id=session_id,
)
# Phase 3: 解析响应
candidate = _parse_v2_response(
raw=response.content,
video_id=video_id,
task_type=task_type,
skill_target=skill_target,
seq=seq,
source_nodes=material.source_nodes,
)
# Phase 4: 附加验证材料(构造新实例,因 frozen=True
candidate = CandidateQuestion(
question_id=candidate.question_id,
video_id=candidate.video_id,
task_type=candidate.task_type,
skill_target=candidate.skill_target,
question=candidate.question,
options=candidate.options,
answer=candidate.answer,
source_nodes=candidate.source_nodes,
difficulty=candidate.difficulty,
subtitle_sentences=tuple(material.subtitle_sentences),
frame_paths=tuple(material.frame_paths),
)
logger.debug(
"generate_one_v2 完成: question_id={}, difficulty={}",
candidate.question_id,
candidate.difficulty,
)
return candidate