"""M5 批量嵌入入口脚本。 输入策略: A. 默认 events 优先 (--input events): data/events/{day}/*.json (M4 ExtractedEvent),用 LLM 摘要 + 事件标签 + 正文组装文本; 同时回查对应 article(M2 输出)以拿正文,如查不到则用 ExtractedEvent.event.summary 作为正文。 B. articles 模式 (--input articles): data/processed/{source}/{day}/*.json (M2 Article 直出),仅用 title + content。 C. deduped 模式 (--input deduped): data/deduped/{day}/uniques/*.json (M3 唯一文章),仅 title + content。 输出: data/embeddings/{day}/{url_hash}.json (含 vector 完整内容) data/embeddings/{day}/index.jsonl (扁平摘要,不含向量,便于检索/调试) data/embeddings/{day}/failed.jsonl (失败列表) """ from __future__ import annotations import argparse import asyncio import json import os import sys import time # 修复树莓派等环境的 ASCII locale 问题(UnicodeEncodeError) os.environ.setdefault("PYTHONUTF8", "1") from datetime import date, datetime from pathlib import Path from typing import Any from dotenv import load_dotenv from loguru import logger from pydantic import ValidationError from embedding import ( EmbeddingError, EmbeddingResult, compose_text, make_async_provider, ) from embedding.base import _from_event_dict from extractor import Article 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") / "embedding.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 _list_source_dirs(p: Path) -> list[str]: if not p.is_dir(): return [] return sorted(d.name for d in p.iterdir() if d.is_dir()) def _load_article_by_hash(processed_root: Path, day: str, url_hash: str) -> Article | None: """根据 url_hash 在 data/processed/*/{day}/ 中查 Article。""" for src_dir in processed_root.iterdir(): if not src_dir.is_dir(): continue candidate = src_dir / day / f"{url_hash}.json" if candidate.is_file(): try: return Article.model_validate( json.loads(candidate.read_text(encoding="utf-8")) ) except (json.JSONDecodeError, ValidationError): return None return None def _build_text_from_event( event_path: Path, processed_root: Path, day: str, ) -> tuple[str, Article, str | None] | None: """从 ExtractedEvent JSON 构造嵌入文本与文章元数据。""" try: obj: dict[str, Any] = json.loads(event_path.read_text(encoding="utf-8")) except json.JSONDecodeError as e: logger.warning("跳过损坏 event 文件 {}: {}", event_path, e) return None article, head, summary = _from_event_dict(obj) # 真正的正文要去 processed/ 找 real_article = _load_article_by_hash(processed_root, day, article.url_hash) if real_article is None: # 退化:用 summary 作为正文(质量降级,但仍可嵌入) real_article = article logger.debug("未找到 Article(url_hash={}),用 summary 作为正文兜底", article.url_hash) else: # 用 processed 的真实数据,但保留 event 元数据(time/source 等) real_article = real_article.model_copy( update={"publish_time": article.publish_time or real_article.publish_time} ) text = compose_text(real_article, head=head, summary=summary) return text, real_article, summary def _build_text_from_article(article_path: Path) -> tuple[str, Article, str | None] | None: try: obj = json.loads(article_path.read_text(encoding="utf-8")) article = Article.model_validate(obj) except (json.JSONDecodeError, ValidationError) as e: logger.warning("跳过损坏 article 文件 {}: {}", article_path, e) return None return compose_text(article), article, None def _collect_inputs(args: argparse.Namespace) -> list[tuple[Path, str]]: """返回 [(input_file, kind), ...],kind 为 'event'/'article'。""" day = args.date files: list[tuple[Path, str]] = [] if args.input == "events": d = Path(args.events_root) / day files = [(p, "event") for p in sorted(d.glob("*.json"))] elif args.input == "deduped": d = Path(args.deduped_root) / day / "uniques" files = [(p, "article") for p in sorted(d.glob("*.json"))] elif args.input == "articles": proc = Path(args.processed_root) srcs = [args.source] if args.source else _list_source_dirs(proc) for src in srcs: sd = proc / src / day if sd.is_dir(): files.extend((p, "article") for p in sorted(sd.glob("*.json"))) else: raise ValueError(f"未知 --input 模式: {args.input}") return files # --------------------------------------------------------------------------- # # 主流程 # --------------------------------------------------------------------------- # async def _run(args: argparse.Namespace) -> int: load_dotenv() files = _collect_inputs(args) if args.limit: files = files[: args.limit] if not files: logger.error("未发现任何输入文件: {} ({})", args.input, args.date) return 2 logger.info("待嵌入文章数: {} (input={})", len(files), args.input) # 准备每篇文本 prepared: list[tuple[str, Article, str | None]] = [] for fp, kind in files: if kind == "event": built = _build_text_from_event(fp, Path(args.processed_root), args.date) else: built = _build_text_from_article(fp) if built is not None: prepared.append(built) if not prepared: logger.error("所有输入文件均无法解析") return 2 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" for p in (index_path, failed_path): if p.exists(): p.unlink() started = time.time() succ_cnt = 0 fail_cnt = 0 async with make_async_provider( args.provider, **({"model": args.model} if args.model else {}), ) as provider: logger.info( "Embedding provider={} model={} dim={}", provider.name, provider.model, provider.dim, ) # 异步分批 batch_size = args.batch_size for i in range(0, len(prepared), batch_size): batch = prepared[i : i + batch_size] texts = [t for t, _, _ in batch] try: vectors = await provider.embed_batch(texts) except EmbeddingError as e: logger.warning("批 {} 嵌入失败: {}", i // batch_size, e) for _, art, _ in batch: fail_cnt += 1 with failed_path.open("a", encoding="utf-8") as f: f.write( json.dumps( {"url_hash": art.url_hash, "url": art.url, "error": str(e)}, ensure_ascii=False, ) + "\n" ) continue for (text, article, _summary), vec in zip(batch, vectors, strict=True): if len(vec) != provider.dim: logger.warning( "维度不一致 url_hash={} 实际={} 预期={}", article.url_hash, len(vec), provider.dim, ) result = EmbeddingResult( url_hash=article.url_hash, source_id=article.source_id, title=article.title, text=text, vector=vec, dim=len(vec), provider=provider.name, model=provider.model, embedded_at=datetime.now(), char_count=len(text), publish_time=article.publish_time, ) # 落盘:完整 + index (out_dir / f"{result.url_hash}.json").write_text( result.model_dump_json(indent=2), encoding="utf-8" ) meta = result.model_dump(exclude={"vector", "text"}, mode="json") meta["text_preview"] = text[:60] with index_path.open("a", encoding="utf-8") as f: f.write(json.dumps(meta, ensure_ascii=False) + "\n") succ_cnt += 1 logger.info( "已处理批 {}: 成功 +{}/{} (累计成功 {})", i // batch_size + 1, len(batch), len(batch), succ_cnt, ) elapsed = time.time() - started total = succ_cnt + fail_cnt rate = succ_cnt / max(total, 1) logger.info( "Embedding 完成: 成功 {}/{} 成功率 {:.1%} 用时 {:.1f}s", succ_cnt, total, rate, elapsed, ) return 0 if rate >= 0.95 or total == 0 else 1 def main() -> int: parser = argparse.ArgumentParser(description="A 股新闻 Embedding 向量化 (M5)") parser.add_argument( "--input", default="events", choices=["events", "deduped", "articles"], help="输入来源:events (M4) / deduped (M3) / articles (M2)", ) parser.add_argument("--events-root", default="data/events") parser.add_argument("--deduped-root", default="data/deduped") parser.add_argument("--processed-root", default="data/processed") parser.add_argument("--source", default=None, help="--input articles 时按源过滤") parser.add_argument("--out-root", default="data/embeddings") parser.add_argument( "--date", default=date.today().strftime("%Y%m%d"), help="日期 YYYYMMDD,默认今日", ) parser.add_argument("--provider", default=None, help="dashscope (默认) / local-bge,可被 .env 覆盖") parser.add_argument("--model", default=None, help="嵌入模型名,覆盖 .env 中的 *_EMBEDDING_MODEL") parser.add_argument("--batch-size", type=int, default=10, help="每批送 embed 的条数(DashScope 上限 10)") parser.add_argument("--limit", type=int, default=0, help="最多处理 N 篇,0=不限") 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())