Skip to content

Commit 0d3fc09

Browse files
authored
Merge pull request exowanderer#5 from philippesaade-wmde/updated_merge_conflicts
- Improved cache db operations and fixed errors involve db insertion - Improved data storage by storing the smaller base64 data per embedding, which results directly from JinaAI's API - Refactored handling of DataStax and JinaAI errors - Fixed duplicate ID errors - updated our caching and merging of cache - created migrate_db routine to update from original to smaller db
2 parents dcbbf87 + cf08708 commit 0d3fc09

5 files changed

Lines changed: 225 additions & 25 deletions

File tree

‎docker/7_Create_Prototype/run.py‎

Lines changed: 12 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -88,12 +88,19 @@ def process_items(queue, progress_bar):
8888
break # Exit condition for worker processes
8989

9090
item_id = item['id']
91-
item_label = textifier.get_label(item_id, json.loads(item['labels']))
91+
92+
item_label = textifier.get_label(
93+
item_id,
94+
json.loads(item['labels'])
95+
)
96+
9297
item_description = textifier.get_description(
9398
item_id,
9499
json.loads(item['descriptions'])
95100
)
96-
item_aliases = textifier.get_aliases(json.loads(item['aliases']))
101+
item_aliases = textifier.get_aliases(
102+
json.loads(item['aliases'])
103+
)
97104

98105
if item_label is not None:
99106
# TODO: Verify: If label does not exist, then skip item
@@ -128,16 +135,15 @@ def process_items(queue, progress_bar):
128135

129136
graph_store.add_document(
130137
id=f"{item_id}_{LANGUAGE}_{chunk_i+1}",
138+
131139
text=chunk,
132140
metadata=metadata
133141
)
134142

135143
progress_bar.value += 1
136144

137-
while True:
138-
# Leftover Maintenance: Ensure that the batch is emptied out
139-
if not graph_store.push_batch(): # Stop when batch is empty
140-
break
145+
146+
graph_store.push_all()
141147

142148

143149
if __name__ == "__main__":

‎src/merge_cache.py‎

Lines changed: 91 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,91 @@
1+
import sqlite3
2+
import glob
3+
import os
4+
import base64
5+
from tqdm import tqdm
6+
7+
# Define database file pattern
8+
db_files = glob.glob("../data/Wikidata/sqlite_cacheembeddings_*.db")
9+
10+
# Define the target merged database
11+
merged_db = "../data/Wikidata/sqlite_cacheembeddings_merged.db"
12+
TABLE_NAME = "wikidata_prototype"
13+
14+
# Batch size for processing
15+
BATCH_SIZE = 1000 # Adjust based on performance needs
16+
17+
# Create the merged database connection
18+
conn_merged = sqlite3.connect(merged_db)
19+
cursor_merged = conn_merged.cursor()
20+
21+
# Create table in the merged database if it doesn't exist
22+
cursor_merged.execute(f"""
23+
CREATE TABLE IF NOT EXISTS {TABLE_NAME} (
24+
id TEXT PRIMARY KEY,
25+
embedding TEXT
26+
);
27+
""")
28+
conn_merged.commit()
29+
30+
# Helper function to check if a string is a valid Base64 encoding
31+
def is_valid_base64(s):
32+
try:
33+
if not s or not isinstance(s, str):
34+
return False
35+
base64.b64decode(s, validate=True)
36+
return True
37+
except Exception:
38+
return False
39+
40+
# Loop through all source databases
41+
for db_file in db_files:
42+
print(f"Processing {db_file}...")
43+
44+
# Connect to the current database
45+
conn_src = sqlite3.connect(db_file)
46+
cursor_src = conn_src.cursor()
47+
48+
# Get total record count for progress tracking
49+
cursor_src.execute(f"SELECT COUNT(*) FROM {TABLE_NAME}")
50+
total_records = cursor_src.fetchone()[0]
51+
52+
# Fetch records in batches
53+
offset = 0
54+
with tqdm(total=total_records,
55+
desc=f"Merging {db_file}", unit="records") as pbar:
56+
while True:
57+
cursor_src.execute(f"SELECT id, embedding FROM {TABLE_NAME} LIMIT {BATCH_SIZE} OFFSET {offset}")
58+
records = cursor_src.fetchall()
59+
if not records:
60+
break # No more records to process
61+
62+
# Prepare batch for insertion
63+
batch_data = []
64+
for id_, embedding in records:
65+
if embedding and embedding.strip() and is_valid_base64(embedding):
66+
cursor_merged.execute(f"SELECT embedding FROM {TABLE_NAME} WHERE id = ?", (id_,))
67+
existing = cursor_merged.fetchone()
68+
69+
if existing is None or not is_valid_base64(existing[0]):
70+
batch_data.append((id_, embedding, embedding)) # Prepare for bulk insert
71+
72+
# Perform batch insert/update
73+
if batch_data:
74+
cursor_merged.executemany(
75+
f"""
76+
INSERT INTO {TABLE_NAME} (id, embedding)
77+
VALUES (?, ?)
78+
ON CONFLICT(id) DO UPDATE SET embedding = ?
79+
""",
80+
batch_data
81+
)
82+
conn_merged.commit()
83+
84+
offset += BATCH_SIZE # Move to the next batch
85+
pbar.update(len(records)) # Update tqdm progress bar
86+
87+
conn_src.close()
88+
89+
# Close merged database connection
90+
conn_merged.close()
91+
print(f"Merge completed! Combined database saved as {merged_db}")

‎src/migrate_cache.py‎

Lines changed: 77 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,77 @@
1+
import sqlite3
2+
import json
3+
import base64
4+
import numpy as np
5+
from tqdm import tqdm
6+
7+
DB_PATH = "../data/Wikidata/sqlite_cacheembeddings.db"
8+
TABLE_NAME = "wikidata_prototype" # Change this to match your actual table name
9+
BATCH_SIZE = 5000 # Process in smaller batches to avoid memory overload
10+
11+
def convert_embeddings():
12+
"""
13+
Convert JSON-stored embeddings into Base64-encoded binary format in batches.
14+
Uses `fetchmany(BATCH_SIZE)` to process records iteratively.
15+
"""
16+
conn = sqlite3.connect(DB_PATH)
17+
cursor = conn.cursor()
18+
19+
# Check if the embedding column exists (sanity check)
20+
cursor.execute(f"PRAGMA table_info({TABLE_NAME})")
21+
columns = [row[1] for row in cursor.fetchall()]
22+
if "embedding" not in columns:
23+
print("Error: 'embedding' column does not exist in the table!")
24+
return
25+
26+
# Count total records for progress tracking
27+
cursor.execute(f"SELECT COUNT(*) FROM {TABLE_NAME}")
28+
total_records = cursor.fetchone()[0]
29+
30+
print(f"Total records to process: {total_records}")
31+
32+
# Fetch records in batches using an iterator
33+
offset = 0
34+
with tqdm(total=total_records, desc="Converting embeddings", unit="record") as pbar:
35+
while True:
36+
cursor.execute(f"SELECT id, embedding FROM {TABLE_NAME} LIMIT {BATCH_SIZE} OFFSET {offset}")
37+
records = cursor.fetchall()
38+
if not records:
39+
break # Stop when there are no more records
40+
41+
updated_records = []
42+
for id, json_embedding in records:
43+
if json_embedding:
44+
try:
45+
# Convert JSON string to list of floats
46+
embedding_list = json.loads(json_embedding)
47+
48+
# Convert list of floats to Base64-encoded binary
49+
binary_data = np.array(embedding_list, dtype=np.float32).tobytes()
50+
base64_embedding = base64.b64encode(binary_data).decode('utf-8')
51+
52+
updated_records.append((base64_embedding, id))
53+
except Exception as e:
54+
pass
55+
56+
pbar.update(1) # Update progress bar for each record processed
57+
58+
# Update database in batches
59+
if updated_records:
60+
cursor.executemany(
61+
f"UPDATE {TABLE_NAME} SET embedding = ? WHERE id = ?",
62+
updated_records
63+
)
64+
conn.commit() # Commit every batch
65+
66+
offset += BATCH_SIZE # Move to next batch
67+
68+
print("Optimizing database with VACUUM...")
69+
cursor.execute("VACUUM;")
70+
conn.commit()
71+
72+
print("Migration completed successfully.")
73+
74+
conn.close()
75+
76+
if __name__ == "__main__":
77+
convert_embeddings()

‎src/wikidataCache.py‎

Lines changed: 24 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,8 @@
55

66
import os
77
import json
8+
import base64
9+
import numpy as np
810

911
"""
1012
SQLite database setup for caching the query embeddings for a faster
@@ -36,19 +38,28 @@
3638
Base = declarative_base()
3739
Session = sessionmaker(bind=engine)
3840

41+
class EmbeddingType(TypeDecorator):
42+
"""Custom SQLAlchemy type for storing embeddings as Base64 strings in SQLite."""
3943

40-
class JSONType(TypeDecorator):
41-
"""Custom SQLAlchemy type for JSON storage in SQLite."""
4244
impl = Text
4345

4446
def process_bind_param(self, value, dialect):
45-
if value is not None:
46-
return json.dumps(value, separators=(',', ':'))
47+
"""Convert a list of floats (embedding) to a Base64 string before storing."""
48+
if value is not None and isinstance(value, list):
49+
# Convert list to binary
50+
binary_data = np.array(value, dtype=np.float32).tobytes()
51+
# Encode to Base64 string
52+
return base64.b64encode(binary_data).decode('utf-8')
4753
return None
4854

4955
def process_result_value(self, value, dialect):
56+
"""Convert a Base64 string back to a list of floats when retrieving."""
5057
if value is not None:
51-
return json.loads(value)
58+
# Decode Base64
59+
binary_data = base64.b64decode(value)
60+
# Convert back to float32 list
61+
embedding_array = np.frombuffer(binary_data, dtype=np.float32)
62+
return embedding_array.tolist()
5263
return None
5364

5465

@@ -59,7 +70,7 @@ class CacheEmbeddings(Base):
5970
__tablename__ = table_name
6071

6172
id = Column(Text, primary_key=True)
62-
embedding = Column(JSONType)
73+
embedding = Column(EmbeddingType)
6374

6475
@staticmethod
6576
def add_cache(id, embedding):
@@ -99,6 +110,13 @@ def add_bulk_cache(data):
99110
- bool: True if the operation was successful, False otherwise.
100111
"""
101112
worked = False
113+
embeddingtype = EmbeddingType()
114+
for i in range(len(data)):
115+
data[i]['embedding'] = embeddingtype.process_bind_param(
116+
data[i]['embedding'],
117+
None
118+
)
119+
102120
with Session() as session:
103121
exec_text = text(
104122
f"""

‎src/wikidataRetriever.py‎

Lines changed: 21 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@ def __init__(self, datastax_token, collection_name, model='jina', batch_size=8,
1717
from langchain_astradb import AstraDBVectorStore
1818
from astrapy.info import CollectionVectorServiceOptions
1919
from astrapy import DataAPIClient
20+
from astrapy.exceptions import InsertManyException
2021
from multiprocessing import Queue
2122

2223
from transformers import AutoTokenizer
@@ -30,6 +31,7 @@ def __init__(self, datastax_token, collection_name, model='jina', batch_size=8,
3031
self.model = model
3132
self.collection_name = collection_name
3233
self.doc_batch = Queue()
34+
self.InsertManyException = InsertManyException
3335

3436
self.cache_on = (cache_embeddings is not None)
3537
if self.cache_on:
@@ -126,23 +128,29 @@ def push_batch(self):
126128
if len(docs) == 0:
127129
return False
128130

129-
try:
130-
vectors = self.embeddings.embed_documents(
131-
[doc['content'] for doc in docs]
132-
)
133-
self.graph_store.insert_many(docs, vectors=vectors)
134-
except Exception as e:
135-
print(e)
136-
137-
# Put the documents back in the Queue and try again later.
138-
for doc in docs:
139-
self.doc_batch.put(doc)
131+
while True:
132+
try:
133+
vectors = self.embeddings.embed_documents(
134+
[doc['content'] for doc in docs]
135+
)
136+
break
137+
except Exception as e:
138+
print(e)
139+
time.sleep(3)
140140

141-
return False
141+
while True:
142+
try:
143+
self.graph_store.insert_many(docs, vectors=vectors)
144+
break
145+
except self.InsertManyException as e:
146+
break
147+
except Exception as e:
148+
print(e)
149+
time.sleep(3)
142150

143151
self.cache_model.add_bulk_cache([{
144152
'id': docs[i]['_id'],
145-
'embedding': json.dumps(vectors[i], separators=(',', ':'))}
153+
'embedding': vectors[i]}
146154
for i in range(len(docs))])
147155

148156
return True

0 commit comments

Comments
 (0)