Initial commit
This commit is contained in:
@@ -0,0 +1,293 @@
|
||||
"""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())
|
||||
Reference in New Issue
Block a user