143 lines
5.7 KiB
Python
143 lines
5.7 KiB
Python
#!/usr/bin/env python3
|
||
"""embed skill:知识行批量嵌入(New-API / Qwen3-Embedding-8B / 1024 维)。
|
||
|
||
合同见同 skill SKILL.md;通道事实见 db/连接信息.md。失败原样报错不静默。
|
||
"""
|
||
import hashlib
|
||
import json
|
||
import sys
|
||
import time
|
||
|
||
import click
|
||
import psycopg
|
||
import requests
|
||
|
||
DSN = ("postgresql://root:f6710e2d0294eb1c10e26a805a64bc54@100.64.0.8:5433/muse-example"
|
||
"?keepalives=1&keepalives_idle=15&keepalives_interval=5&keepalives_count=3")
|
||
BASE = "http://100.64.0.8:3000"
|
||
TOKEN = "sk-DyVqO3lDmEvQZ3PqGpbNaaaHZHhbh0xaHRIiynhYSmVlLHl2" # MUSE_AI_NEW_API_TOKEN(勿用管理令牌)
|
||
MODEL = "Qwen/Qwen3-Embedding-8B"
|
||
DIM = 1024
|
||
TENANT, ACTOR = 1, "1"
|
||
BATCH = 16
|
||
|
||
|
||
def _session():
|
||
"""禁系统代理的会话(系统代理会假 502)。"""
|
||
s = requests.Session()
|
||
s.trust_env = False
|
||
s.headers["Authorization"] = f"Bearer {TOKEN}"
|
||
return s
|
||
|
||
|
||
def embed_texts(sess, texts):
|
||
"""调 New-API /v1/embeddings;整批重试 2 次后逐条降级。返回 (向量列表, 失败索引集)。"""
|
||
def call(batch):
|
||
r = sess.post(f"{BASE}/v1/embeddings", json={
|
||
"model": MODEL, "input": batch, "dimensions": DIM}, timeout=120)
|
||
r.raise_for_status()
|
||
data = r.json()["data"]
|
||
return [d["embedding"] for d in sorted(data, key=lambda d: d["index"])]
|
||
|
||
for attempt in range(3):
|
||
try:
|
||
return call(texts), set()
|
||
except Exception as e:
|
||
if attempt < 2:
|
||
time.sleep(2 ** attempt)
|
||
continue
|
||
# 整批三败 → 逐条降级,坏行记错不断批
|
||
vecs, bad = [], set()
|
||
for i, t in enumerate(texts):
|
||
try:
|
||
vecs.append(call([t])[0])
|
||
except Exception as ee:
|
||
vecs.append(None)
|
||
bad.add(i)
|
||
click.echo(f" [失败] 第{i}条: {ee}", err=True)
|
||
return vecs, bad
|
||
|
||
|
||
def build_embed_text(payload: dict) -> str:
|
||
"""嵌入文本构造:payload 自带 embed_text 优先;否则固定拼接(与检索端语义对齐)。"""
|
||
if payload.get("embed_text"):
|
||
return payload["embed_text"]
|
||
t = payload.get("型") or payload.get("target_type", "")
|
||
name = payload.get("名称", "")
|
||
brief = payload.get("一句话摘要", "")
|
||
fields = payload.get("字段") or {}
|
||
body = "\n".join(f"{k}:{v}" for k, v in fields.items() if v and k not in ("名称", "一句话摘要"))
|
||
return f"【{t}】{name}:{brief}\n{body}"[:4000]
|
||
|
||
|
||
@click.command()
|
||
@click.option("--work-id", type=int, help="限定拆书批次的 work(draft.work_id=0 为全局行,用 source_id 关联参考书)")
|
||
@click.option("--limit", type=int, default=0, help="最多处理条数(0=不限)")
|
||
@click.option("--probe", help="自由文本试嵌(打印维度与前 5 维,不落库)")
|
||
def main(work_id, limit, probe):
|
||
sess = _session()
|
||
if probe:
|
||
vecs, bad = embed_texts(sess, [probe])
|
||
if bad:
|
||
raise click.ClickException("试嵌失败")
|
||
v = vecs[0]
|
||
click.echo(f"维度={len(v)} 前5维={[round(x, 4) for x in v[:5]]}")
|
||
return
|
||
|
||
with psycopg.connect(DSN) as conn:
|
||
# 待嵌=pending 草稿且无嵌入行
|
||
sql = """SELECT d.id, d.draft_payload FROM muse_knowledge_draft d
|
||
WHERE d.tenant_id=%s AND d.deleted=FALSE AND d.status='pending'
|
||
AND NOT EXISTS (SELECT 1 FROM example_knowledge_embedding e
|
||
WHERE e.tenant_id=%s AND e.draft_id=d.id AND e.deleted=FALSE)"""
|
||
args = [TENANT, TENANT]
|
||
if work_id is not None:
|
||
sql += " AND d.source_id=%s"
|
||
args.append(work_id)
|
||
sql += " ORDER BY d.id"
|
||
if limit:
|
||
sql += f" LIMIT {int(limit)}"
|
||
rows = conn.execute(sql, args).fetchall()
|
||
click.echo(f"待嵌草稿: {len(rows)} 条")
|
||
|
||
done = skip = fail = 0
|
||
for i in range(0, len(rows), BATCH):
|
||
chunk = rows[i:i + BATCH]
|
||
texts, metas = [], []
|
||
for did, payload in chunk:
|
||
text = build_embed_text(payload or {})
|
||
h = hashlib.sha256(f"{text}|{MODEL}".encode()).hexdigest()
|
||
if conn.execute(
|
||
"SELECT 1 FROM example_knowledge_embedding WHERE tenant_id=%s AND content_hash=%s AND model=%s",
|
||
(TENANT, h, MODEL)).fetchone():
|
||
skip += 1 # 幂等:同文同模型不重嵌
|
||
continue
|
||
texts.append(text)
|
||
metas.append((did, h, text))
|
||
if not texts:
|
||
continue
|
||
vecs, bad = embed_texts(sess, texts)
|
||
for j, (did, h, text) in enumerate(metas):
|
||
if j in bad:
|
||
fail += 1
|
||
continue
|
||
conn.execute(
|
||
"""INSERT INTO example_knowledge_embedding
|
||
(draft_id, content_hash, embed_text, model, dimensions, embedding,
|
||
creator, updater, tenant_id)
|
||
VALUES (%s,%s,%s,%s,%s,%s,%s,%s,%s)
|
||
ON CONFLICT (tenant_id, content_hash, model) DO NOTHING""",
|
||
(did, h, text, MODEL, DIM, json.dumps(vecs[j]), ACTOR, ACTOR, TENANT))
|
||
done += 1
|
||
conn.commit()
|
||
click.echo(f" 进度 {min(i + BATCH, len(rows))}/{len(rows)}(新嵌{done} 跳过{skip} 失败{fail})")
|
||
click.echo(f"完成:新嵌 {done}、幂等跳过 {skip}、失败 {fail}")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
try:
|
||
main()
|
||
except (psycopg.Error, requests.RequestException) as e:
|
||
click.echo(f"[错误] {type(e).__name__}: {e}", err=True)
|
||
sys.exit(1)
|