c49d0ff12f
Add ValueError guard at the top of sample_material_v2 for level not in
{1, 2, 3}, preventing silent fallthrough to L1 sampling. Add unit test
for the new validation.
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
598 lines
17 KiB
Python
598 lines
17 KiB
Python
"""v2 素材采样器 — 基于采样约束的树节点采样与上下文收集。
|
||
|
||
在 v1 synthesizer 的基础上引入 SamplingConstraint 约束验证,
|
||
为每次出题提供更丰富的素材上下文(字幕、跨 L2 上下文、帧路径)。
|
||
|
||
典型调用路径::
|
||
|
||
material = sample_material_v2(
|
||
tree=tree_index,
|
||
task_type="Action Reasoning",
|
||
used_node_ids=already_used,
|
||
rng=rng,
|
||
level=2,
|
||
constraint=my_constraint,
|
||
)
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from dataclasses import dataclass
|
||
from typing import TYPE_CHECKING
|
||
|
||
from loguru import logger
|
||
|
||
if TYPE_CHECKING:
|
||
import random
|
||
|
||
from app.question_gen.families import SamplingConstraint
|
||
from app.tree.index import L1Node, L2Node, TreeIndex
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 题型 → 采样层级映射
|
||
# ---------------------------------------------------------------------------
|
||
|
||
_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,
|
||
"Object Reasoning": 2,
|
||
"Information Synopsis": 2,
|
||
"Counting Problem": 2,
|
||
# Level 1(段落/场景级)
|
||
"Temporal Reasoning": 1,
|
||
"Temporal Perception": 1,
|
||
"Spatial Reasoning": 1,
|
||
"Spatial Perception": 1,
|
||
}
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 数据类型
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class AnchorContext:
|
||
"""采样锚点上下文。
|
||
|
||
属性:
|
||
node_id: 锚节点 ID。
|
||
level: 锚节点所在层级(1/2/3)。
|
||
l2_id: 锚节点所属的 L2 节点 ID(若自身为 L2 则等于 node_id;
|
||
若为 L1 则取其第一个 L2 子节点 ID)。
|
||
"""
|
||
|
||
node_id: str
|
||
level: int
|
||
l2_id: str
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class MaterialContext:
|
||
"""采样素材上下文 — 出题所需的全部素材打包。
|
||
|
||
属性:
|
||
anchor: 采样锚点信息。
|
||
source_nodes: 参与采样的节点 ID 元组。
|
||
subtitle_sentences: 锚节点子树中收集的字幕句列表。
|
||
frame_paths: 锚节点子树中可用的帧路径列表。
|
||
cross_l2_texts: 跨 L2 段的上下文文本列表(仅 cross_l2_span 时填充)。
|
||
"""
|
||
|
||
anchor: AnchorContext
|
||
source_nodes: tuple[str, ...]
|
||
subtitle_sentences: list[str]
|
||
frame_paths: list[str]
|
||
cross_l2_texts: list[str]
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 内部索引辅助
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class _ResolvedSubtree:
|
||
"""节点解析结果 — 将任意层级 node_id 统一解析为 L2 节点列表 + 父 L1。
|
||
|
||
属性:
|
||
l2_nodes: 与目标节点关联的 L2 节点列表。
|
||
parent_l1: 父 L1 节点(用于跨 L2 判断)。
|
||
"""
|
||
|
||
l2_nodes: list[L2Node]
|
||
parent_l1: L1Node | None
|
||
|
||
|
||
def _find_l2_node(tree: TreeIndex, l2_id: str) -> tuple[L2Node, L1Node] | None:
|
||
"""按 ID 定位 L2 节点及其父 L1。
|
||
|
||
参数:
|
||
tree: 三层树索引。
|
||
l2_id: L2 节点 ID。
|
||
|
||
返回:
|
||
(L2Node, 父L1Node) 元组;未找到返回 None。
|
||
"""
|
||
for l1 in tree.roots:
|
||
for l2 in l1.children:
|
||
if l2.id == l2_id:
|
||
return (l2, l1)
|
||
return None
|
||
|
||
|
||
def _find_l1_node(tree: TreeIndex, l1_id: str) -> L1Node | None:
|
||
"""按 ID 定位 L1 节点。
|
||
|
||
参数:
|
||
tree: 三层树索引。
|
||
l1_id: L1 节点 ID。
|
||
|
||
返回:
|
||
L1Node;未找到返回 None。
|
||
"""
|
||
for l1 in tree.roots:
|
||
if l1.id == l1_id:
|
||
return l1
|
||
return None
|
||
|
||
|
||
def _find_l3_parent(tree: TreeIndex, l3_id: str) -> tuple[L2Node, L1Node] | None:
|
||
"""按 L3 ID 定位其所属的 L2 和 L1 节点。
|
||
|
||
参数:
|
||
tree: 三层树索引。
|
||
l3_id: L3 节点 ID。
|
||
|
||
返回:
|
||
(L2Node, L1Node) 元组;未找到返回 None。
|
||
"""
|
||
for l1 in tree.roots:
|
||
for l2 in l1.children:
|
||
for l3 in l2.children:
|
||
if l3.id == l3_id:
|
||
return (l2, l1)
|
||
return None
|
||
|
||
|
||
def _resolve_subtree(tree: TreeIndex, node_id: str) -> _ResolvedSubtree | None:
|
||
"""将任意层级节点 ID 解析为关联的 L2 节点列表和父 L1。
|
||
|
||
参数:
|
||
tree: 三层树索引。
|
||
node_id: 任意层级的节点 ID。
|
||
|
||
返回:
|
||
_ResolvedSubtree 实例;节点不存在时返回 None。
|
||
"""
|
||
# 尝试作为 L2
|
||
result = _find_l2_node(tree, node_id)
|
||
if result is not None:
|
||
l2_node, parent_l1 = result
|
||
return _ResolvedSubtree(l2_nodes=[l2_node], parent_l1=parent_l1)
|
||
|
||
# 尝试作为 L1
|
||
l1_node = _find_l1_node(tree, node_id)
|
||
if l1_node is not None:
|
||
return _ResolvedSubtree(l2_nodes=list(l1_node.children), parent_l1=l1_node)
|
||
|
||
# 尝试作为 L3
|
||
l3_result = _find_l3_parent(tree, node_id)
|
||
if l3_result is not None:
|
||
l2_node, parent_l1 = l3_result
|
||
return _ResolvedSubtree(l2_nodes=[l2_node], parent_l1=parent_l1)
|
||
|
||
return None
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 约束检查辅助(单一职责)
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
def _count_l3_descendants(l2_nodes: list[L2Node]) -> int:
|
||
"""统计 L2 节点列表下的 L3 总数。
|
||
|
||
参数:
|
||
l2_nodes: L2 节点列表。
|
||
|
||
返回:
|
||
L3 节点总数。
|
||
"""
|
||
return sum(len(l2.children) for l2 in l2_nodes)
|
||
|
||
|
||
def _has_frames(l2_nodes: list[L2Node]) -> bool:
|
||
"""检查 L2 节点列表的子树中是否有可用帧。
|
||
|
||
参数:
|
||
l2_nodes: L2 节点列表。
|
||
|
||
返回:
|
||
True 表示至少有一个 L3 有 frame_path。
|
||
"""
|
||
return any(l3.frame_path for l2 in l2_nodes for l3 in l2.children)
|
||
|
||
|
||
def _count_subtitles(l2_nodes: list[L2Node]) -> int:
|
||
"""统计 L2 节点列表中全部字幕数(L2 + L3)。
|
||
|
||
参数:
|
||
l2_nodes: L2 节点列表。
|
||
|
||
返回:
|
||
非空字幕总数。
|
||
"""
|
||
count = 0
|
||
for l2 in l2_nodes:
|
||
if l2.card.subtitle:
|
||
count += 1
|
||
count += sum(1 for l3 in l2.children if l3.card.subtitle)
|
||
return count
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 公开辅助函数
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
def _validate_sampling_constraints(
|
||
tree: TreeIndex, node_id: str, constraint: SamplingConstraint
|
||
) -> bool:
|
||
"""校验指定节点是否满足采样约束。
|
||
|
||
根据节点层级自动判断检查范围:
|
||
- L2 节点:检查其子 L3 的帧/字幕数量。
|
||
- L1 节点:检查其下全部 L2/L3 的帧/字幕总数。
|
||
- L3 节点:检查其所属 L2 的子树。
|
||
|
||
参数:
|
||
tree: 三层树索引。
|
||
node_id: 待检查节点 ID。
|
||
constraint: 采样约束条件。
|
||
|
||
返回:
|
||
True 表示满足所有约束,False 表示至少一项不满足。
|
||
"""
|
||
resolved = _resolve_subtree(tree, node_id)
|
||
if resolved is None:
|
||
return False
|
||
|
||
if _count_l3_descendants(resolved.l2_nodes) < constraint.min_l3_nodes:
|
||
return False
|
||
if constraint.require_frames and not _has_frames(resolved.l2_nodes):
|
||
return False
|
||
if _count_subtitles(resolved.l2_nodes) < constraint.min_subtitles:
|
||
return False
|
||
return not (
|
||
constraint.cross_l2_span
|
||
and (resolved.parent_l1 is None or len(resolved.parent_l1.children) < 2)
|
||
)
|
||
|
||
|
||
def _subtitles_from_l2_list(l2_nodes: list[L2Node]) -> list[str]:
|
||
"""从 L2 节点列表收集全部非空字幕。
|
||
|
||
参数:
|
||
l2_nodes: L2 节点列表。
|
||
|
||
返回:
|
||
非空字幕字符串列表。
|
||
"""
|
||
sentences: list[str] = []
|
||
for l2 in l2_nodes:
|
||
if l2.card.subtitle:
|
||
sentences.append(l2.card.subtitle)
|
||
for l3 in l2.children:
|
||
if l3.card.subtitle:
|
||
sentences.append(l3.card.subtitle)
|
||
return sentences
|
||
|
||
|
||
def _collect_subtitle_sentences(tree: TreeIndex, node_ids: tuple[str, ...]) -> list[str]:
|
||
"""从指定节点集合中收集字幕句。
|
||
|
||
遍历每个 node_id 对应的子树,提取非空字幕。
|
||
对 L2 节点提取自身 + 子 L3 字幕;对 L1 提取下属全部。
|
||
|
||
参数:
|
||
tree: 三层树索引。
|
||
node_ids: 待收集字幕的节点 ID 元组。
|
||
|
||
返回:
|
||
非空字幕句列表。
|
||
"""
|
||
sentences: list[str] = []
|
||
for nid in node_ids:
|
||
resolved = _resolve_subtree(tree, nid)
|
||
if resolved is not None:
|
||
sentences.extend(_subtitles_from_l2_list(resolved.l2_nodes))
|
||
return sentences
|
||
|
||
|
||
def _collect_cross_l2_context(tree: TreeIndex, anchor_l2_id: str, max_peers: int = 3) -> list[str]:
|
||
"""收集锚 L2 的同级 L2 节点描述文本(跨 L2 上下文)。
|
||
|
||
找到锚 L2 所属的 L1 父节点,取该父节点下除锚 L2 之外的其他 L2 描述。
|
||
|
||
参数:
|
||
tree: 三层树索引。
|
||
anchor_l2_id: 锚 L2 节点 ID。
|
||
max_peers: 最多返回的同级 L2 描述数量。
|
||
|
||
返回:
|
||
同级 L2 的 event_description 列表(最多 max_peers 条)。
|
||
"""
|
||
result = _find_l2_node(tree, anchor_l2_id)
|
||
if result is None:
|
||
return []
|
||
|
||
_, parent_l1 = result
|
||
peers: list[str] = []
|
||
for l2 in parent_l1.children:
|
||
if l2.id != anchor_l2_id:
|
||
peers.append(l2.card.event_description)
|
||
if len(peers) >= max_peers:
|
||
break
|
||
|
||
return peers
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 层级采样策略
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
def _sample_l3_node(
|
||
tree: TreeIndex,
|
||
used_node_ids: set[str],
|
||
rng: random.Random,
|
||
) -> tuple[str, str] | None:
|
||
"""随机采样一个未使用的 L3 节点,返回 (l3_id, 所属l2_id)。
|
||
|
||
参数:
|
||
tree: 三层树索引。
|
||
used_node_ids: 已用节点 ID 集合。
|
||
rng: 随机数生成器。
|
||
|
||
返回:
|
||
(l3_id, l2_id) 元组;无候选返回 None。
|
||
"""
|
||
candidates: list[tuple[str, str]] = []
|
||
for l1 in tree.roots:
|
||
for l2 in l1.children:
|
||
for l3 in l2.children:
|
||
if l3.id not in used_node_ids:
|
||
candidates.append((l3.id, l2.id))
|
||
|
||
if not candidates:
|
||
return None
|
||
return rng.choice(candidates)
|
||
|
||
|
||
def _sample_l2_node(
|
||
tree: TreeIndex,
|
||
used_node_ids: set[str],
|
||
rng: random.Random,
|
||
) -> tuple[str, str] | None:
|
||
"""随机采样一个未使用的 L2 节点,返回 (l2_id, l2_id)。
|
||
|
||
参数:
|
||
tree: 三层树索引。
|
||
used_node_ids: 已用节点 ID 集合。
|
||
rng: 随机数生成器。
|
||
|
||
返回:
|
||
(l2_id, l2_id) 元组;无候选返回 None。
|
||
"""
|
||
candidates: list[str] = []
|
||
for l1 in tree.roots:
|
||
for l2 in l1.children:
|
||
if l2.id not in used_node_ids:
|
||
candidates.append(l2.id)
|
||
|
||
if not candidates:
|
||
return None
|
||
chosen = rng.choice(candidates)
|
||
return (chosen, chosen)
|
||
|
||
|
||
def _sample_l1_node(
|
||
tree: TreeIndex,
|
||
used_node_ids: set[str],
|
||
rng: random.Random,
|
||
) -> tuple[str, str] | None:
|
||
"""随机采样一个未使用的 L1 节点,返回 (l1_id, 首个子l2_id)。
|
||
|
||
参数:
|
||
tree: 三层树索引。
|
||
used_node_ids: 已用节点 ID 集合。
|
||
rng: 随机数生成器。
|
||
|
||
返回:
|
||
(l1_id, first_l2_id) 元组;无候选返回 None。
|
||
"""
|
||
candidates: list[tuple[str, str]] = []
|
||
for l1 in tree.roots:
|
||
if l1.id not in used_node_ids and l1.children:
|
||
candidates.append((l1.id, l1.children[0].id))
|
||
|
||
if not candidates:
|
||
return None
|
||
return rng.choice(candidates)
|
||
|
||
|
||
def _frames_from_l2_list(l2_nodes: list[L2Node]) -> list[str]:
|
||
"""从 L2 节点列表收集全部可用帧路径。
|
||
|
||
参数:
|
||
l2_nodes: L2 节点列表。
|
||
|
||
返回:
|
||
帧路径字符串列表。
|
||
"""
|
||
return [l3.frame_path for l2 in l2_nodes for l3 in l2.children if l3.frame_path]
|
||
|
||
|
||
def _collect_frame_paths(tree: TreeIndex, node_id: str) -> list[str]:
|
||
"""收集节点子树下的所有可用帧路径。
|
||
|
||
参数:
|
||
tree: 三层树索引。
|
||
node_id: 目标节点 ID。
|
||
|
||
返回:
|
||
帧路径列表。
|
||
"""
|
||
resolved = _resolve_subtree(tree, node_id)
|
||
if resolved is None:
|
||
return []
|
||
return _frames_from_l2_list(resolved.l2_nodes)
|
||
|
||
|
||
def _collect_source_nodes(tree: TreeIndex, node_id: str) -> tuple[str, ...]:
|
||
"""收集节点子树涉及的全部节点 ID(包含自身)。
|
||
|
||
参数:
|
||
tree: 三层树索引。
|
||
node_id: 目标节点 ID。
|
||
|
||
返回:
|
||
相关节点 ID 元组。
|
||
"""
|
||
ids: list[str] = [node_id]
|
||
|
||
# L2 节点:加入子 L3
|
||
result = _find_l2_node(tree, node_id)
|
||
if result is not None:
|
||
l2_node, _ = result
|
||
for l3 in l2_node.children:
|
||
ids.append(l3.id)
|
||
return tuple(ids)
|
||
|
||
# L1 节点:加入子 L2 + L3
|
||
l1_node = _find_l1_node(tree, node_id)
|
||
if l1_node is not None:
|
||
for l2 in l1_node.children:
|
||
ids.append(l2.id)
|
||
for l3 in l2.children:
|
||
ids.append(l3.id)
|
||
return tuple(ids)
|
||
|
||
# L3 节点:仅自身
|
||
return tuple(ids)
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 主入口
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
def sample_material_v2(
|
||
tree: TreeIndex,
|
||
task_type: str,
|
||
used_node_ids: set[str],
|
||
rng: random.Random,
|
||
*,
|
||
level: int,
|
||
constraint: SamplingConstraint,
|
||
max_attempts: int = 10,
|
||
) -> MaterialContext:
|
||
"""基于采样约束从视频树中采样素材上下文。
|
||
|
||
采样流程:
|
||
1. 按指定 level 确定采样层级
|
||
2. 随机选取候选节点(排除 used_node_ids)
|
||
3. 验证 SamplingConstraint 约束
|
||
4. 约束不满足则重试(最多 max_attempts 次)
|
||
5. 收集字幕、帧路径、跨 L2 上下文
|
||
|
||
参数:
|
||
tree: 三层树索引。
|
||
task_type: 任务类型字符串。
|
||
used_node_ids: 本轮已用节点 ID 集合。
|
||
rng: 可控随机数生成器。
|
||
level: 采样层级(1/2/3)。
|
||
constraint: 采样约束条件。
|
||
max_attempts: 最大尝试次数。
|
||
|
||
返回:
|
||
MaterialContext 实例。
|
||
|
||
异常:
|
||
ValueError: level 不在 {1, 2, 3} 中。
|
||
RuntimeError: 耗尽 max_attempts 次尝试仍无法满足约束。
|
||
"""
|
||
if level not in (1, 2, 3):
|
||
msg = f"level 必须为 1、2 或 3,收到: {level}"
|
||
raise ValueError(msg)
|
||
|
||
for attempt in range(max_attempts):
|
||
# Phase 1: 按层级采样候选节点
|
||
if level == 3:
|
||
sampled = _sample_l3_node(tree, used_node_ids, rng)
|
||
elif level == 2:
|
||
sampled = _sample_l2_node(tree, used_node_ids, rng)
|
||
else:
|
||
sampled = _sample_l1_node(tree, used_node_ids, rng)
|
||
|
||
if sampled is None:
|
||
logger.debug(
|
||
"sample_material_v2 尝试 {}/{}: 无可用候选节点 (level={})",
|
||
attempt + 1,
|
||
max_attempts,
|
||
level,
|
||
)
|
||
continue
|
||
|
||
node_id, l2_id = sampled
|
||
|
||
# Phase 2: 验证约束
|
||
if not _validate_sampling_constraints(tree, node_id, constraint):
|
||
logger.debug(
|
||
"sample_material_v2 尝试 {}/{}: 约束违反 (node={})",
|
||
attempt + 1,
|
||
max_attempts,
|
||
node_id,
|
||
)
|
||
continue
|
||
|
||
# Phase 3: 构造 AnchorContext
|
||
anchor = AnchorContext(node_id=node_id, level=level, l2_id=l2_id)
|
||
|
||
# Phase 4: 收集素材
|
||
source_nodes = _collect_source_nodes(tree, node_id)
|
||
subtitle_sentences = _collect_subtitle_sentences(tree, (node_id,))
|
||
frame_paths = _collect_frame_paths(tree, node_id)
|
||
|
||
# Phase 5: 跨 L2 上下文(仅 cross_l2_span 时收集)
|
||
cross_l2_texts: list[str] = []
|
||
if constraint.cross_l2_span:
|
||
cross_l2_texts = _collect_cross_l2_context(tree, l2_id)
|
||
|
||
logger.debug(
|
||
"sample_material_v2 成功: node={}, level={}, attempt={}/{}",
|
||
node_id,
|
||
level,
|
||
attempt + 1,
|
||
max_attempts,
|
||
)
|
||
|
||
return MaterialContext(
|
||
anchor=anchor,
|
||
source_nodes=source_nodes,
|
||
subtitle_sentences=subtitle_sentences,
|
||
frame_paths=frame_paths,
|
||
cross_l2_texts=cross_l2_texts,
|
||
)
|
||
|
||
raise RuntimeError(
|
||
f"sample_material_v2: 耗尽 max_attempts={max_attempts} 次尝试,"
|
||
f"无法为 task_type='{task_type}' (level={level}) 满足采样约束"
|
||
)
|