feat(cli): wire train mode with PoolStrategy selection and task_types in RunConfig

- Add --pool-split-mode, --train-ratio, --test-questions CLI arguments
- Include task_types in cli_overrides (convert list to tuple for RunConfig)
- Wire train mode branch: select strategy based on pool_split_mode,
  resolve workspace paths, build pools, and call runner.train()
- Infer mode now reads task_types from config instead of raw args

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-07-12 22:56:18 -04:00
parent c66a00c924
commit 5563d0d675
+29 -3
View File
@@ -86,6 +86,7 @@ def _build_adapters(settings: InfraSettings, embed_cfg: dict) -> _Adapters:
if settings.redis_url: if settings.redis_url:
try: try:
import redis.asyncio as aioredis import redis.asyncio as aioredis
from adapters.redis_cache import RedisResponseCache from adapters.redis_cache import RedisResponseCache
redis_client = aioredis.from_url(settings.redis_url, decode_responses=True) redis_client = aioredis.from_url(settings.redis_url, decode_responses=True)
@@ -195,6 +196,13 @@ def _build_parser() -> argparse.ArgumentParser:
parser.add_argument("--fresh", action="store_true", dest="fresh") parser.add_argument("--fresh", action="store_true", dest="fresh")
parser.add_argument("--seed", type=str, dest="seed") parser.add_argument("--seed", type=str, dest="seed")
parser.add_argument("--epochs", type=int) parser.add_argument("--epochs", type=int)
parser.add_argument(
"--pool-split-mode",
choices=["global", "per_category"],
dest="pool_split_mode",
)
parser.add_argument("--train-ratio", type=float, dest="train_ratio")
parser.add_argument("--test-questions", type=str, dest="test_questions")
return parser return parser
@@ -255,7 +263,10 @@ def main() -> None:
from app.harness.config import load_config from app.harness.config import load_config
cli_overrides = {k: v for k, v in vars(args).items() if k not in ("config", "task_types")} cli_args = vars(args)
if cli_args.get("task_types") is not None:
cli_args["task_types"] = tuple(cli_args["task_types"])
cli_overrides = {k: v for k, v in cli_args.items() if k != "config"}
config = load_config(args.config, cli_overrides) config = load_config(args.config, cli_overrides)
logger.info("配置加载完成: mode={}, workspace={}", config.mode, config.workspace_dir) logger.info("配置加载完成: mode={}, workspace={}", config.mode, config.workspace_dir)
@@ -291,9 +302,24 @@ def main() -> None:
) )
if config.mode == "infer": if config.mode == "infer":
task_types = getattr(args, "task_types", None) result = asyncio.run(runner.infer(task_types=config.task_types))
result = asyncio.run(runner.infer(task_types=task_types))
_log_result(result) _log_result(result)
elif config.mode == "train":
from app.harness.pools import (
GlobalPoolStrategy,
PerCategoryPoolStrategy,
build_or_load_pools,
)
from app.harness.workspace import resolve_paths
strategy = (
PerCategoryPoolStrategy()
if config.pool_split_mode == "per_category"
else GlobalPoolStrategy()
)
paths = resolve_paths(config.workspace_dir)
pools = build_or_load_pools(config, strategy, paths.db_path)
asyncio.run(runner.train(pools))
else: else:
raise SystemExit(f"模式 {config.mode!r} 尚未实现") raise SystemExit(f"模式 {config.mode!r} 尚未实现")