feat: add avatar chat and knowledge workflow
This commit is contained in:
136
digital-avatar-app/backend/embeddings.py
Normal file
136
digital-avatar-app/backend/embeddings.py
Normal file
@@ -0,0 +1,136 @@
|
||||
"""
|
||||
向量化服务:对文档/查询文本生成向量。
|
||||
|
||||
优先级:
|
||||
1. 若配置了环境变量 EMBEDDING_API_URL,则调用第三方「OpenAI 兼容」的 /embeddings 接口
|
||||
(需配置 EMBEDDING_API_KEY、EMBEDDING_MODEL,默认 text-embedding-3-small)。
|
||||
2. 否则使用本地「哈希 TF 嵌入」兜底,使向量检索在无外部依赖时也能端到端跑通,
|
||||
且相似文本(共享词汇)会得到更高余弦相似度,便于演示召回效果。
|
||||
"""
|
||||
import os
|
||||
import re
|
||||
import math
|
||||
import json
|
||||
import hashlib
|
||||
import urllib.request
|
||||
|
||||
EMBED_DIM = 256
|
||||
MODEL = os.getenv("EMBEDDING_MODEL", "mock-hash-embed-v1")
|
||||
|
||||
|
||||
def _tokenize(text):
|
||||
text = (text or "").lower()
|
||||
# 英文/数字按词,CJK 逐字(中文无空格,需拆到字级才能命中子词)
|
||||
tokens = re.findall(r"[a-z0-9]+", text)
|
||||
tokens += re.findall(r"[一-鿿]", text)
|
||||
return tokens
|
||||
|
||||
|
||||
def _hash_embedding(texts, dim=EMBED_DIM):
|
||||
vecs = []
|
||||
for text in texts:
|
||||
vec = [0.0] * dim
|
||||
tokens = _tokenize(text)
|
||||
if not tokens:
|
||||
tokens = list(text or "")
|
||||
for tok in tokens:
|
||||
h = int(hashlib.md5(tok.encode("utf-8")).hexdigest(), 16)
|
||||
vec[h % dim] += 1.0
|
||||
norm = math.sqrt(sum(v * v for v in vec))
|
||||
if norm > 0:
|
||||
vec = [v / norm for v in vec]
|
||||
vecs.append(vec)
|
||||
return vecs
|
||||
|
||||
|
||||
def embed(texts):
|
||||
"""返回 list[list[float]],与输入顺序一致。"""
|
||||
if not texts:
|
||||
return []
|
||||
api_url = os.getenv("EMBEDDING_API_URL")
|
||||
if api_url:
|
||||
api_key = os.getenv("EMBEDDING_API_KEY", "")
|
||||
model = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small")
|
||||
payload = json.dumps({"input": texts, "model": model}).encode("utf-8")
|
||||
req = urllib.request.Request(
|
||||
api_url,
|
||||
data=payload,
|
||||
headers={
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {api_key}" if api_key else "",
|
||||
},
|
||||
method="POST",
|
||||
)
|
||||
with urllib.request.urlopen(req, timeout=30) as resp:
|
||||
data = json.loads(resp.read().decode("utf-8"))
|
||||
items = data["data"]
|
||||
if items and "index" in items[0]:
|
||||
items = sorted(items, key=lambda x: x["index"])
|
||||
return [item["embedding"] for item in items]
|
||||
return _hash_embedding(texts)
|
||||
|
||||
|
||||
def cosine(a, b):
|
||||
dot = sum(x * y for x, y in zip(a, b))
|
||||
na = math.sqrt(sum(x * x for x in a))
|
||||
nb = math.sqrt(sum(y * y for y in b))
|
||||
if na == 0 or nb == 0:
|
||||
return 0.0
|
||||
return dot / (na * nb)
|
||||
|
||||
|
||||
def chunk_text(text, size=400, overlap=50):
|
||||
text = (text or "").strip()
|
||||
if not text:
|
||||
return []
|
||||
if len(text) <= size:
|
||||
return [text]
|
||||
chunks = []
|
||||
start = 0
|
||||
while start < len(text):
|
||||
end = min(start + size, len(text))
|
||||
chunks.append(text[start:end])
|
||||
if end == len(text):
|
||||
break
|
||||
start = end - overlap
|
||||
return chunks
|
||||
|
||||
|
||||
def extract_text(path, ext):
|
||||
"""抽取文档纯文本;未知格式拒绝,已知格式解析失败时保留占位文本。"""
|
||||
if ext not in {".txt", ".md", ".docx", ".xlsx", ".pdf", ".doc"}:
|
||||
raise ValueError(f"unsupported file extension: {ext}")
|
||||
try:
|
||||
if ext in {".txt", ".md"}:
|
||||
with open(path, "r", encoding="utf-8", errors="replace") as f:
|
||||
return f.read()
|
||||
if ext == ".docx":
|
||||
from docx import Document
|
||||
|
||||
doc = Document(path)
|
||||
return "\n".join(p.text for p in doc.paragraphs)
|
||||
if ext == ".xlsx":
|
||||
import openpyxl
|
||||
|
||||
wb = openpyxl.load_workbook(path, data_only=True, read_only=True)
|
||||
rows = []
|
||||
for ws in wb.worksheets:
|
||||
for row in ws.iter_rows(values_only=True):
|
||||
cells = [str(c) for c in row if c is not None]
|
||||
if cells:
|
||||
rows.append(" ".join(cells))
|
||||
return "\n".join(rows)
|
||||
if ext == ".pdf":
|
||||
try:
|
||||
from pypdf import PdfReader
|
||||
except ImportError:
|
||||
from PyPDF2 import PdfReader
|
||||
reader = PdfReader(path)
|
||||
return "\n".join((p.extract_text() or "") for p in reader.pages)
|
||||
if ext == ".doc":
|
||||
with open(path, "rb") as f:
|
||||
raw = f.read().decode("utf-8", errors="ignore")
|
||||
return re.sub(r"[\x00-\x08\x0b\x0c\x0e-\x1f]+", " ", raw)
|
||||
except Exception as e: # 解析失败时回退
|
||||
print("extract_text failed:", e)
|
||||
return f"文档:{os.path.basename(path)} 类型 {ext}"
|
||||
Reference in New Issue
Block a user