Text Generation
PEFT
Chinese
English
preference-learning
qlora
agent
personalization
association-engine
Instructions to use feiertu/hermes-association-engine with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use feiertu/hermes-association-engine with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
File size: 4,442 Bytes
00138da | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 | """测试动态语义聚类."""
import tempfile
import sqlite3
import os
import pytest
from hermes_core.db import init_db, upsert_scope, insert_record, get_active_scopes, get_scope
from hermes_core.embedder import Embedder
from hermes_core.types import Record, Scope
from hermes_core.cluster import (
assign_scope,
recluster_scope,
check_merge,
check_split,
DEFAULT_MATCH_THRESHOLD,
DEFAULT_MERGE_THRESHOLD,
DEFAULT_SPLIT_THRESHOLD,
)
pytestmark = pytest.mark.requires_embedder
@pytest.fixture
def conn():
tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False)
conn = sqlite3.connect(tmp.name)
init_db(conn)
yield conn
conn.close()
try:
os.unlink(tmp.name)
except PermissionError:
pass # Windows may hold file lock briefly
@pytest.fixture
def embedder():
return Embedder()
class TestAssignScope:
def test_creates_new_scope_when_none_exist(self, conn, embedder):
sid, label = assign_scope(conn, embedder, "后端API开发")
assert sid.startswith("scope_")
assert label == "后端API开发"
scopes = get_active_scopes(conn)
assert len(scopes) == 1
def test_matches_existing_scope_on_high_similarity(self, conn, embedder):
# 先创建一个 scope
s = Scope(id="scope_existing", label="后端服务开发",
centroid=embedder.encode("后端服务开发"), record_count=5)
upsert_scope(conn, s)
# 赋一条语义相近的
sid, label = assign_scope(conn, embedder, "后端API开发")
# 应该匹配到已有 scope
assert sid == "scope_existing"
def test_creates_new_scope_on_low_similarity(self, conn, embedder):
s = Scope(id="scope_existing", label="后端开发",
centroid=embedder.encode("后端开发"), record_count=5)
upsert_scope(conn, s)
sid, label = assign_scope(conn, embedder, "周末去哪玩")
assert sid != "scope_existing"
assert len(get_active_scopes(conn)) == 2
class TestRecluster:
def test_recluster_updates_centroid(self, conn, embedder):
s = Scope(id="scope_a", label="后端",
centroid=embedder.encode("后端开发"), record_count=1)
upsert_scope(conn, s)
insert_record(conn, Record(
id="r1", user_id="u1", scope_id="scope_a",
scope_label="后端", dimensions=[]
))
insert_record(conn, Record(
id="r2", user_id="u1", scope_id="scope_a",
scope_label="后端服务", dimensions=[]
))
recluster_scope(conn, embedder, "scope_a")
updated = get_scope(conn, "scope_a")
# 新的 centroid 应该是两条记录 label 的平均
assert updated.centroid is not None
assert updated.record_count == 2
class TestCheckMerge:
def test_similar_scopes_flagged_for_merge(self, conn, embedder):
s1 = Scope(id="scope_1", label="Python数据分析",
centroid=embedder.encode("Python数据分析"), record_count=3)
s2 = Scope(id="scope_2", label="Python数据处理",
centroid=embedder.encode("Python数据处理"), record_count=3)
upsert_scope(conn, s1)
upsert_scope(conn, s2)
merges = check_merge(conn, embedder, threshold=0.85)
assert len(merges) > 0
def test_dissimilar_scopes_not_merged(self, conn, embedder):
s1 = Scope(id="scope_1", label="后端开发",
centroid=embedder.encode("后端开发"), record_count=3)
s2 = Scope(id="scope_2", label="周末去哪玩",
centroid=embedder.encode("周末去哪玩"), record_count=3)
upsert_scope(conn, s1)
upsert_scope(conn, s2)
merges = check_merge(conn, embedder, threshold=0.85)
assert len(merges) == 0
class TestCheckSplit:
def test_coherent_scope_not_flagged(self, conn, embedder):
label = "后端开发"
s = Scope(id="scope_a", label=label,
centroid=embedder.encode(label), record_count=1)
upsert_scope(conn, s)
for i in range(5):
insert_record(conn, Record(
id=f"r{i}", user_id="u1", scope_id="scope_a",
scope_label=f"后端开发-{i}", dimensions=[]
))
needs_split = check_split(conn, embedder, "scope_a", threshold=0.40)
assert not needs_split
|