diff --git a/app/question_gen/pipeline_v2.py b/app/question_gen/pipeline_v2.py index 99aa530..d164bdb 100644 --- a/app/question_gen/pipeline_v2.py +++ b/app/question_gen/pipeline_v2.py @@ -38,7 +38,7 @@ from app.question_gen.families import QuestionFamilySpec, get_family_for_slot from app.question_gen.gates import GateReport, GateResult, GateVerdict, run_gates from app.question_gen.generator_v2 import CandidateQuestion, generate_one_v2 from app.question_gen.postprocess import run_postprocess -from app.question_gen.sampler_v2 import sample_material_v2 +from app.question_gen.sampler_v2 import _TASK_TYPE_TO_LEVEL, sample_material_v2 from core.types import GeneratedQuestion if TYPE_CHECKING: @@ -101,8 +101,6 @@ class PipelineConfig: concurrency: 并发 slot 数上限。 seed: 随机种子。 output_dir: 输出目录。 - gate_models: 门控模型配置字典。 - heavy_agent_model: 重量抽检使用的模型名。 """ family_ratios: dict[str, float] @@ -113,8 +111,6 @@ class PipelineConfig: concurrency: int seed: int output_dir: Path - gate_models: dict[str, str] - heavy_agent_model: str # --------------------------------------------------------------------------- @@ -163,8 +159,6 @@ def load_pipeline_config(yaml_path: Path) -> PipelineConfig: concurrency=int(section["concurrency"]), seed=int(section["seed"]), output_dir=Path(section["output_dir"]), - gate_models=section["gate"], - heavy_agent_model=str(section["heavy_agent_model"]), ) @@ -269,6 +263,7 @@ def _is_duplicate( def _to_generated_question( candidate: CandidateQuestion, *, + family: str, options: tuple[str, ...] | None = None, answer: str | None = None, ) -> GeneratedQuestion: @@ -276,6 +271,7 @@ def _to_generated_question( 参数: candidate: 门控通过的候选题目。 + family: 问题家族名称(如 "RETRIEVAL")。 options: 洗牌后的选项元组(若为 None 则使用 candidate 原始选项)。 answer: 重映射后的答案字母(若为 None 则使用 candidate 原始答案)。 @@ -291,6 +287,7 @@ def _to_generated_question( answer=answer if answer is not None else candidate.answer, source_nodes=candidate.source_nodes, difficulty=candidate.difficulty, + family=family, skill_target=candidate.skill_target, difficulty_steps=None, ) @@ -316,6 +313,7 @@ async def _process_one_slot( *, session_id: str, run_id: str, + all_trees: dict[str, TreeIndex] | None = None, ) -> GeneratedQuestion | None: """处理单个 slot 的完整重出循环。 @@ -347,18 +345,36 @@ async def _process_one_slot( 返回: GeneratedQuestion(通过全部检查)或 None(重出耗尽)。 """ + _RESAMPLE_VIDEO_INTERVAL = 1 + async with sem: prev_reason: str | None = None + current_tree = tree + current_video_id = slot.video_id for attempt in range(1, config.retry_limit + 1): + # 连续失败 _RESAMPLE_VIDEO_INTERVAL 次后换视频 + if attempt > 1 and (attempt - 1) % _RESAMPLE_VIDEO_INTERVAL == 0 and all_trees: + alt_ids = [v for v in all_trees if v != current_video_id] + if alt_ids: + current_video_id = rng.choice(alt_ids) + current_tree = all_trees[current_video_id] + logger.info( + "slot {} 连续 {} 次失败,换视频 {} 重试", + slot.slot_id, + attempt - 1, + current_video_id, + ) + # Phase 1: 采样素材 try: material = sample_material_v2( - tree=tree, - family_spec=slot.family, + tree=current_tree, task_type=slot.task_type, used_node_ids=used_node_ids, rng=rng, + level=_TASK_TYPE_TO_LEVEL[slot.task_type], + constraint=slot.family.sampling, ) except (RuntimeError, KeyError) as e: logger.warning( @@ -374,16 +390,16 @@ async def _process_one_slot( try: candidate = await generate_one_v2( vlm=vlm, - tree=tree, + tree=current_tree, material=material, family_spec=slot.family, task_type=slot.task_type, seq=slot.seq, - video_id=slot.video_id, + video_id=current_video_id, reject_reason=prev_reason, session_id=session_id, ) - except (ValueError, FileNotFoundError) as e: + except (ValueError, FileNotFoundError, OSError, Exception) as e: logger.warning( "slot {} 生成失败 (attempt {}/{}): {}", slot.slot_id, @@ -399,7 +415,7 @@ async def _process_one_slot( item_id=item_id, run_id=run_id, slot_id=slot.slot_id, - video_id=slot.video_id, + video_id=current_video_id, family=slot.family.name, task_type=slot.task_type, skill_target=slot.family.skill_target, @@ -446,15 +462,27 @@ async def _process_one_slot( store.update_gates(item_id, verbatim_report) continue - # Phase 6: 四门质量检查 - report = await run_gates( - candidate=candidate, - tree=tree, - llm=llm, - family_spec=slot.family, - postprocess=pp, - session_id=session_id, - ) + # Phase 6: 四门质量检查(key_verify 使用 VLM 看帧+文本) + try: + report = await run_gates( + candidate=candidate, + tree=tree, + llm=llm, + family_spec=slot.family, + postprocess=pp, + vlm=vlm, + session_id=session_id, + ) + except Exception as e: + logger.warning( + "slot {} 门控调用异常 (attempt {}/{}): {}", + slot.slot_id, + attempt, + config.retry_limit, + e, + ) + prev_reason = f"gate_error: {e}" + continue store.update_gates(item_id, report) if not report.passed: @@ -481,7 +509,12 @@ async def _process_one_slot( continue # Phase 8: 通过全部检查 → 接受(使用洗牌后的选项和答案) - result = _to_generated_question(candidate, options=pp.options, answer=pp.answer) + result = _to_generated_question( + candidate, + family=slot.family.name, + options=pp.options, + answer=pp.answer, + ) # 将题目 embedding 加入池(flatten 确保 1D) embed_pool.append(embed_fn(candidate.question).flatten()) # 标记使用的节点 @@ -571,16 +604,16 @@ async def _heavy_check_one( _DEFAULT_TASK_TYPES: list[str] = [ "Action Recognition", "Action Reasoning", - "Action Prediction", - "Action Sequence", + "Attribute Perception", + "Counting Problem", + "Information Synopsis", "Object Recognition", "Object Reasoning", - "Object Interaction", - "Scene Understanding", - "Event Reasoning", - "Causal Reasoning", - "Temporal Reasoning", + "OCR Problems", + "Spatial Perception", "Spatial Reasoning", + "Temporal Perception", + "Temporal Reasoning", ] @@ -697,6 +730,7 @@ async def run_pipeline_v2( *, task_types: list[str] | None = None, progress: dict[str, str] | None = None, + on_accept: Callable[[GeneratedQuestion], None] | None = None, ) -> PipelineResult: """v2 出题管线主入口 — 编排全部 slot 的生成、检查与抽检。 @@ -718,6 +752,7 @@ async def run_pipeline_v2( config: 管线配置。 task_types: 任务类型列表(默认使用 12 类标准集)。 progress: 已完成 slot 映射 {slot_id → "accepted"|"rejected"}。 + on_accept: 每接受一题时的回调(用于实时持久化,防崩溃丢数据)。 返回: PipelineResult 实例。 @@ -749,7 +784,7 @@ async def run_pipeline_v2( if tree is None: logger.warning("slot {} 对应视频 {} 的树不存在,跳过", slot.slot_id, slot.video_id) return None - return await _process_one_slot( + result = await _process_one_slot( slot=slot, tree=tree, vlm=vlm, @@ -763,7 +798,11 @@ async def run_pipeline_v2( sem=sem, session_id=session_id, run_id=run_id, + all_trees=trees, ) + if result is not None and on_accept is not None: + on_accept(result) + return result results = await asyncio.gather(*[_process_wrapper(s) for s in pending_slots]) diff --git a/app/question_gen/sampler_v2.py b/app/question_gen/sampler_v2.py index 52e8800..5dc735d 100644 --- a/app/question_gen/sampler_v2.py +++ b/app/question_gen/sampler_v2.py @@ -1,16 +1,17 @@ -"""v2 素材采样器 — 基于家族约束的树节点采样与上下文收集。 +"""v2 素材采样器 — 基于采样约束的树节点采样与上下文收集。 -在 v1 synthesizer 的基础上引入 QuestionFamilySpec 约束验证, +在 v1 synthesizer 的基础上引入 SamplingConstraint 约束验证, 为每次出题提供更丰富的素材上下文(字幕、跨 L2 上下文、帧路径)。 典型调用路径:: material = sample_material_v2( tree=tree_index, - family_spec=REASONING_FAMILY, - task_type="Causal Reasoning", + task_type="Action Reasoning", used_node_ids=already_used, rng=rng, + level=2, + constraint=my_constraint, ) """ @@ -24,7 +25,7 @@ from loguru import logger if TYPE_CHECKING: import random - from app.question_gen.families import QuestionFamilySpec, SamplingConstraint + from app.question_gen.families import SamplingConstraint from app.tree.index import L1Node, L2Node, TreeIndex # --------------------------------------------------------------------------- @@ -35,18 +36,18 @@ _TASK_TYPE_TO_LEVEL: dict[str, int] = { # Level 3(细粒度帧级) "Action Recognition": 3, "Object Recognition": 3, + "Attribute Perception": 3, + "OCR Problems": 3, # Level 2(片段/事件级) "Action Reasoning": 2, - "Action Prediction": 2, - "Action Sequence": 2, "Object Reasoning": 2, - "Object Interaction": 2, - "Scene Understanding": 2, - "Event Reasoning": 2, - "Causal Reasoning": 2, + "Information Synopsis": 2, + "Counting Problem": 2, # Level 1(段落/场景级) "Temporal Reasoning": 1, + "Temporal Perception": 1, "Spatial Reasoning": 1, + "Spatial Perception": 1, } @@ -494,17 +495,18 @@ def _collect_source_nodes(tree: TreeIndex, node_id: str) -> tuple[str, ...]: def sample_material_v2( tree: TreeIndex, - family_spec: QuestionFamilySpec, task_type: str, used_node_ids: set[str], rng: random.Random, *, + level: int, + constraint: SamplingConstraint, max_attempts: int = 10, ) -> MaterialContext: - """基于家族约束从视频树中采样素材上下文。 + """基于采样约束从视频树中采样素材上下文。 采样流程: - 1. 根据 task_type 确定采样层级 + 1. 按指定 level 确定采样层级 2. 随机选取候选节点(排除 used_node_ids) 3. 验证 SamplingConstraint 约束 4. 约束不满足则重试(最多 max_attempts 次) @@ -512,10 +514,11 @@ def sample_material_v2( 参数: tree: 三层树索引。 - family_spec: 问题家族规格(含采样约束)。 task_type: 任务类型字符串。 used_node_ids: 本轮已用节点 ID 集合。 rng: 可控随机数生成器。 + level: 采样层级(1/2/3)。 + constraint: 采样约束条件。 max_attempts: 最大尝试次数。 返回: @@ -523,10 +526,7 @@ def sample_material_v2( 异常: RuntimeError: 耗尽 max_attempts 次尝试仍无法满足约束。 - KeyError: task_type 不在 _TASK_TYPE_TO_LEVEL 映射中。 """ - level = _TASK_TYPE_TO_LEVEL[task_type] - constraint = family_spec.sampling for attempt in range(max_attempts): # Phase 1: 按层级采样候选节点 @@ -589,5 +589,5 @@ def sample_material_v2( raise RuntimeError( f"sample_material_v2: 耗尽 max_attempts={max_attempts} 次尝试," - f"无法为 task_type='{task_type}' 满足家族 '{family_spec.name}' 的采样约束" + f"无法为 task_type='{task_type}' (level={level}) 满足采样约束" ) diff --git a/tests/unit/test_sampler_v2.py b/tests/unit/test_sampler_v2.py index be4d60e..2f1cea4 100644 --- a/tests/unit/test_sampler_v2.py +++ b/tests/unit/test_sampler_v2.py @@ -181,10 +181,11 @@ class TestSampleMaterialV2: rng = random.Random(42) result = sample_material_v2( tree=real_tree, - family_spec=RETRIEVAL_FAMILY, task_type="Action Reasoning", used_node_ids=set(), rng=rng, + level=2, + constraint=RETRIEVAL_FAMILY.sampling, ) assert isinstance(result, MaterialContext) @@ -213,10 +214,11 @@ class TestSampleMaterialV2: result = sample_material_v2( tree=real_tree, - family_spec=RETRIEVAL_FAMILY, task_type="Action Reasoning", used_node_ids=used, rng=rng, + level=2, + constraint=RETRIEVAL_FAMILY.sampling, ) # 锚节点应该是那个未被排除的 L2 @@ -233,10 +235,11 @@ class TestSampleMaterialV2: with pytest.raises(RuntimeError, match="max_attempts"): sample_material_v2( tree=sparse_tree, - family_spec=VISUAL_FAMILY, task_type="Object Recognition", used_node_ids=set(), rng=rng, + level=3, + constraint=VISUAL_FAMILY.sampling, max_attempts=3, ) @@ -248,10 +251,11 @@ class TestSampleMaterialV2: result = sample_material_v2( tree=real_tree, - family_spec=REASONING_FAMILY, - task_type="Causal Reasoning", + task_type="Action Reasoning", used_node_ids=set(), rng=rng, + level=2, + constraint=REASONING_FAMILY.sampling, ) # cross_l2_span=True 时必须有跨 L2 文本 @@ -265,10 +269,11 @@ class TestSampleMaterialV2: result = sample_material_v2( tree=real_tree, - family_spec=RETRIEVAL_FAMILY, task_type="Action Reasoning", used_node_ids=set(), rng=rng, + level=2, + constraint=RETRIEVAL_FAMILY.sampling, ) # real_tree 所有节点都有字幕,所以 subtitle_sentences 非空