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