()
| 2260 | # ==================== 主函数 ==================== |
| 2261 | |
| 2262 | def main(): |
| 2263 | parser = argparse.ArgumentParser(description="Simulate Historical Episodes with Interest Drift") |
| 2264 | parser.add_argument("--start-date", type=str, required=True, help="开始日期 YYYYMMDD") |
| 2265 | parser.add_argument("--end-date", type=str, required=True, help="结束日期 YYYYMMDD") |
| 2266 | parser.add_argument( |
| 2267 | "--llm-model", |
| 2268 | type=str, |
| 2269 | default=os.environ.get("LLM_PARSER_OPENAI_MODEL", "gemini-3-flash-preview"), |
| 2270 | help="LLM model label used in benchmark summaries", |
| 2271 | ) |
| 2272 | parser.add_argument("--embedding-model", type=str, default="Qwen/Qwen3-Embedding-8B", help="Embedding 模型") |
| 2273 | parser.add_argument("--drift-probability", type=float, default=0.5, help="漂移触发概率 (0-1)") |
| 2274 | parser.add_argument("--sources", nargs="*", default=None, help="每天收集的论文源 (arxiv, openreview, journal)") |
| 2275 | parser.add_argument("--limit-per-source", type=int, default=None, help="每天每个源的可选上限") |
| 2276 | parser.add_argument("--skip-paper-collection", action="store_true", help="跳过每日论文收集,只使用数据库已有论文池") |
| 2277 | parser.add_argument("--output-dir", type=str, default=None, help="输出目录") |
| 2278 | parser.add_argument("--seed", type=int, default=42, help="随机种子") |
| 2279 | parser.add_argument("--skip-reading-reports", action="store_true", help="Skip reading report generation for recommendation-only simulation") |
| 2280 | parser.add_argument("--show-count", type=int, default=DEFAULT_SIMULATION_SHOW_COUNT, help="Displayed papers per episode after real-ranking fallback fill") |
| 2281 | parser.add_argument("--user-count", type=int, default=None, help="Run only the first N users after stable user-id sorting") |
| 2282 | parser.add_argument("--user-ids", nargs="*", default=None, help="Run only these explicit user ids, e.g. user_role1 user_role9 user_role24") |
| 2283 | args = parser.parse_args() |
| 2284 | |
| 2285 | random.seed(args.seed) |
| 2286 | embedding_module._default_service = None |
| 2287 | if hasattr(reading_agent, "READING_REPORT_EVIDENCE_CACHE_ENABLED"): |
| 2288 | reading_agent.READING_REPORT_EVIDENCE_CACHE_ENABLED = False |
| 2289 | _patch_real_usage_logging() |
| 2290 | |
| 2291 | start_date = datetime.strptime(args.start_date, "%Y%m%d") |
| 2292 | end_date = datetime.strptime(args.end_date, "%Y%m%d") |
| 2293 | |
| 2294 | print(f"Simulating from {start_date.date()} to {end_date.date()}") |
| 2295 | print(f"Drift probability: {args.drift_probability}") |
| 2296 | print() |
| 2297 | |
| 2298 | # 初始化 |
| 2299 | conn = sqlite3.connect(DB_PATH) |
| 2300 | users = get_all_users(conn) |
| 2301 | papers = get_all_papers(conn) |
| 2302 | conn.close() |
| 2303 | users = select_users(users, user_ids=args.user_ids, user_count=args.user_count) |
| 2304 | |
| 2305 | print(f"Loaded {len(users)} users, {len(papers)} papers") |
| 2306 | print(f"Selected users: {', '.join(str(user.get('user_id')) for user in users)}") |
| 2307 | |
| 2308 | # 加载漂移 checkfile |
| 2309 | checkfiles = load_checkfiles(DRIFT_CHECKFILES_DIR) |
| 2310 | print(f"Loaded {len(checkfiles)} drift checkfiles") |
| 2311 | |
| 2312 | drift_engine = DriftEngine(checkfiles) |
| 2313 | |
| 2314 | # 输出管理器 |
| 2315 | output_dir = Path(args.output_dir) if args.output_dir else (PROJECT_ROOT / "data" / "simulation_output") |
| 2316 | if not output_dir.is_absolute(): |
| 2317 | output_dir = PROJECT_ROOT / output_dir |
| 2318 | resume_state = load_resume_state(output_dir, start_date) |
| 2319 | if resume_state["resume"]: |
no test coverage detected