feat: wire real agent runner, backfill assembly and adversarial-filter CLI
This commit is contained in:
@@ -1173,6 +1173,176 @@ async def _run_generate_v2(args: argparse.Namespace) -> None:
|
||||
logger.warning("超过 50% 的 slot 被拒绝,建议检查 VLM/门控配置")
|
||||
|
||||
|
||||
def _load_trees_abs(videos_dir: Path, video_ids: list[str]) -> dict:
|
||||
"""加载视频树并将相对帧路径解析为绝对路径(复用 generate-v2 逻辑)。
|
||||
|
||||
参数:
|
||||
videos_dir: store/videos 目录。
|
||||
video_ids: 待加载的 video_id 列表。
|
||||
|
||||
返回:
|
||||
video_id → TreeIndex 映射(加载失败的视频被跳过)。
|
||||
"""
|
||||
from app.tree.index import TreeIndex
|
||||
|
||||
trees: dict = {}
|
||||
for vid in video_ids:
|
||||
tree_path = videos_dir / vid / "tree.json"
|
||||
try:
|
||||
tree = TreeIndex.load_json(str(tree_path))
|
||||
except (OSError, ValueError, KeyError) as exc:
|
||||
logger.warning("加载树 {} 失败,跳过: {}", tree_path, exc)
|
||||
continue
|
||||
video_dir = videos_dir / vid
|
||||
for l1 in tree.roots:
|
||||
for l2 in l1.children:
|
||||
for l3 in l2.children:
|
||||
if l3.frame_path and not Path(l3.frame_path).is_absolute():
|
||||
l3.frame_path = str(video_dir / l3.frame_path)
|
||||
trees[vid] = tree
|
||||
return trees
|
||||
|
||||
|
||||
def _add_adversarial_filter_parser(subparsers: argparse._SubParsersAction) -> None:
|
||||
"""注册 adversarial-filter 子命令(Phase B 后置对抗过滤 CLI 入口)。
|
||||
|
||||
参数:
|
||||
subparsers: argparse 子命令注册器。
|
||||
"""
|
||||
p = subparsers.add_parser(
|
||||
"adversarial-filter",
|
||||
help="Phase B 后置对抗过滤(作弊门 + 翻转门 + 缺额补生成)",
|
||||
)
|
||||
p.add_argument("--config", type=Path, required=True, help="YAML 配置(含 question_gen_v2 / adversarial_filter / embed 段)")
|
||||
p.add_argument("--store-dir", type=Path, required=True, help="store 根目录(含 videos/ prompts/ skills/)")
|
||||
p.add_argument("--accepted-path", type=Path, required=True, help="Phase A 产物 accepted_questions.json 路径(只读)")
|
||||
p.add_argument("--final-path", type=Path, default=None, help="最终题库输出路径(默认 accepted 同目录 accepted_questions_final.json)")
|
||||
p.add_argument("--db-path", type=Path, default=Path("logs/question_gen.db"), help="QuestionGenStore SQLite 路径")
|
||||
p.add_argument("--harness-db", type=Path, default=Path("logs/adversarial_harness.db"), help="agent 推理 HarnessLog SQLite 路径")
|
||||
p.add_argument("--prompts-version", type=str, default="v1", help="推理 prompt 版本目录名(store/prompts/<version>)")
|
||||
p.add_argument("--skills-version", type=str, default="v1", help="推理 skill 版本目录名(store/skills/<version>)")
|
||||
p.add_argument("--skill-mode", type=str, choices=["auto", "manual", "none"], default="auto", help="skill 模式")
|
||||
p.add_argument("--concurrency", type=int, default=4, help="agent 推理并发数")
|
||||
p.add_argument("--session-id", type=str, default="adversarial", help="遥测会话 ID(派生各门 run_id)")
|
||||
|
||||
|
||||
async def _run_adversarial_filter(args: argparse.Namespace) -> None:
|
||||
"""adversarial-filter 子命令主流程。
|
||||
|
||||
装配 adapters(同 main._build_adapters)、InferenceDepsRouter(同 main.py 参数)、
|
||||
QuestionGenStore、视频树(帧路径绝对化)、真实 _RealAgentRunner 与真实 backfill,
|
||||
调 run_adversarial_filter 跑两门 + 缺额补生成迭代。
|
||||
|
||||
参数:
|
||||
args: CLI 参数(config, store_dir, accepted_path, final_path, db_path,
|
||||
harness_db, prompts_version, skills_version, skill_mode, concurrency,
|
||||
session_id)。
|
||||
"""
|
||||
import yaml
|
||||
|
||||
from app.harness.deps_router import InferenceDepsRouter
|
||||
from app.question_gen.adversarial_config import load_adversarial_config
|
||||
from app.question_gen.adversarial_filter import (
|
||||
_RealAgentRunner,
|
||||
build_backfill,
|
||||
run_adversarial_filter,
|
||||
)
|
||||
from app.question_gen.pipeline_v2 import load_pipeline_config
|
||||
from app.question_gen.run_store import QuestionGenStore
|
||||
from main import InfraSettings, _build_adapters
|
||||
|
||||
config_path = args.config.resolve()
|
||||
store_dir = args.store_dir.resolve()
|
||||
|
||||
# Phase 1: 加载配置(对抗过滤 + 出题管线 + embed 段)
|
||||
with config_path.open(encoding="utf-8") as f:
|
||||
raw_yaml = yaml.safe_load(f) or {}
|
||||
embed_cfg = raw_yaml.get("embed", {})
|
||||
filter_config = load_adversarial_config(config_path)
|
||||
pipeline_config = load_pipeline_config(config_path)
|
||||
|
||||
# Phase 2: 装配 adapters
|
||||
settings = InfraSettings()
|
||||
adapters = _build_adapters(settings, embed_cfg)
|
||||
|
||||
# Phase 3: 加载视频树(帧路径绝对化)
|
||||
videos_dir = store_dir / "videos"
|
||||
if not videos_dir.exists():
|
||||
logger.error("视频目录不存在: {}", videos_dir)
|
||||
sys.exit(1)
|
||||
video_ids = sorted(
|
||||
d.name for d in videos_dir.iterdir() if d.is_dir() and (d / "tree.json").exists()
|
||||
)
|
||||
trees = _load_trees_abs(videos_dir, video_ids)
|
||||
if not trees:
|
||||
logger.error("所有视频树加载失败,无法继续")
|
||||
sys.exit(1)
|
||||
logger.info("成功加载 {} / {} 棵视频树", len(trees), len(video_ids))
|
||||
|
||||
# Phase 4: InferenceDepsRouter(同 main.py 参数)
|
||||
router = InferenceDepsRouter(
|
||||
store_dir=store_dir,
|
||||
embed_provider=adapters.embed,
|
||||
llm=adapters.llm,
|
||||
vlm=adapters.vlm,
|
||||
ocr=adapters.ocr,
|
||||
default_prompts_dir=store_dir / "prompts" / args.prompts_version,
|
||||
default_skills_dir=store_dir / "skills" / args.skills_version,
|
||||
skill_mode=args.skill_mode,
|
||||
verify_vision=True,
|
||||
anchor=True,
|
||||
assemble_mode="ids_expand",
|
||||
)
|
||||
|
||||
# Phase 5: QuestionGenStore + 真实 agent + 真实 backfill
|
||||
db_path = args.db_path.resolve()
|
||||
db_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
store = QuestionGenStore(str(db_path))
|
||||
|
||||
harness_db = args.harness_db.resolve()
|
||||
harness_db.parent.mkdir(parents=True, exist_ok=True)
|
||||
agent = _RealAgentRunner(
|
||||
llm=adapters.llm,
|
||||
tool_dispatch_fn=router.create_dispatch(),
|
||||
prompt_builder=router.create_prompt_builder(),
|
||||
db_path=str(harness_db),
|
||||
concurrency=args.concurrency,
|
||||
skill_mode=args.skill_mode,
|
||||
model=settings.search_llm_model,
|
||||
)
|
||||
backfill = build_backfill(
|
||||
trees=trees,
|
||||
vlm=adapters.vlm,
|
||||
llm=adapters.llm,
|
||||
embed_fn=adapters.embed.embed,
|
||||
store=store,
|
||||
pipeline_config=pipeline_config,
|
||||
filter_task_types=filter_config.filter_task_types,
|
||||
session_id=args.session_id,
|
||||
)
|
||||
|
||||
# Phase 6: 运行对抗过滤
|
||||
accepted_path = args.accepted_path.resolve()
|
||||
final_path = (
|
||||
args.final_path.resolve()
|
||||
if args.final_path is not None
|
||||
else accepted_path.parent / "accepted_questions_final.json"
|
||||
)
|
||||
await run_adversarial_filter(
|
||||
accepted_path=accepted_path,
|
||||
final_path=final_path,
|
||||
agent=agent,
|
||||
vlm=adapters.vlm,
|
||||
trees=trees,
|
||||
store=store,
|
||||
filter_config=filter_config,
|
||||
backfill=backfill,
|
||||
session_id=args.session_id,
|
||||
)
|
||||
store.close()
|
||||
logger.info("对抗过滤完成,final 已写入: {}", final_path)
|
||||
|
||||
|
||||
def _parse_args() -> argparse.Namespace:
|
||||
"""解析命令行参数。"""
|
||||
parser = argparse.ArgumentParser(description="赛题生成工具:generate + calibrate + generate-v2")
|
||||
@@ -1181,6 +1351,9 @@ def _parse_args() -> argparse.Namespace:
|
||||
# generate-v2 子命令
|
||||
_add_generate_v2_parser(subparsers)
|
||||
|
||||
# adversarial-filter 子命令(Phase B 后置对抗过滤)
|
||||
_add_adversarial_filter_parser(subparsers)
|
||||
|
||||
# generate 子命令
|
||||
gen_parser = subparsers.add_parser("generate", help="生成新题目(v1 传统模式)")
|
||||
gen_parser.add_argument(
|
||||
@@ -1280,6 +1453,8 @@ def main() -> None:
|
||||
asyncio.run(_run_generate(args))
|
||||
elif args.command == "generate-v2":
|
||||
asyncio.run(_run_generate_v2(args))
|
||||
elif args.command == "adversarial-filter":
|
||||
asyncio.run(_run_adversarial_filter(args))
|
||||
elif args.command == "calibrate":
|
||||
_run_calibrate(args)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user