feat(question_gen): is_duplicate + generate_one — 去重判定与单题生成编排

- is_duplicate: 余弦相似度去重,空池短路
- generate_one: 异步重试循环,不含去重(由调用方汇总点原子执行)

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-07-09 05:32:39 -04:00
parent 90f17e330e
commit 5aa7cc48c5
2 changed files with 255 additions and 1 deletions
+134
View File
@@ -12,10 +12,15 @@ import re
from dataclasses import dataclass
from typing import TYPE_CHECKING
import numpy as np
from loguru import logger
if TYPE_CHECKING:
import random
from collections.abc import Callable
from app.tree.index import L2Node, L3Node, TreeIndex
from core.protocols import VLMProvider
from core.types import GeneratedQuestion
@@ -624,3 +629,132 @@ def parse_vlm_response(
"options": list(options),
"answer": answer,
}
# ---------------------------------------------------------------------------
# Embedding 去重
# ---------------------------------------------------------------------------
def is_duplicate(
question_text: str,
pool_embeddings: np.ndarray,
embed_fn: Callable[[str | list[str]], np.ndarray],
threshold: float,
) -> bool:
"""embedding 去重判定。
参数:
question_text: 待检查的题目文本。
pool_embeddings: 已有题目的 embedding 矩阵 [N, D]L2 归一化)。
embed_fn: 文本嵌入函数,返回 [N, D] ndarrayL2 归一化)。
threshold: 余弦相似度阈值。
返回:
True 表示与池中某题重复。空池永远返回 False。
"""
if pool_embeddings.shape[0] == 0:
return False
query = embed_fn(question_text) # [1, D]
query = query.squeeze(0) # [D]
similarities = pool_embeddings @ query # [N]
return bool(np.max(similarities) >= threshold)
# ---------------------------------------------------------------------------
# 单题生成
# ---------------------------------------------------------------------------
async def generate_one(
vlm: VLMProvider,
embed_fn: Callable[[str | list[str]], np.ndarray],
tree: TreeIndex,
video_id: str,
task_type: str,
seq: int,
*,
exemplars: list[GeneratedQuestion],
used_node_ids: set[str],
max_retries: int,
similarity_threshold: float,
rng: random.Random,
session_id: str,
) -> GeneratedQuestion | None:
"""生成单道候选题(不含去重——去重在调用方汇总点原子执行)。
循环最多 max_retries 次尝试生成。每次尝试:
1. 采样锚节点
2. 构造 prompt
3. 调用 VLM
4. 解析响应
5. 构造 GeneratedQuestion
返回 None 表示耗尽重试。
参数:
vlm: VLM 调用端口。
embed_fn: 文本嵌入函数(本函数内未使用,由调用方统一去重)。
tree: 三层树索引。
video_id: 所属视频标识。
task_type: 12 种 Video-MME 题型之一。
seq: 序列号,用于生成 question_id。
exemplars: 少样本示例列表。
used_node_ids: 已用节点 ID 集合。
max_retries: 最大重试次数。
similarity_threshold: 余弦相似度阈值(本函数内未使用)。
rng: 可控随机数生成器。
session_id: 会话 ID(传递给 VLM 遥测)。
返回:
GeneratedQuestion 实例,或 None(耗尽重试)。
"""
from core.types import GeneratedQuestion as _GeneratedQuestion
for attempt in range(max_retries):
try:
# Phase 1: 采样锚节点
anchor = sample_anchor(tree, task_type, used_node_ids, rng)
# Phase 2: 构造 prompt
messages, images = build_generation_prompt(task_type, anchor, exemplars)
# Phase 3: 调用 VLM
response = await vlm.chat_with_images(
messages,
images,
session_id=session_id,
)
# Phase 4: 解析响应
parsed = parse_vlm_response(response.content, video_id, task_type, seq)
# Phase 5: 构造 GeneratedQuestion
return _GeneratedQuestion(
question_id=parsed["question_id"],
video_id=video_id,
task_type=task_type,
question=parsed["question"],
options=tuple(parsed["options"]),
answer=parsed["answer"],
source_nodes=(anchor.node_id,),
difficulty="medium",
)
except (ValueError, KeyError) as exc:
logger.warning(
"generate_one 尝试 {}/{} 失败 ({}): {}",
attempt + 1,
max_retries,
task_type,
exc,
)
continue
logger.warning(
"generate_one 耗尽 {} 次重试 (video={}, task_type={})",
max_retries,
video_id,
task_type,
)
return None