Files
intl_news/embedding/pipeline.py
T
2026-07-18 16:13:52 +08:00

180 lines
5.3 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""批量向量生成管道。
输入: data/events/{YYYYMMDD}/{url_hash}.jsonM4 翻译+事件输出)
输出: data/embeddings/{YYYYMMDD}/{url_hash}.json
"""
import json
import logging
from datetime import datetime
from pathlib import Path
from crawler.utils import get_news_day
from embedding.client import (
EmbeddingConfig,
load_embedding_config,
make_embedding_client,
)
from embedding.embedder import embed_articles
from embedding.models import EmbeddingError
from llm.models import EnTranslatedArticle
logger = logging.getLogger(__name__)
def _load_event_articles(date_str: str) -> list[EnTranslatedArticle]:
"""加载指定日期的翻译+事件文章。
Args:
date_str: 日期 YYYYMMDD
Returns:
EnTranslatedArticle 列表
"""
base_dir = Path(f"data/events/{date_str}")
if not base_dir.exists():
return []
articles: list[EnTranslatedArticle] = []
for json_file in sorted(base_dir.glob("*.json")):
if json_file.name == "index.json":
continue
try:
data = json.loads(json_file.read_text(encoding="utf-8"))
articles.append(EnTranslatedArticle(**data))
except (json.JSONDecodeError, Exception) as e:
logger.warning("解析事件文章失败 %s: %s", json_file, e)
return articles
def embed_all_events(
date_str: str | None = None,
*,
model: str | None = None,
) -> dict:
"""对所有 M4 输出的文章执行向量化。
Args:
date_str: 日期 YYYYMMDD,默认当前新闻日
model: Embedding 模型名,默认从 system.yaml 读取
Returns:
统计 dict
"""
if date_str is None:
date_str = get_news_day()
logger.info("══════ 开始向量生成,日期: %s ══════", date_str)
# 加载文章
articles = _load_event_articles(date_str)
if not articles:
logger.warning("事件目录无文章: data/events/%s/", date_str)
return {"date": date_str, "total": 0, "success": 0, "failed": 0, "elapsed_sec": 0}
# 初始化 Embedding 客户端
config = load_embedding_config(model=model)
client = make_embedding_client(config)
# 输出目录
out_dir = Path(f"data/embeddings/{date_str}")
out_dir.mkdir(parents=True, exist_ok=True)
# 增量:跳过已向量化的文章
new_articles = []
skipped = 0
for a in articles:
if (out_dir / f"{a.url_hash}.json").exists():
skipped += 1
else:
new_articles.append(a)
if skipped > 0:
logger.info("增量跳过 %d 篇已向量化,剩余 %d 篇待处理", skipped, len(new_articles))
articles = new_articles
start_time = datetime.now()
success = 0
failed = 0
# 批量嵌入(按 batch_size 分块,每批输出进度)
batch_size = config.batch_size
total = len(articles)
logger.info("开始向量化 %d 篇文章(batch_size=%d, model=%s",
total, batch_size, config.model)
for batch_start in range(0, total, batch_size):
batch_end = min(batch_start + batch_size, total)
batch = articles[batch_start:batch_end]
try:
results = embed_articles(client, config, batch)
for result in results:
out_file = out_dir / f"{result.url_hash}.json"
out_file.write_text(
result.model_dump_json(indent=2, ensure_ascii=False),
encoding="utf-8",
)
success += 1
logger.info(" [%d/%d] ✅ %d 篇 → %d 维向量",
batch_end, total, len(results), config.dimension)
except EmbeddingError as e:
failed += len(batch)
logger.error("批量嵌入失败 [%d-%d]: %s", batch_start, batch_end, e.reason)
except Exception as e:
failed += len(batch)
logger.exception("批量嵌入异常 [%d-%d]: %s", batch_start, batch_end, e)
elapsed = (datetime.now() - start_time).total_seconds()
# 写入索引
_write_embedding_index(date_str, success, failed, elapsed, config)
logger.info(
"══════ 向量生成完成: 成功 %d / 失败 %d / 总计 %d,耗时 %.1f 秒 ══════",
success, failed, len(articles), elapsed,
)
return {
"date": date_str,
"total": len(articles),
"success": success,
"failed": failed,
"elapsed_sec": elapsed,
"provider": config.provider,
"model": config.model,
"dimension": config.dimension,
}
def _write_embedding_index(
date_str: str,
success: int,
failed: int,
elapsed_sec: float,
config: EmbeddingConfig,
) -> None:
"""写入向量索引文件。"""
out_dir = Path(f"data/embeddings/{date_str}")
out_dir.mkdir(parents=True, exist_ok=True)
index_data = {
"date": date_str,
"success": success,
"failed": failed,
"elapsed_sec": round(elapsed_sec, 1),
"provider": config.provider,
"model": config.model,
"dimension": config.dimension,
"generated_at": datetime.now().isoformat(),
}
index_path = out_dir / "index.json"
index_path.write_text(
json.dumps(index_data, indent=2, ensure_ascii=False),
encoding="utf-8",
)
logger.info("向量索引已写入: %s", index_path)