框架: embedding 按行落卡并用 blob COUNT 对账
This commit is contained in:
parent
a4b77f8baa
commit
1401408c7a
@ -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])
|
||||||
|
|||||||
@ -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()
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user