"""M4 LLM 投资事件抽取批处理。 输入: data/deduped/{YYYYMMDD}/uniques/*.json (M3 唯一文章产物) 或者 data/processed/{source}/{YYYYMMDD}/*.json (M2 直出,跳过去重时使用) 输出: data/events/{YYYYMMDD}/{url_hash}.json (ExtractedEvent) data/events/{YYYYMMDD}/index.jsonl (扁平摘要) data/events/{YYYYMMDD}/failed.jsonl (失败列表) 增量: 默认跳过已抽取的文章(输出目录已有 {url_hash}.json 视为已处理), 断点续跑/失败重试不会重复调用 LLM API;--force 强制全量重抽。 用法: uv run python -m scripts.run_event_extraction uv run python -m scripts.run_event_extraction --date 20260616 uv run python -m scripts.run_event_extraction --provider qwen --model qwen-plus uv run python -m scripts.run_event_extraction --concurrency 5 --limit 10 uv run python -m scripts.run_event_extraction --input-root data/processed --no-deduped uv run python -m scripts.run_event_extraction --force # 全量重抽 """ from __future__ import annotations import argparse import asyncio import json import sys import time from collections import Counter from datetime import date from pathlib import Path from dotenv import load_dotenv from loguru import logger from pydantic import ValidationError from extractor import Article from llm import ( SCENE_EVENT_EXTRACTION, ExtractedEvent, LLMCallError, PromptTemplate, extract_event_async, load_llm_config, make_async_client, ) def _setup_logger(level: str) -> None: logger.remove() logger.add( sys.stderr, level=level, format="{time:YYYY-MM-DD HH:mm:ss} | {level} | {name} | {message}", ) log_path = Path("logs") / "llm.log" log_path.parent.mkdir(parents=True, exist_ok=True) logger.add(log_path, level="DEBUG", rotation="10 MB", retention=5, encoding="utf-8") def _collect_inputs( input_root: Path, day: str, use_deduped: bool, source_filter: str | None, ) -> list[Path]: """根据是否走 dedup 选择输入文件清单。""" if use_deduped: # data/deduped/{day}/uniques/*.json d = input_root / day / "uniques" if not d.is_dir(): return [] return sorted(d.glob("*.json")) # data/processed/{source}/{day}/*.json files: list[Path] = [] if source_filter: srcs = [source_filter] else: srcs = sorted(p.name for p in input_root.iterdir() if p.is_dir()) for src in srcs: d = input_root / src / day if d.is_dir(): files.extend(sorted(d.glob("*.json"))) return files def _filter_existing(files: list[Path], out_dir: Path) -> tuple[list[Path], int]: """过滤掉已有产物(输出目录存在同名 {url_hash}.json)的输入。 输入文件名即 url_hash(如 {url_hash}.json),与 M4 产物命名一致。 返回 (待处理文件, 跳过数);断点续跑/失败重试借此避免重复调用 LLM API。 """ pending: list[Path] = [] skipped = 0 for fp in files: if (out_dir / f"{fp.stem}.json").exists(): skipped += 1 else: pending.append(fp) if skipped: logger.info("跳过已处理 {} 篇(产物已存在),待处理 {}", skipped, len(pending)) return pending, skipped def _load_article(p: Path) -> Article | None: try: return Article.model_validate(json.loads(p.read_text(encoding="utf-8"))) except (json.JSONDecodeError, ValidationError) as e: logger.warning("跳过无法解析的 article 文件 {}: {}", p, e) return None def _load_sources(p: Path) -> list[str] | None: """从输入 JSON 读取 sources 多源字段(deduped uniques 才有);无则返回 None。 None 表示无多源记录,由 extract_event 兜底为 [source_id]。 """ try: data = json.loads(p.read_text(encoding="utf-8")) srcs = data.get("sources") if isinstance(srcs, list) and srcs: return [s for s in srcs if s] except (json.JSONDecodeError, OSError): pass return None async def _run(args: argparse.Namespace) -> int: load_dotenv() # 读 .env 到 os.environ # scene=event_extraction: 读取 configs/llm_models.yaml 场景 1 配置,未配置字段回退 .env config = load_llm_config( provider=args.provider, model=args.model, scene=SCENE_EVENT_EXTRACTION ) logger.info( "LLM provider={} model={} base_url={}", config.provider, config.model, config.base_url, ) input_root = Path(args.input_root) use_deduped = not args.no_deduped files = _collect_inputs(input_root, args.date, use_deduped, args.source) if not files: logger.error( "{} 下未发现 {} 的文章(use_deduped={})", input_root, args.date, use_deduped, ) return 2 out_dir = Path(args.out_root) / args.date out_dir.mkdir(parents=True, exist_ok=True) # 增量:跳过已有产物(断点续跑/失败重试不重复调用 LLM API),--force 全量 skipped = 0 if not args.force: files, skipped = _filter_existing(files, out_dir) if args.limit: files = files[: args.limit] if not files: logger.info( "无待处理文章(全部已抽取,跳过 {} 篇),如需重抽请加 --force", skipped ) return 0 logger.info("待处理文章数: {} (跳过已处理 {})", len(files), skipped) index_path = out_dir / "index.jsonl" failed_path = out_dir / "failed.jsonl" if args.force: # 全量模式:重建 index / failed for p in (index_path, failed_path): if p.exists(): p.unlink() else: # 增量模式:index 累积追加;failed 只保留本次运行失败的 if failed_path.exists(): failed_path.unlink() template = PromptTemplate(args.prompt) semaphore = asyncio.Semaphore(args.concurrency) succ_cnt = 0 fail_cnt = 0 sentiment_cnt: Counter = Counter() layer_attempts: Counter = Counter() started = time.time() async with make_async_client(config) as client: async def _process(fp: Path) -> tuple[Path, ExtractedEvent | None, str | None]: article = _load_article(fp) if article is None: return fp, None, "无法解析输入" sources = _load_sources(fp) # 多源记录(去重层),无则 None try: event = await extract_event_async( client=client, config=config, article=article, template=template, max_attempts=args.max_attempts, sources=sources, semaphore=semaphore, ) return fp, event, None except LLMCallError as e: return fp, None, str(e) tasks = [_process(fp) for fp in files] for coro in asyncio.as_completed(tasks): fp, event, err = await coro if event is None: fail_cnt += 1 with failed_path.open("a", encoding="utf-8") as f: f.write( json.dumps( {"file": str(fp), "error": err}, ensure_ascii=False, ) + "\n" ) continue succ_cnt += 1 sentiment_cnt[event.event.sentiment.value] += 1 layer_attempts[event.attempts] += 1 event_path = out_dir / f"{event.url_hash}.json" event_path.write_text(event.model_dump_json(indent=2), encoding="utf-8") meta = event.model_dump(mode="json") # 扁平 index 不含完整 content,但保留 event 字段 with index_path.open("a", encoding="utf-8") as f: f.write(json.dumps(meta, ensure_ascii=False) + "\n") logger.info(event.short_summary()) elapsed = time.time() - started total = succ_cnt + fail_cnt rate = succ_cnt / max(total, 1) logger.info( "完成: 成功 {}/{} 成功率 {:.1%} 用时 {:.1f}s 情绪={} 重试分布={}", succ_cnt, total, rate, elapsed, dict(sentiment_cnt), dict(layer_attempts), ) # 验收门槛: ≥ 95% JSON 成功率 return 0 if rate >= 0.95 or total == 0 else 1 def main() -> int: parser = argparse.ArgumentParser(description="A 股 LLM 投资事件抽取 (M4)") parser.add_argument( "--input-root", default="data/deduped", help="输入根目录(默认 dedup 输出);配合 --no-deduped 时改为 data/processed", ) parser.add_argument("--no-deduped", action="store_true", help="跳过 M3,直接读 M2 处理产物") parser.add_argument("--source", default=None, help="--no-deduped 时按源过滤") parser.add_argument("--out-root", default="data/events") parser.add_argument( "--date", default=date.today().strftime("%Y%m%d"), help="日期 YYYYMMDD,默认今日", ) parser.add_argument("--provider", default=None, help="LLM provider: deepseek / qwen,默认读 .env") parser.add_argument("--model", default=None, help="模型名,默认 provider 默认值") parser.add_argument("--concurrency", type=int, default=3, help="LLM 异步并发上限") parser.add_argument("--max-attempts", type=int, default=3, help="单篇文章最大重试次数") parser.add_argument("--force", action="store_true", help="强制全量重抽(默认跳过已抽取文章)") parser.add_argument("--limit", type=int, default=0, help="最多处理 N 篇,0=不限制(用于联调)") parser.add_argument("--prompt", default="prompts/event_extraction.md", help="Prompt 模板路径") parser.add_argument("--log-level", default="INFO") args = parser.parse_args() _setup_logger(args.log_level) return asyncio.run(_run(args)) if __name__ == "__main__": raise SystemExit(main())