File size: 4,195 Bytes
c96a6bc
 
a1252ee
c96a6bc
0ac07c8
c96a6bc
a1252ee
c96a6bc
a1252ee
 
c96a6bc
 
a1252ee
 
c96a6bc
 
 
a1252ee
c96a6bc
 
a1252ee
 
da9b599
 
 
 
 
 
 
 
 
a1252ee
da9b599
 
 
 
a1252ee
da9b599
 
c96a6bc
a1252ee
 
 
 
 
c96a6bc
 
 
a1252ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c96a6bc
 
 
 
 
 
a1252ee
 
 
 
 
 
 
c96a6bc
a1252ee
c96a6bc
 
 
a1252ee
c96a6bc
a1252ee
c96a6bc
a1252ee
 
 
c96a6bc
 
a1252ee
0ac07c8
c96a6bc
 
a1252ee
 
 
 
 
c96a6bc
 
 
a1252ee
 
c96a6bc
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
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()