"""Grounded 干扰项 selector — 候选池 + VLM 视觉打分 + 区间选择(仅 AR 路径)。 把干扰项从"VLM 主观写得像"下沉到机制层:VLM 生成 N 个候选干扰项,再对 候选 + 正解逐一打"视觉可信度"分,按 [正解分-δ_high, 正解分-δ_low] 区间 选 3 个 grounded near-miss,从机制上消灭 Easy-Options Bias。 设计: research-wiki/designs/2026-07-14-grounded-question-gen-phaseA-design.md §3 """ from __future__ import annotations import json from dataclasses import dataclass 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 core.protocols import VLMProvider _PROMPTS_DIR = Path(__file__).resolve().parent.parent.parent / "store" / "prompts" / "question_gen" @dataclass(frozen=True) class SelectorConfig: """selector 科研参数。 属性: candidate_pool_size: 首轮候选干扰项数 N。 delta_low: 干扰项视觉分与正解的最小差(上界,太近=真歧义)。 delta_high: 干扰项视觉分与正解的最大差(下界,太低=负空间)。 max_delta_relax: δ_high 放宽次数上限(退火)。 delta_relax_step: 每次放宽 δ_high 的增量。 """ candidate_pool_size: int delta_low: float delta_high: float max_delta_relax: int = 2 delta_relax_step: float = 0.1 @dataclass(frozen=True) class SelectorOutcome: """selector 产出。observation 始终存在(含 hard-fail),供 run_store 落库。 属性: observation: 打分观测 dict(correct_score/chosen/pool_size/anneal_rounds/hard_fail)。 options: 重组四选项(A=正解),hard-fail 时为 None。 answer: 正解字母(恒 "A"),hard-fail 时为 None。 """ observation: dict[str, object] options: tuple[str, ...] | None = None answer: str | None = None @property def hard_fail(self) -> bool: """是否硬失败(凑不齐 3 个 grounded 干扰项)。""" return self.options is None def _select_in_interval( correct_score: float, candidates: list[str], candidate_scores: list[float], delta_low: float, delta_high: float, ) -> list[str] | None: """从候选中选 3 个视觉分落 [correct-δ_high, correct-δ_low] 区间的干扰项。 落区间者按分数降序取前 3(分数越高越接近正解=越难)。不足 3 个返回 None。 参数: correct_score: 正解视觉可信度分。 candidates: 候选干扰项文本列表。 candidate_scores: 与 candidates 对齐的视觉分列表。 delta_low: 最小差(上界 = correct - delta_low)。 delta_high: 最大差(下界 = correct - delta_high)。 返回: 选中的 3 个候选文本(降序)或 None(不足 3 个)。 """ upper = correct_score - delta_low lower = correct_score - delta_high eligible = [ (c, s) for c, s in zip(candidates, candidate_scores, strict=True) if lower <= s <= upper ] if len(eligible) < 3: return None eligible.sort(key=lambda cs: cs[1], reverse=True) return [c for c, _ in eligible[:3]] def _load_prompt(name: str) -> str: path = _PROMPTS_DIR / name if not path.exists(): msg = f"Prompt 模板不存在: {path}" raise FileNotFoundError(msg) return path.read_text(encoding="utf-8") def _material_context_block(question: str, correct_text: str, material: MaterialContext) -> str: parts = [f"## Question\n{question}", f"## Correct Answer\n{correct_text}"] if material.subtitle_sentences: parts.append("## Subtitles") parts.extend(f" - {s}" for s in material.subtitle_sentences) if getattr(material, "cross_l2_texts", None): parts.append("## Cross-Segment Context") parts.extend(f" - {t}" for t in material.cross_l2_texts) return "\n".join(parts) def _parse_json_object(raw: str) -> dict: content = raw.strip() if "```" in content: for part in content.split("```"): stripped = part.strip() if stripped.startswith("json"): stripped = stripped[4:].strip() if stripped.startswith("{"): content = stripped break data = json.loads(repair_json(content, return_objects=False)) if not isinstance(data, dict): msg = f"selector 响应顶层非 JSON 对象: {type(data).__name__}" raise ValueError(msg) return data async def _generate_pool( vlm: VLMProvider, question: str, correct_text: str, material: MaterialContext, n: int, *, session_id: str, ) -> list[str]: """VLM 生成 n 个候选干扰项文本。""" system = _load_prompt("ar_distractor_pool.md") user = _material_context_block(question, correct_text, material) + f"\n## N\nGenerate exactly {n} distractors." messages = [{"role": "system", "content": system}, {"role": "user", "content": user}] resp = await vlm.chat_with_images(messages, list(material.frame_paths), session_id=session_id) # graceful 降级:解析爆炸时返回空池,让退火/hard_fail 接管(区别于契约违反的打分校验) try: data = _parse_json_object(resp.content) except (ValueError, json.JSONDecodeError) as exc: logger.warning("grounded selector 候选池响应解析失败,降级为空池: {}", exc) return [] raw = data.get("distractors", []) if not isinstance(raw, list): return [] # 防御:去空、去重、剔除与正解字面相同者 seen: set[str] = set() out: list[str] = [] for item in raw: text = str(item).strip() if not text or text == correct_text.strip() or text in seen: continue seen.add(text) out.append(text) return out async def _score_options( vlm: VLMProvider, question: str, options: list[str], material: MaterialContext, *, session_id: str, ) -> list[float]: """VLM 对 options(首个为正解)逐一打视觉可信度分 [0,1],返回对齐分数列表。""" system = _load_prompt("ar_distractor_score.md") numbered = "\n".join(f"{i}. {opt}" for i, opt in enumerate(options, 1)) user = f"## Question\n{question}\n\n## Candidates\n{numbered}" messages = [{"role": "system", "content": system}, {"role": "user", "content": user}] resp = await vlm.chat_with_images(messages, list(material.frame_paths), session_id=session_id) data = _parse_json_object(resp.content) scores_raw = data.get("scores", []) if not isinstance(scores_raw, list) or len(scores_raw) != len(options): msg = f"打分数量({len(scores_raw) if isinstance(scores_raw, list) else 'NA'}) != 选项数({len(options)})" raise ValueError(msg) return [max(0.0, min(1.0, float(s))) for s in scores_raw] async def build_grounded_options( vlm: VLMProvider, question: str, correct_text: str, material: MaterialContext, config: SelectorConfig, *, session_id: str, ) -> SelectorOutcome: """生成候选池 → 视觉打分 → 区间选 3 干扰项 → 重组四选项。 退火(凑不齐 3 个时按序):① 最多追加 N 个新候选(去重后实际增量可能更少)再打分;② 逐步放宽 δ_high(纯重选,不再调 VLM);③ 仍不足则 hard_fail(调用方走重出)。 参数: vlm: VLM 端口。 question: 题干。 correct_text: 正解文本(无字母前缀)。 material: 采样素材(提供 frame_paths / subtitles)。 config: selector 科研参数。 session_id: 遥测会话 ID。 返回: SelectorOutcome。成功时 options=A 正解+3 grounded 干扰项;hard_fail 时 options=None,但 observation 始终存在供落库。 """ candidates = await _generate_pool( vlm, question, correct_text, material, config.candidate_pool_size, session_id=session_id ) # options[0] 恒为正解 scored = await _score_options(vlm, question, [correct_text, *candidates], material, session_id=session_id) correct_score, cand_scores = scored[0], scored[1:] anneal_rounds = 0 chosen = _select_in_interval( correct_score, candidates, cand_scores, config.delta_low, config.delta_high ) # 退火 1: 最多追加 N 个新候选(去重后实际增量可能更少),仅对新增候选打分,正解分保持首轮值 if chosen is None: anneal_rounds += 1 more = await _generate_pool( vlm, question, correct_text, material, config.candidate_pool_size, session_id=session_id ) more = [m for m in more if m not in candidates] if more: more_scores = await _score_options( vlm, question, [correct_text, *more], material, session_id=session_id ) candidates = candidates + more cand_scores = cand_scores + more_scores[1:] chosen = _select_in_interval( correct_score, candidates, cand_scores, config.delta_low, config.delta_high ) # 退火 2: 放宽 δ_high(下界下移,纳入更低分候选),δ_low 不动 relax = 0 delta_high = config.delta_high while chosen is None and relax < config.max_delta_relax: relax += 1 anneal_rounds += 1 delta_high = delta_high + config.delta_relax_step chosen = _select_in_interval( correct_score, candidates, cand_scores, config.delta_low, delta_high ) hard_fail = chosen is None observation: dict[str, object] = { "correct_score": correct_score, "chosen": [ cand_scores[candidates.index(c)] for c in (chosen or []) ], "pool_size": len(candidates), "anneal_rounds": anneal_rounds, "delta_high_final": delta_high, "hard_fail": hard_fail, } if hard_fail: logger.warning( "grounded selector 硬失败: correct={:.3f}, pool={}, anneal={}", correct_score, len(candidates), anneal_rounds, ) # observation 仍返回,供 pipeline 落 selector_scores(设计 §3.3 退化观测) return SelectorOutcome(observation=observation) options = ( f"A. {correct_text}", f"B. {chosen[0]}", f"C. {chosen[1]}", f"D. {chosen[2]}", ) return SelectorOutcome(observation=observation, options=options, answer="A")