embed_only.py 3.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106
  1. #!/usr/bin/env python3
  2. """仅向量化:读取PG中无向量的chunk,批量embedding后回写"""
  3. import json
  4. import asyncio
  5. import os
  6. import sys
  7. from pathlib import Path
  8. import httpx
  9. from sqlalchemy.ext.asyncio import create_async_engine
  10. from sqlalchemy import text
  11. ROOT = Path(__file__).resolve().parent.parent
  12. env_file = ROOT / ".env"
  13. if env_file.exists():
  14. for line in open(env_file):
  15. line = line.strip()
  16. if line and not line.startswith("#") and "=" in line:
  17. key, _, val = line.partition("=")
  18. os.environ.setdefault(key.strip(), val.strip())
  19. API_KEY = os.environ.get("QWEN_API_KEY", "")
  20. EMBEDDING_URL = "https://dashscope.aliyuncs.com/api/v1/services/embeddings/text-embedding/text-embedding"
  21. DB_PW = os.environ.get("POSTGRES_PASSWORD", "pharma2025")
  22. DB_USER = os.environ.get("POSTGRES_USER", "postgres")
  23. DB_HOST = os.environ.get("POSTGRES_HOST", "localhost")
  24. DB_PORT = os.environ.get("POSTGRES_PORT", "5432")
  25. DB_NAME = os.environ.get("POSTGRES_DB", "pharmacopoeia")
  26. DB_URL = f"postgresql+asyncpg://{DB_USER}:{DB_PW}@{DB_HOST}:{DB_PORT}/{DB_NAME}"
  27. async def embed_batch(texts: list[str]) -> list[list[float]]:
  28. async with httpx.AsyncClient(timeout=120) as client:
  29. resp = await client.post(
  30. EMBEDDING_URL,
  31. headers={
  32. "Content-Type": "application/json",
  33. "Authorization": f"Bearer {API_KEY}",
  34. },
  35. json={
  36. "model": "text-embedding-v3",
  37. "input": {"texts": texts},
  38. "parameters": {"text_type": "document"},
  39. },
  40. )
  41. data = resp.json()
  42. if data.get("code") and data["code"] != "":
  43. raise RuntimeError(f"Embedding API error: {data.get('message', data)}")
  44. return [item["embedding"] for item in data["output"]["embeddings"]]
  45. async def main():
  46. print("=" * 50)
  47. print("Embedding-only: 仅向量化 + 写回 PG")
  48. print("=" * 50)
  49. if not API_KEY:
  50. print("❌ QWEN_API_KEY not set")
  51. sys.exit(1)
  52. engine = create_async_engine(DB_URL)
  53. async with engine.begin() as conn:
  54. rows = (await conn.execute(text(
  55. "SELECT id, content FROM drug_chunks WHERE embedding IS NULL OR embedding::text = '[]' ORDER BY id"
  56. ))).fetchall()
  57. total = len(rows)
  58. print(f"\n📊 待处理 chunk: {total}")
  59. if total == 0:
  60. print("✅ 所有 chunk 已有向量,无需处理")
  61. await engine.dispose()
  62. return
  63. ids = [r[0] for r in rows]
  64. contents = [r[1] for r in rows]
  65. BATCH = 10
  66. count = 0
  67. async with engine.begin() as conn:
  68. for i in range(0, total, BATCH):
  69. batch_texts = contents[i : i + BATCH]
  70. batch_ids = ids[i : i + BATCH]
  71. vecs = await embed_batch(batch_texts)
  72. for cid, vec in zip(batch_ids, vecs):
  73. vec_str = json.dumps(vec)
  74. await conn.execute(
  75. text("UPDATE drug_chunks SET embedding = CAST(:emb AS jsonb) WHERE id = :id"),
  76. {"id": cid, "emb": vec_str},
  77. )
  78. count += len(batch_texts)
  79. print(f" 进度: {count}/{total}")
  80. await engine.dispose()
  81. print(f"\n🎉 向量化完成!{count} 个 chunk 已更新")
  82. if __name__ == "__main__":
  83. asyncio.run(main())