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
| """测试数据库层.""" | |
| import os | |
| import sqlite3 | |
| import tempfile | |
| import pytest | |
| from hermes_core.types import Record, Scope, TrainingRun, RecordState, TrainingStatus | |
| from hermes_core.db import ( | |
| init_db, | |
| insert_record, | |
| get_records_by_scope, | |
| get_active_records, | |
| get_record, | |
| update_record_state, | |
| update_record_label, | |
| upsert_scope, | |
| get_scope, | |
| get_active_scopes, | |
| get_scopes_needing_training, | |
| insert_training_run, | |
| update_training_run, | |
| get_latest_checkpoint, | |
| get_max_dims, | |
| set_max_dims, | |
| ) | |
| 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 | |
| class TestRecords: | |
| def test_insert_and_get_record(self, conn): | |
| rec = Record( | |
| id="rec_001", user_id="u1", scope_id="scope_a", | |
| scope_label="后端开发", dimensions=[], confidence=0.8, | |
| ) | |
| insert_record(conn, rec) | |
| got = get_record(conn, "rec_001") | |
| assert got is not None | |
| assert got.id == "rec_001" | |
| assert got.confidence == 0.8 | |
| def test_insert_upserts_on_duplicate(self, conn): | |
| rec1 = Record(id="rec_001", user_id="u1", scope_id="scope_a", | |
| scope_label="后端", dimensions=[], confidence=0.5) | |
| insert_record(conn, rec1) | |
| rec2 = Record(id="rec_001", user_id="u1", scope_id="scope_a", | |
| scope_label="后端", dimensions=[], confidence=0.9, occurrences=2) | |
| insert_record(conn, rec2) | |
| got = get_record(conn, "rec_001") | |
| assert got.confidence == 0.9 | |
| assert got.occurrences == 2 | |
| def test_get_records_by_scope(self, conn): | |
| for i in range(3): | |
| rec = Record(id=f"rec_{i}", user_id="u1", scope_id="scope_a", | |
| scope_label="后端", dimensions=[]) | |
| insert_record(conn, rec) | |
| records = get_records_by_scope(conn, "scope_a") | |
| assert len(records) == 3 | |
| def test_get_active_records_filters_rejected(self, conn): | |
| rec1 = Record(id="rec_1", user_id="u1", scope_id="scope_a", | |
| scope_label="后端", dimensions=[], state=RecordState.active) | |
| rec2 = Record(id="rec_2", user_id="u1", scope_id="scope_a", | |
| scope_label="后端", dimensions=[], state=RecordState.rejected) | |
| insert_record(conn, rec1) | |
| insert_record(conn, rec2) | |
| active = get_active_records(conn, "scope_a") | |
| assert len(active) == 1 | |
| assert active[0].id == "rec_1" | |
| def test_update_record_state(self, conn): | |
| rec = Record(id="rec_001", user_id="u1", scope_id="scope_a", | |
| scope_label="后端", dimensions=[], state=RecordState.active) | |
| insert_record(conn, rec) | |
| update_record_state(conn, "rec_001", RecordState.rejected) | |
| got = get_record(conn, "rec_001") | |
| assert got.state == RecordState.rejected | |
| def test_update_record_label(self, conn): | |
| rec = Record(id="rec_001", user_id="u1", scope_id="scope_a", | |
| scope_label="后端", dimensions=[]) | |
| insert_record(conn, rec) | |
| update_record_label(conn, "rec_001", "scope_b", "系统编程") | |
| got = get_record(conn, "rec_001") | |
| assert got.scope_id == "scope_b" | |
| assert got.scope_label == "系统编程" | |
| class TestScopes: | |
| def test_upsert_and_get_scope(self, conn): | |
| scope = Scope(id="scope_a", label="后端开发", centroid=[0.1, 0.2]) | |
| upsert_scope(conn, scope) | |
| got = get_scope(conn, "scope_a") | |
| assert got is not None | |
| assert got.label == "后端开发" | |
| def test_upsert_scope_merges(self, conn): | |
| s1 = Scope(id="scope_a", label="后端", centroid=[0.1], record_count=5) | |
| upsert_scope(conn, s1) | |
| s2 = Scope(id="scope_a", label="后端开发", centroid=[0.2], record_count=10, coherence=0.9) | |
| upsert_scope(conn, s2) | |
| got = get_scope(conn, "scope_a") | |
| assert got.label == "后端开发" | |
| assert got.record_count == 10 | |
| assert got.coherence == 0.9 | |
| def test_get_active_scopes(self, conn): | |
| s1 = Scope(id="scope_a", label="活跃", status="active") | |
| s2 = Scope(id="scope_b", label="归档", status="archived") | |
| upsert_scope(conn, s1) | |
| upsert_scope(conn, s2) | |
| scopes = get_active_scopes(conn) | |
| assert len(scopes) == 1 | |
| assert scopes[0].id == "scope_a" | |
| def test_get_scopes_needing_training(self, conn): | |
| s1 = Scope(id="scope_a", label="a", needs_training=True) | |
| s2 = Scope(id="scope_b", label="b", needs_training=False) | |
| s3 = Scope(id="scope_c", label="c", needs_training=True, status="archived") | |
| upsert_scope(conn, s1) | |
| upsert_scope(conn, s2) | |
| upsert_scope(conn, s3) | |
| needing = get_scopes_needing_training(conn) | |
| assert len(needing) == 1 | |
| assert needing[0].id == "scope_a" | |
| class TestTrainingRuns: | |
| def test_insert_and_get_latest_checkpoint(self, conn): | |
| run = TrainingRun(id="run_1", scope_id="scope_a", version=1, | |
| status=TrainingStatus.done, checkpoint_path="/tmp/v1") | |
| insert_training_run(conn, run) | |
| latest = get_latest_checkpoint(conn, "scope_a") | |
| assert latest is not None | |
| assert latest.version == 1 | |
| def test_update_training_run(self, conn): | |
| run = TrainingRun(id="run_1", scope_id="scope_a", version=1) | |
| insert_training_run(conn, run) | |
| update_training_run(conn, "run_1", | |
| status=TrainingStatus.training, | |
| progress={"epoch": 1, "loss": 0.5}) | |
| # verify by re-reading | |
| import json | |
| cur = conn.execute("SELECT status, progress FROM training_runs WHERE id='run_1'") | |
| row = cur.fetchone() | |
| assert row[0] == "training" | |
| assert json.loads(row[1]) == {"epoch": 1, "loss": 0.5} | |
| class TestDimensionConstraints: | |
| def test_get_set_max_dims(self, conn): | |
| max_dims = get_max_dims(conn, "scope_a") | |
| assert max_dims is None # 无记录时返回 None | |
| set_max_dims(conn, "scope_a", "u1", 3) | |
| assert get_max_dims(conn, "scope_a") == 3 | |
| set_max_dims(conn, "scope_a", "u1", 5) | |
| assert get_max_dims(conn, "scope_a") == 5 | |