feat: wire grounded selector into AR slot processing

将 Task 5 的 grounded selector 织入 AR 出题路径(Phase 3.5,位于
record_item 与 postprocess 之间),仅在 strategy.uses_grounded_selector
为真时进入。observation 始终落库(含 hard-fail),硬失败走重出。
PipelineConfig 新增 candidate_pool_size/selector_delta_low/
selector_delta_high 三参,YAML 与 CLI seed override 同步。
This commit is contained in:
2026-07-14 14:15:50 -04:00
parent d0194f5840
commit b13eab0659
5 changed files with 171 additions and 0 deletions
+105
View File
@@ -24,6 +24,7 @@
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import json
import random import random
import uuid import uuid
from dataclasses import dataclass from dataclasses import dataclass
@@ -98,6 +99,9 @@ class PipelineConfig:
concurrency: 并发 slot 数上限。 concurrency: 并发 slot 数上限。
seed: 随机种子。 seed: 随机种子。
output_dir: 输出目录。 output_dir: 输出目录。
candidate_pool_size: grounded selector 首轮候选干扰项数 N。
selector_delta_low: 干扰项视觉分与正解的最小差(区间上界)。
selector_delta_high: 干扰项视觉分与正解的最大差(区间下界)。
""" """
per_type: int per_type: int
@@ -107,6 +111,9 @@ class PipelineConfig:
concurrency: int concurrency: int
seed: int seed: int
output_dir: Path output_dir: Path
candidate_pool_size: int = 24
selector_delta_low: float = 0.05
selector_delta_high: float = 0.35
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -150,6 +157,9 @@ def load_pipeline_config(yaml_path: Path) -> PipelineConfig:
concurrency=int(section["concurrency"]), concurrency=int(section["concurrency"]),
seed=int(section["seed"]), seed=int(section["seed"]),
output_dir=Path(section["output_dir"]), output_dir=Path(section["output_dir"]),
candidate_pool_size=int(section.get("candidate_pool_size", 24)),
selector_delta_low=float(section.get("selector_delta_low", 0.05)),
selector_delta_high=float(section.get("selector_delta_high", 0.35)),
) )
@@ -281,6 +291,56 @@ def _to_generated_question(
) )
def _extract_correct_text(options: tuple[str, ...], answer: str) -> str:
"""从四选项中取正解文本(去掉 "X. " 字母前缀)。
参数:
options: 选项元组,格式 ("A. ...", "B. ...", ...)。
answer: 正解字母(大小写不敏感)。
返回:
正解选项去前缀后的文本。
异常:
ValueError: answer 对应索引超出选项范围。
"""
idx = ord(answer.strip().upper()) - ord("A")
if not 0 <= idx < len(options):
msg = f"answer '{answer}' 超出选项范围 (n={len(options)})"
raise ValueError(msg)
opt = options[idx]
prefix = f"{answer.strip().upper()}. "
return opt[len(prefix) :] if opt.startswith(prefix) else opt
def _replace_candidate_options(
candidate: CandidateQuestion, options: tuple[str, ...], answer: str
) -> CandidateQuestion:
"""用 selector 重组的选项/答案替换候选(CandidateQuestion frozen)。
参数:
candidate: 原候选题目。
options: grounded selector 重组后的四选项。
answer: 重组后的正解字母(恒 "A")。
返回:
仅替换 options/answer、其余字段照搬的新 CandidateQuestion。
"""
return 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=options,
answer=answer,
source_nodes=candidate.source_nodes,
difficulty=candidate.difficulty,
subtitle_sentences=candidate.subtitle_sentences,
frame_paths=candidate.frame_paths,
)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# 单 Slot 处理(重出循环) # 单 Slot 处理(重出循环)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -428,6 +488,51 @@ async def _process_one_slot(
sub_pattern=sub_pattern.name if sub_pattern else None, sub_pattern=sub_pattern.name if sub_pattern else None,
) )
# Phase 3.5: grounded selector(仅 AR 路径)
if strategy.uses_grounded_selector:
from app.question_gen.distractor_selector import (
SelectorConfig,
build_grounded_options,
)
correct_text = _extract_correct_text(candidate.options, candidate.answer)
selector_cfg = SelectorConfig(
candidate_pool_size=config.candidate_pool_size,
delta_low=config.selector_delta_low,
delta_high=config.selector_delta_high,
)
try:
outcome = await build_grounded_options(
vlm=vlm,
question=candidate.question,
correct_text=correct_text,
material=material,
config=selector_cfg,
session_id=session_id,
)
except (ValueError, FileNotFoundError) as e:
logger.warning(
"slot {} selector 异常 (attempt {}): {}", slot.slot_id, attempt, e
)
prev_reason = f"selector_error: {e}"
continue
# observation 始终落库(含 hard-fail),供 EOB 退化观测与调参
store.update_selector_scores(
item_id, json.dumps(outcome.observation, ensure_ascii=False)
)
if outcome.hard_fail:
prev_reason = "grounded 干扰项不足(selector 硬失败)"
store.mark_item_rejected(item_id, prev_reason)
logger.info("slot {} selector 硬失败 (attempt {})", slot.slot_id, attempt)
continue
# 用 grounded 四选项替换候选(frozen → 构造新实例)
candidate = _replace_candidate_options(
candidate, outcome.options, outcome.answer
)
# Phase 4: 后处理 # Phase 4: 后处理
source_texts = list(material.subtitle_sentences) source_texts = list(material.subtitle_sentences)
pp = run_postprocess( pp = run_postprocess(
+14
View File
@@ -0,0 +1,14 @@
# AR 专项出题配置:使用 ActionRecognitionStrategy + 6 SubPattern 靶向生成
# 目标:生成 30 道 Action Recognition 题
question_gen_v2:
per_type: 30 # 只跑 AR 一类,30 题
retry_limit: 15 # AR 约束更严,给更多重试机会
heavy_sample_rate: 0.0 # 不需要重量抽检
dedup_threshold: 0.85
concurrency: 8 # AR 需要帧,适度并发
seed: 2024
output_dir: "store/questions/generated-ar30"
candidate_pool_size: 24 # grounded selector 首轮候选干扰项数 N
selector_delta_low: 0.05 # 干扰项视觉分与正解的最小差(区间上界)
selector_delta_high: 0.35 # 干扰项视觉分与正解的最大差(区间下界)
+7
View File
@@ -194,8 +194,15 @@ class MockVLM:
parent_call_id: str | None = None, parent_call_id: str | None = None,
) -> LLMResponse: ) -> LLMResponse:
prompt_text = str(messages) prompt_text = str(messages)
system_text = messages[0].get("content", "") if messages else ""
if "verdict" in prompt_text.lower(): if "verdict" in prompt_text.lower():
return _make_llm_response(self._gate_response) return _make_llm_response(self._gate_response)
if "distractor" in system_text.lower() and "grader" not in system_text.lower():
# 候选池请求:返回 4 个 grounded 干扰项
return _make_llm_response('{"distractors": ["", "", "", ""]}')
if "grader" in system_text.lower():
# 打分请求:正解高分、3 个落区间、1 个负空间
return _make_llm_response('{"scores": [0.90, 0.80, 0.70, 0.60, 0.20]}')
idx = min(self._gen_count, len(self._responses) - 1) idx = min(self._gen_count, len(self._responses) - 1)
self._gen_count += 1 self._gen_count += 1
return _make_llm_response(self._responses[idx]) return _make_llm_response(self._responses[idx])
@@ -0,0 +1,42 @@
"""selector 织入辅助:正解文本提取 + 分流。"""
import pytest
from app.question_gen.pipeline_v2 import _extract_correct_text
def test_extract_correct_text_strips_prefix():
options = ("A. 蒸", "B. 炒", "C. 煮", "D. 炸")
assert _extract_correct_text(options, "C") == ""
def test_extract_correct_text_handles_lowercase_answer():
options = ("A. run", "B. walk", "C. jump", "D. sit")
assert _extract_correct_text(options, "b") == "walk"
def test_extract_correct_text_out_of_range_raises():
options = ("A. a", "B. b", "C. c", "D. d")
with pytest.raises(ValueError):
_extract_correct_text(options, "E")
def test_replace_candidate_options():
from app.question_gen.generator_v2 import CandidateQuestion
from app.question_gen.pipeline_v2 import _replace_candidate_options
c = CandidateQuestion(
question_id="q",
video_id="v",
task_type="Action Recognition",
skill_target="M1_AR",
question="?",
options=("A. a", "B. b", "C. c", "D. d"),
answer="A",
source_nodes=("n1",),
difficulty="hard",
)
new = _replace_candidate_options(c, ("A. 蒸", "B. 炒", "C. 煮", "D. 炸"), "A")
assert new.options == ("A. 蒸", "B. 炒", "C. 煮", "D. 炸")
assert new.question == "?" # 其余字段不变
assert new.source_nodes == ("n1",)
+3
View File
@@ -912,6 +912,9 @@ async def _run_generate_v2(args: argparse.Namespace) -> None:
concurrency=config.concurrency, concurrency=config.concurrency,
seed=args.seed, seed=args.seed,
output_dir=config.output_dir, output_dir=config.output_dir,
candidate_pool_size=config.candidate_pool_size,
selector_delta_low=config.selector_delta_low,
selector_delta_high=config.selector_delta_high,
) )
logger.info("seed 覆盖为: {}", args.seed) logger.info("seed 覆盖为: {}", args.seed)