QuoteSearch / offline_processing.py
ruidiao's picture
Switch to potion-base-4M static embedding model
a1252ee
Raw History Blame Contribute Delete
4.2 kB
import pandas as pd
import numpy as np
from model2vec import StaticModel
import json
import gzip
import struct
import os
MODEL_NAME = 'minishlab/potion-base-4M'
EMBEDDING_DIM = 128
INPUT_CSV_PATH = 'data/quotes.csv'
OUTPUT_BINARY_PATH = 'data/quotes_index.bin'
WEIGHTS_PATH = 'data/potion_weights.bin'
TOKENIZER_PATH = 'data/potion_tokenizer.json'
def load_quotes(file_path):
column_names = ['quote', 'author', 'category']
df = pd.read_csv(file_path, sep=',', quotechar='"', header=None, names=column_names, on_bad_lines='skip')
initial_rows = len(df)
df = df[df['category'].apply(lambda x: isinstance(x, str) and not any(c.isupper() for c in x))]
if initial_rows - len(df) > 0:
print(f"Ignored {initial_rows - len(df)} rows due to uppercase letters in category.")
df['author'] = df['author'].fillna('').astype(str)
grouped = {}
for _, row in df.iterrows():
quote = row['quote']
author = row['author']
quote_key = quote.lower().strip() if isinstance(quote, str) else ''
author_key = author.lower().strip() if isinstance(author, str) else ''
key = (quote_key, author_key)
if key not in grouped:
grouped[key] = {'quote': quote, 'author': author}
records = []
for key, data in grouped.items():
records.append({
'quote': data['quote'],
'author': data['author'] if data['author'] != '' else None
})
return records
def generate_embeddings(quotes, model):
texts = [q['quote'] for q in quotes]
embeddings = model.encode(texts)
return embeddings
def quantize_embeddings(embeddings):
abs_max = np.abs(embeddings).max()
scale = 127.0 / abs_max if abs_max != 0 else 0
quantized = np.clip(embeddings * scale, -127, 127).astype(np.int8)
return quantized, scale
def export_model_weights(model):
weights = model.embedding
weights_f16 = weights.astype(np.float16)
with open(WEIGHTS_PATH, 'wb') as f:
f.write(weights_f16.tobytes())
file_size = os.path.getsize(WEIGHTS_PATH)
print(f"Exported model weights: {weights_f16.shape} ({file_size / 1e6:.1f} MB)")
def export_tokenizer(model):
vocab = {token: i for i, token in enumerate(model.tokens)}
tokenizer_data = {
"vocab": vocab,
"unk_id": model.unk_token_id
}
with open(TOKENIZER_PATH, 'w', encoding='utf-8') as f:
json.dump(tokenizer_data, f, separators=(',', ':'))
file_size = os.path.getsize(TOKENIZER_PATH)
print(f"Exported tokenizer vocab ({len(vocab)} tokens, {file_size / 1e6:.2f} MB)")
def main():
print("Starting offline processing...")
quotes = load_quotes(INPUT_CSV_PATH)
print(f"Loaded {len(quotes)} quotes.")
print("Loading model...")
model = StaticModel.from_pretrained(MODEL_NAME)
print(f"Model dimension: {model.dim}")
export_model_weights(model)
export_tokenizer(model)
print("Generating embeddings...")
float_embeddings = generate_embeddings(quotes, model)
print(f"Generated float embeddings with shape: {float_embeddings.shape}")
quantized_embeddings, scale = quantize_embeddings(float_embeddings)
print(f"Quantized embeddings: {quantized_embeddings.shape}, scale: {scale}")
metadata = [{"quote": q["quote"], "author": q["author"]} for q in quotes]
for item in metadata:
for k, v in list(item.items()):
if isinstance(v, float) and np.isnan(v):
item[k] = None
metadata_json = json.dumps(metadata, separators=(",", ":"))
metadata_bytes = gzip.compress(metadata_json.encode('utf-8'))
metadata_format = 1
with open(OUTPUT_BINARY_PATH, 'wb') as f:
f.write(struct.pack('<I', len(quotes)))
f.write(struct.pack('<H', EMBEDDING_DIM))
f.write(struct.pack('<f', scale))
f.write(struct.pack('<I', len(metadata_bytes)))
f.write(struct.pack('<B', metadata_format))
f.write(metadata_bytes)
f.write(quantized_embeddings.tobytes())
index_size = os.path.getsize(OUTPUT_BINARY_PATH)
print(f"Offline processing complete. Index: {OUTPUT_BINARY_PATH} ({index_size / 1e6:.1f} MB)")
if __name__ == "__main__":
main()