框架: embedding 按行落卡并用 blob COUNT 对账
This commit is contained in:
parent
a4b77f8baa
commit
1401408c7a
@ -96,9 +96,12 @@ def migrate(
|
||||
_prefix_count(sqlite, "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(
|
||||
_count(pg, "example_knowledge_embedding"),
|
||||
_count(pg, "example_knowledge_embedding"),
|
||||
pg_embeddings,
|
||||
_blob_count(sqlite),
|
||||
"example_knowledge_embedding",
|
||||
)
|
||||
report["ledger"]["muse_knowledge_entity"] = _pair(
|
||||
@ -143,7 +146,9 @@ def migrate(
|
||||
)
|
||||
embeddings = _copy_embeddings(pg, sqlite)
|
||||
report["ledger"]["example_knowledge_embedding"] = _pair(
|
||||
embeddings, embeddings, "example_knowledge_embedding"
|
||||
embeddings,
|
||||
_blob_count(sqlite),
|
||||
"example_knowledge_embedding",
|
||||
)
|
||||
entities = _copy_entities(pg, sqlite)
|
||||
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]:
|
||||
sqlite.execute("DELETE FROM cards WHERE id LIKE ?", (f"{prefix}:%",))
|
||||
pg_count = _count(pg, table)
|
||||
@ -342,6 +355,12 @@ def _copy_drafts(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")
|
||||
cur = pg.execute(
|
||||
"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:
|
||||
vector = _parse_vector(row[4])
|
||||
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(
|
||||
"""UPDATE cards SET embedding=?, content_hash=? WHERE id=?""",
|
||||
(blob, row[3], card_id),
|
||||
"""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, 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
|
||||
if n % 1000 == 0:
|
||||
sqlite.commit()
|
||||
sqlite.commit()
|
||||
if n != pg_count:
|
||||
raise RuntimeError(f"example_knowledge_embedding 对账失败 pg={pg_count} sqlite={n}")
|
||||
blobs = _blob_count(sqlite)
|
||||
if n != pg_count or blobs != pg_count:
|
||||
raise RuntimeError(
|
||||
f"example_knowledge_embedding 对账失败 pg={pg_count} copied={n} blobs={blobs}"
|
||||
)
|
||||
return n
|
||||
|
||||
|
||||
@ -435,16 +459,10 @@ def _cosine_samples(pg, sqlite, n: int = 5) -> list[dict[str, Any]]:
|
||||
samples = []
|
||||
for row in rows:
|
||||
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(
|
||||
"SELECT embedding FROM cards WHERE id=? AND embedding IS NOT NULL",
|
||||
(card_id,),
|
||||
(f"embedding:{row[0]}",),
|
||||
).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:
|
||||
raise RuntimeError(f"抽样向量缺失 embedding id={row[0]}")
|
||||
sqlite_vec = unpack_embedding(blob_row[0])
|
||||
|
||||
@ -17,6 +17,7 @@ ROOT = next(
|
||||
if str(ROOT) not in sys.path:
|
||||
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
|
||||
add_review,
|
||||
connect,
|
||||
@ -106,6 +107,43 @@ class MuseStoreTest(unittest.TestCase):
|
||||
self.assertEqual(revision["before_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__":
|
||||
unittest.main()
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user