feat(rag): 新增 Embedder 抽象与离线 FakeEmbedder
This commit is contained in:
@@ -0,0 +1,51 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import List, Protocol
|
||||||
|
|
||||||
|
|
||||||
|
class Embedder(Protocol):
|
||||||
|
def embed(self, texts: List[str]) -> List[List[float]]: ...
|
||||||
|
|
||||||
|
|
||||||
|
_DIM = 64
|
||||||
|
|
||||||
|
|
||||||
|
def _tokenize(text: str) -> List[str]:
|
||||||
|
toks = []
|
||||||
|
cur = ""
|
||||||
|
for ch in text.lower():
|
||||||
|
if ch.isalnum():
|
||||||
|
cur += ch
|
||||||
|
else:
|
||||||
|
if cur:
|
||||||
|
toks.append(cur)
|
||||||
|
cur = ""
|
||||||
|
if cur:
|
||||||
|
toks.append(cur)
|
||||||
|
out = []
|
||||||
|
for t in toks:
|
||||||
|
idx = 0
|
||||||
|
for i, c in enumerate(t):
|
||||||
|
if i > 0 and c.isupper():
|
||||||
|
out.append(t[idx:i])
|
||||||
|
idx = i
|
||||||
|
out.append(t[idx:])
|
||||||
|
return [x for x in out if x]
|
||||||
|
|
||||||
|
|
||||||
|
class FakeEmbedder:
|
||||||
|
def embed(self, texts: List[str]) -> List[List[float]]:
|
||||||
|
vecs = []
|
||||||
|
for t in texts:
|
||||||
|
v = [0.0] * _DIM
|
||||||
|
for tok in _tokenize(t):
|
||||||
|
h = __import__("hashlib").md5(tok.encode("utf-8")).digest()
|
||||||
|
idx = h[0] % _DIM
|
||||||
|
v[idx] += 1.0
|
||||||
|
norm = __import__("math").sqrt(sum(x * x for x in v)) or 1.0
|
||||||
|
vecs.append([x / norm for x in v])
|
||||||
|
return vecs
|
||||||
|
|
||||||
|
|
||||||
|
def get_embedder(engine) -> Embedder:
|
||||||
|
return FakeEmbedder()
|
||||||
@@ -0,0 +1,30 @@
|
|||||||
|
# tests/test_rag_embeddings.py
|
||||||
|
from genesis.rag.embeddings import FakeEmbedder, get_embedder
|
||||||
|
|
||||||
|
|
||||||
|
def test_fake_embedder_deterministic():
|
||||||
|
e = FakeEmbedder()
|
||||||
|
a = e.embed(["OrderController 创建订单"])[0]
|
||||||
|
b = e.embed(["OrderController 创建订单"])[0]
|
||||||
|
assert a == b
|
||||||
|
|
||||||
|
|
||||||
|
def test_fake_embedder_similar_closer_than_unrelated():
|
||||||
|
e = FakeEmbedder()
|
||||||
|
base = e.embed(["OrderController 处理创建订单请求"])[0]
|
||||||
|
sim = e.embed(["OrderController 保存订单到数据库"])[0]
|
||||||
|
dif = e.embed(["用户登录认证模块"])[0]
|
||||||
|
import math
|
||||||
|
|
||||||
|
def cos(x, y):
|
||||||
|
dot = sum(p * q for p, q in zip(x, y))
|
||||||
|
nx = math.sqrt(sum(p * p for p in x))
|
||||||
|
ny = math.sqrt(sum(q * q for q in y))
|
||||||
|
return dot / (nx * ny or 1.0)
|
||||||
|
|
||||||
|
assert cos(base, sim) > cos(base, dif)
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_embedder_fake_engine_returns_fake():
|
||||||
|
e = get_embedder("fake")
|
||||||
|
assert e.__class__.__name__ == "FakeEmbedder"
|
||||||
Reference in New Issue
Block a user