Files
news/scripts/run_event_extraction.py
T
2026-07-18 15:51:01 +08:00

221 lines
7.6 KiB
Python

"""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 (失败列表)
用法:
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
"""
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 (
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 _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
async def _run(args: argparse.Namespace) -> int:
load_dotenv() # 读 .env 到 os.environ
config = load_llm_config(provider=args.provider, model=args.model)
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 args.limit:
files = files[: args.limit]
if not files:
logger.error(
"{} 下未发现 {} 的文章(use_deduped={})",
input_root, args.date, use_deduped,
)
return 2
logger.info("待处理文章数: {}", len(files))
out_dir = Path(args.out_root) / args.date
out_dir.mkdir(parents=True, exist_ok=True)
index_path = out_dir / "index.jsonl"
failed_path = out_dir / "failed.jsonl"
# 重跑时清掉旧的 jsonl,避免重复追加
for p in (index_path, failed_path):
if p.exists():
p.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, "无法解析输入"
try:
event = await extract_event_async(
client=client,
config=config,
article=article,
template=template,
max_attempts=args.max_attempts,
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("--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())