From f46e87258c818258ad18e4b41d064b639cc7f707 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Sat, 11 Jul 2026 23:57:35 -0400 Subject: [PATCH] refactor(question_gen): extract helpers to reduce pipeline_v2 CC below grade C Extract _get_git_sha, _filter_pending_slots, and _run_heavy_sampling from run_pipeline_v2. Reduces cyclomatic complexity from C(15) to B(6). Co-Authored-By: Claude Opus 4.6 (1M context) --- app/question_gen/pipeline_v2.py | 205 ++++++++++++++++++++------------ 1 file changed, 128 insertions(+), 77 deletions(-) diff --git a/app/question_gen/pipeline_v2.py b/app/question_gen/pipeline_v2.py index 12afcc1..8fd7fc0 100644 --- a/app/question_gen/pipeline_v2.py +++ b/app/question_gen/pipeline_v2.py @@ -556,6 +556,123 @@ async def _heavy_check_one( return 1 +# --------------------------------------------------------------------------- +# 管线辅助函数 +# --------------------------------------------------------------------------- + +_DEFAULT_TASK_TYPES: list[str] = [ + "Action Recognition", + "Action Reasoning", + "Action Prediction", + "Action Sequence", + "Object Recognition", + "Object Reasoning", + "Object Interaction", + "Scene Understanding", + "Event Reasoning", + "Causal Reasoning", + "Temporal Reasoning", + "Spatial Reasoning", +] + + +def _get_git_sha() -> str: + """获取当前 Git HEAD 短 SHA。 + + 返回: + 短 SHA 字符串;获取失败时返回 "unknown"。 + """ + import subprocess + + try: + return subprocess.check_output( + ["git", "rev-parse", "--short", "HEAD"], + text=True, + timeout=5, + ).strip() + except (subprocess.SubprocessError, FileNotFoundError): + return "unknown" + + +def _filter_pending_slots( + slots: list[SlotAssignment], + progress: dict[str, str], +) -> list[SlotAssignment]: + """过滤出未完成的 slot(跳过 progress 中已记录的)。 + + 参数: + slots: 全部 slot 列表。 + progress: 已完成 slot 映射 {slot_id → "accepted"|"rejected"}。 + + 返回: + 待处理的 slot 列表。 + """ + pending = [s for s in slots if s.slot_id not in progress] + logger.info( + "待处理 slots: {} / {} (已跳过 {})", + len(pending), + len(slots), + len(slots) - len(pending), + ) + return pending + + +async def _run_heavy_sampling( + accepted: list[GeneratedQuestion], + trees: dict[str, TreeIndex], + store: QuestionGenStore, + llm: LLMProvider, + config: PipelineConfig, + *, + session_id: str, +) -> list[tuple[str, int]]: + """对接受题目按比例随机抽检,计算难度步数。 + + 参数: + accepted: 已接受的题目列表。 + trees: video_id → TreeIndex 映射。 + store: 日志记录器。 + llm: LLM 调用端口。 + config: 管线配置(含 heavy_sample_rate 和 seed)。 + session_id: 会话 ID。 + + 返回: + (question_id, difficulty_steps) 元组列表。 + """ + if config.heavy_sample_rate <= 0 or not accepted: + return [] + + sample_count = max(1, int(len(accepted) * config.heavy_sample_rate)) + sample_count = min(sample_count, len(accepted)) + heavy_rng = random.Random(config.seed + 1) + sampled_questions = heavy_rng.sample(accepted, sample_count) + + logger.info("重量抽检: {} / {} 题", len(sampled_questions), len(accepted)) + + heavy_tasks = [] + for q in sampled_questions: + tree = trees.get(q.video_id) + if tree is None: + continue + heavy_tasks.append(_heavy_check_one(q, tree, llm, session_id=session_id)) + + heavy_results = await asyncio.gather(*heavy_tasks) + + heavy_sampled: list[tuple[str, int]] = [] + for q, steps in zip(sampled_questions, heavy_results, strict=True): + heavy_sampled.append((q.question_id, steps)) + cursor = store._conn.execute( + "SELECT item_id FROM question_gen_items " + "WHERE slot_id LIKE ? AND final_status='accepted' LIMIT 1", + (f"%{q.question_id.split('_')[-1]}%",), + ) + row = cursor.fetchone() + if row: + store.update_difficulty(row[0], steps) + + return heavy_sampled + + # --------------------------------------------------------------------------- # 管线主入口 # --------------------------------------------------------------------------- @@ -597,57 +714,23 @@ async def run_pipeline_v2( 返回: PipelineResult 实例。 """ - # 默认 12 类标准任务类型 - if task_types is None: - task_types = [ - "Action Recognition", - "Action Reasoning", - "Action Prediction", - "Action Sequence", - "Object Recognition", - "Object Reasoning", - "Object Interaction", - "Scene Understanding", - "Event Reasoning", - "Causal Reasoning", - "Temporal Reasoning", - "Spatial Reasoning", - ] - + task_types = task_types or _DEFAULT_TASK_TYPES progress = progress or {} rng = random.Random(config.seed) - # Phase 1: 分配 slot + # Phase 1: 分配 slot + 创建 run 记录 slots = _assign_slots(video_ids, task_types, config.per_type, config.family_ratios, rng) logger.info( "管线启动: {} slots, {} 视频, {} 任务类型", len(slots), len(video_ids), len(task_types) ) - # Phase 2: 创建 run 记录 run_id = uuid.uuid4().hex - import subprocess + store.record_run_start(run_id, _get_git_sha(), str(config)) - try: - git_sha = subprocess.check_output( - ["git", "rev-parse", "--short", "HEAD"], - text=True, - timeout=5, - ).strip() - except (subprocess.SubprocessError, FileNotFoundError): - git_sha = "unknown" + # Phase 2: 过滤已完成 slot + pending_slots = _filter_pending_slots(slots, progress) - store.record_run_start(run_id, git_sha, str(config)) - - # Phase 3: 过滤已完成 slot - pending_slots = [s for s in slots if s.slot_id not in progress] - logger.info( - "待处理 slots: {} / {} (已跳过 {})", - len(pending_slots), - len(slots), - len(slots) - len(pending_slots), - ) - - # Phase 4: 并发处理 + # Phase 3: 并发处理 sem = asyncio.Semaphore(config.concurrency) embed_pool: list[np.ndarray] = [] used_node_ids: set[str] = set() @@ -676,48 +759,16 @@ async def run_pipeline_v2( results = await asyncio.gather(*[_process_wrapper(s) for s in pending_slots]) - # Phase 5: 统计结果 + # Phase 4: 统计 + 重量抽检 accepted: list[GeneratedQuestion] = [r for r in results if r is not None] rejected_count = len(pending_slots) - len(accepted) + logger.info("管线生成完成: accepted={}, rejected={}", len(accepted), rejected_count) - logger.info( - "管线生成完成: accepted={}, rejected={}", - len(accepted), - rejected_count, + heavy_sampled = await _run_heavy_sampling( + accepted, trees, store, llm, config, session_id=session_id ) - # Phase 6: 重量抽检 - heavy_sampled: list[tuple[str, int]] = [] - if config.heavy_sample_rate > 0 and accepted: - sample_count = max(1, int(len(accepted) * config.heavy_sample_rate)) - sample_count = min(sample_count, len(accepted)) - heavy_rng = random.Random(config.seed + 1) - sampled_questions = heavy_rng.sample(accepted, sample_count) - - logger.info("重量抽检: {} / {} 题", len(sampled_questions), len(accepted)) - - heavy_tasks = [] - for q in sampled_questions: - tree = trees.get(q.video_id) - if tree is None: - continue - heavy_tasks.append(_heavy_check_one(q, tree, llm, session_id=session_id)) - - heavy_results = await asyncio.gather(*heavy_tasks) - - for q, steps in zip(sampled_questions, heavy_results, strict=True): - heavy_sampled.append((q.question_id, steps)) - # 更新 store - # 找到对应的 item_id(最后一次 attempt 的记录) - cursor = store._conn.execute( - "SELECT item_id FROM question_gen_items WHERE slot_id LIKE ? AND final_status='accepted' LIMIT 1", - (f"%{q.question_id.split('_')[-1]}%",), - ) - row = cursor.fetchone() - if row: - store.update_difficulty(row[0], steps) - - # Phase 7: 更新 run 统计 + # Phase 5: 更新 run 统计 from app.question_gen.run_store import RunStats stats = RunStats(