feat: add AdversarialFilterConfig for Phase B post-hoc filter layer

This commit is contained in:
2026-07-14 15:47:21 -04:00
parent ac115d96fb
commit 334fbbc94d
3 changed files with 104 additions and 0 deletions
+59
View File
@@ -0,0 +1,59 @@
"""Phase B 对抗过滤层配置 — filter 层配置(非 strategy 属性)。
设计: research-wiki/designs/2026-07-14-adversarial-question-gen-phaseB-design.md §8
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING
import yaml
if TYPE_CHECKING:
from pathlib import Path
@dataclass(frozen=True)
class AdversarialFilterConfig:
"""后置对抗过滤配置。
属性:
filter_task_types: 被过滤的题型(仅这些走 agent 门),默认仅 AR。
adversarial_max_rounds: 补生成迭代上限。
adversarial_agent_max_steps: agent 试答步数上限。
difficulty_warn_threshold: 批次 agent 正确率告警阈值。
"""
filter_task_types: tuple[str, ...] = ("Action Recognition",)
adversarial_max_rounds: int = 5
adversarial_agent_max_steps: int = 40
difficulty_warn_threshold: float = 0.85
def load_adversarial_config(config_path: Path) -> AdversarialFilterConfig:
"""从 YAML 的 adversarial_filter 区段加载配置,缺段/缺键用默认值。
参数:
config_path: YAML 配置文件路径。
返回:
AdversarialFilterConfig 实例。
"""
with open(config_path, encoding="utf-8") as f:
raw = yaml.safe_load(f) or {}
section = raw.get("adversarial_filter", {}) or {}
default = AdversarialFilterConfig()
types = section.get("filter_task_types")
return AdversarialFilterConfig(
filter_task_types=tuple(types) if types else default.filter_task_types,
adversarial_max_rounds=int(
section.get("adversarial_max_rounds", default.adversarial_max_rounds)
),
adversarial_agent_max_steps=int(
section.get("adversarial_agent_max_steps", default.adversarial_agent_max_steps)
),
difficulty_warn_threshold=float(
section.get("difficulty_warn_threshold", default.difficulty_warn_threshold)
),
)
+7
View File
@@ -12,3 +12,10 @@ question_gen_v2:
candidate_pool_size: 24 # grounded selector 首轮候选干扰项数 N
selector_delta_low: 0.05 # 干扰项视觉分与正解的最小差(区间上界)
selector_delta_high: 0.35 # 干扰项视觉分与正解的最大差(区间下界)
# Phase B 后置对抗过滤层配置(filter 层,非 strategy 属性)
adversarial_filter:
filter_task_types: [Action Recognition] # 仅这些题型走 agent 门(路径隔离开关)
adversarial_max_rounds: 5 # 补生成迭代上限
adversarial_agent_max_steps: 40 # agent 试答步数上限
difficulty_warn_threshold: 0.85 # 批次 agent 正确率告警阈值
+38
View File
@@ -0,0 +1,38 @@
"""AdversarialFilterConfig 默认值 + YAML 加载。"""
from app.question_gen.adversarial_config import (
AdversarialFilterConfig,
load_adversarial_config,
)
def test_defaults():
cfg = AdversarialFilterConfig()
assert cfg.filter_task_types == ("Action Recognition",)
assert cfg.adversarial_max_rounds == 5
assert cfg.adversarial_agent_max_steps == 40
assert cfg.difficulty_warn_threshold == 0.85
def test_load_from_yaml(tmp_path):
p = tmp_path / "c.yaml"
p.write_text(
"adversarial_filter:\n"
" filter_task_types: [Action Recognition, Object Recognition]\n"
" adversarial_max_rounds: 3\n"
" adversarial_agent_max_steps: 20\n"
" difficulty_warn_threshold: 0.7\n",
encoding="utf-8",
)
cfg = load_adversarial_config(p)
assert cfg.filter_task_types == ("Action Recognition", "Object Recognition")
assert cfg.adversarial_max_rounds == 3
assert cfg.adversarial_agent_max_steps == 20
assert cfg.difficulty_warn_threshold == 0.7
def test_load_missing_section_uses_defaults(tmp_path):
p = tmp_path / "c.yaml"
p.write_text("question_gen_v2:\n per_type: 3\n", encoding="utf-8")
cfg = load_adversarial_config(p)
assert cfg == AdversarialFilterConfig()