框架: embedding 按行落卡并用 blob COUNT 对账

This commit is contained in:
zizi 2026-08-28 07:55:44 +08:00
parent a4b77f8baa
commit 1401408c7a
2 changed files with 78 additions and 22 deletions

View File

@ -96,9 +96,12 @@ def migrate(
_prefix_count(sqlite, "draft"), _prefix_count(sqlite, "draft"),
"muse_knowledge_draft", "muse_knowledge_draft",
) )
pg_embeddings = _count(pg, "example_knowledge_embedding")
if _blob_count(sqlite) != pg_embeddings:
_copy_embeddings(pg, sqlite)
report["ledger"]["example_knowledge_embedding"] = _pair( report["ledger"]["example_knowledge_embedding"] = _pair(
_count(pg, "example_knowledge_embedding"), pg_embeddings,
_count(pg, "example_knowledge_embedding"), _blob_count(sqlite),
"example_knowledge_embedding", "example_knowledge_embedding",
) )
report["ledger"]["muse_knowledge_entity"] = _pair( report["ledger"]["muse_knowledge_entity"] = _pair(
@ -143,7 +146,9 @@ def migrate(
) )
embeddings = _copy_embeddings(pg, sqlite) embeddings = _copy_embeddings(pg, sqlite)
report["ledger"]["example_knowledge_embedding"] = _pair( report["ledger"]["example_knowledge_embedding"] = _pair(
embeddings, embeddings, "example_knowledge_embedding" embeddings,
_blob_count(sqlite),
"example_knowledge_embedding",
) )
entities = _copy_entities(pg, sqlite) entities = _copy_entities(pg, sqlite)
report["ledger"]["muse_knowledge_entity"] = _pair( report["ledger"]["muse_knowledge_entity"] = _pair(
@ -232,6 +237,14 @@ def _prefix_count(sqlite, prefix: str, table: str = "cards") -> int:
) )
def _blob_count(sqlite) -> int:
return int(
sqlite.execute(
"SELECT COUNT(*) FROM cards WHERE embedding IS NOT NULL"
).fetchone()[0]
)
def _copy_prefixed_cards(pg, sqlite, table: str, prefix: str) -> dict[str, int]: def _copy_prefixed_cards(pg, sqlite, table: str, prefix: str) -> dict[str, int]:
sqlite.execute("DELETE FROM cards WHERE id LIKE ?", (f"{prefix}:%",)) sqlite.execute("DELETE FROM cards WHERE id LIKE ?", (f"{prefix}:%",))
pg_count = _count(pg, table) pg_count = _count(pg, table)
@ -342,6 +355,12 @@ def _copy_drafts(pg, sqlite) -> int:
def _copy_embeddings(pg, sqlite) -> int: def _copy_embeddings(pg, sqlite) -> int:
"""一行 embedding 一张卡,禁止按 draft_id 折叠。"""
sqlite.execute("DELETE FROM cards WHERE id LIKE 'embedding:%'")
sqlite.execute(
"UPDATE cards SET embedding=NULL WHERE id LIKE 'draft:%' OR id LIKE 'entity:%'"
)
pg_count = _count(pg, "example_knowledge_embedding") pg_count = _count(pg, "example_knowledge_embedding")
cur = pg.execute( cur = pg.execute(
"SELECT id, draft_id, entity_id, content_hash, embedding::text FROM example_knowledge_embedding" "SELECT id, draft_id, entity_id, content_hash, embedding::text FROM example_knowledge_embedding"
@ -350,24 +369,29 @@ def _copy_embeddings(pg, sqlite) -> int:
for row in cur: for row in cur:
vector = _parse_vector(row[4]) vector = _parse_vector(row[4])
blob = pack_embedding(vector) if vector else None blob = pack_embedding(vector) if vector else None
card_id = f"draft:{row[1]}" if row[1] is not None else f"entity:{row[2]}" owner = f"draft:{row[1]}" if row[1] is not None else f"entity:{row[2]}"
sqlite.execute( sqlite.execute(
"""UPDATE cards SET embedding=?, content_hash=? WHERE id=?""", """INSERT INTO cards(id, kind, title, payload_json, embedding, content_hash, created_at)
(blob, row[3], card_id), VALUES (?, 'embedding', ?, ?, ?, ?, datetime('now'))
ON CONFLICT(id) DO UPDATE SET
embedding=excluded.embedding, content_hash=excluded.content_hash""",
(
f"embedding:{row[0]}",
owner,
_json({"draft_id": row[1], "entity_id": row[2]}),
blob,
row[3],
),
) )
if sqlite.execute("SELECT changes()").fetchone()[0] == 0 and blob is not None:
sqlite.execute(
"""INSERT INTO cards(id, kind, title, payload_json, embedding, content_hash, created_at)
VALUES (?, 'embedding', ?, '{}', ?, ?, datetime('now'))
ON CONFLICT(id) DO UPDATE SET embedding=excluded.embedding""",
(f"embedding:{row[0]}", card_id, blob, row[3]),
)
n += 1 n += 1
if n % 1000 == 0: if n % 1000 == 0:
sqlite.commit() sqlite.commit()
sqlite.commit() sqlite.commit()
if n != pg_count: blobs = _blob_count(sqlite)
raise RuntimeError(f"example_knowledge_embedding 对账失败 pg={pg_count} sqlite={n}") if n != pg_count or blobs != pg_count:
raise RuntimeError(
f"example_knowledge_embedding 对账失败 pg={pg_count} copied={n} blobs={blobs}"
)
return n return n
@ -435,16 +459,10 @@ def _cosine_samples(pg, sqlite, n: int = 5) -> list[dict[str, Any]]:
samples = [] samples = []
for row in rows: for row in rows:
pg_vec = _parse_vector(row[3]) pg_vec = _parse_vector(row[3])
card_id = f"draft:{row[1]}" if row[1] is not None else f"entity:{row[2]}"
blob_row = sqlite.execute( blob_row = sqlite.execute(
"SELECT embedding FROM cards WHERE id=? AND embedding IS NOT NULL", "SELECT embedding FROM cards WHERE id=? AND embedding IS NOT NULL",
(card_id,), (f"embedding:{row[0]}",),
).fetchone() ).fetchone()
if blob_row is None:
blob_row = sqlite.execute(
"SELECT embedding FROM cards WHERE id=?",
(f"embedding:{row[0]}",),
).fetchone()
if blob_row is None or blob_row[0] is None: if blob_row is None or blob_row[0] is None:
raise RuntimeError(f"抽样向量缺失 embedding id={row[0]}") raise RuntimeError(f"抽样向量缺失 embedding id={row[0]}")
sqlite_vec = unpack_embedding(blob_row[0]) sqlite_vec = unpack_embedding(blob_row[0])

View File

@ -17,6 +17,7 @@ ROOT = next(
if str(ROOT) not in sys.path: if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT)) sys.path.insert(0, str(ROOT))
from muse.migrate_pg_to_sqlite import _blob_count, _copy_embeddings # noqa: E402
from muse.store import ( # noqa: E402 from muse.store import ( # noqa: E402
add_review, add_review,
connect, connect,
@ -106,6 +107,43 @@ class MuseStoreTest(unittest.TestCase):
self.assertEqual(revision["before_text"], "旧稿") self.assertEqual(revision["before_text"], "旧稿")
self.assertEqual(revision["after_text"], "新稿") self.assertEqual(revision["after_text"], "新稿")
def test_copy_embeddings_persists_one_card_per_row_not_collapsed_by_draft(self) -> None:
class _Rows:
def __init__(self, rows):
self._rows = rows
def fetchone(self):
return self._rows[0] if self._rows else None
def __iter__(self):
return iter(self._rows)
class _FakePG:
def execute(self, sql, params=None):
if "COUNT(*)" in sql:
return _Rows([(3,)])
return _Rows(
[
(101, 10, None, "h1", "[1.0, 0.0]"),
(102, 10, None, "h2", "[0.0, 1.0]"),
(103, 10, None, "h3", "[0.5, 0.5]"),
]
)
conn = connect(self.db)
copied = _copy_embeddings(_FakePG(), conn)
blobs = _blob_count(conn)
ids = [
row[0]
for row in conn.execute(
"SELECT id FROM cards WHERE embedding IS NOT NULL ORDER BY id"
)
]
conn.close()
self.assertEqual(copied, 3)
self.assertEqual(blobs, 3)
self.assertEqual(ids, ["embedding:101", "embedding:102", "embedding:103"])
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()