chatbot / embedding.py
ogx786's picture
Create embedding.py
3d7b0fd verified
Raw
History Blame Contribute Delete
4.2 kB
import os
import pickle
import numpy as np
import torch
import faiss
from tqdm import tqdm
from sentence_transformers import SentenceTransformer
# ==================================================
# CONFIG
# ==================================================
INPUT_FILE = "hbl_final_chunks.pkl"
OUTPUT_DIR = "hbl_vector_store"
MODEL_NAME = "BAAI/bge-m3"
BATCH_SIZE = 64
os.makedirs(
OUTPUT_DIR,
exist_ok=True
)
# ==================================================
# GPU CHECK
# ==================================================
print("=" * 60)
print(
"PyTorch:",
torch.__version__
)
print(
"CUDA available:",
torch.cuda.is_available()
)
if torch.cuda.is_available():
device = "cuda"
print(
"GPU:",
torch.cuda.get_device_name(0)
)
else:
device = "cpu"
print(
"Running on CPU"
)
print("=" * 60)
# ==================================================
# LOAD CHUNKS
# ==================================================
print("\nLoading chunks...")
with open(
INPUT_FILE,
"rb"
) as f:
chunks = pickle.load(f)
print(
"Total chunks:",
len(chunks)
)
# ==================================================
# PREPARE TEXT
# ==================================================
print("\nPreparing text...")
texts = []
for c in chunks:
text = c.get(
"text",
""
)
if not text.strip():
text = "empty document"
texts.append(text)
print(
"Texts:",
len(texts)
)
# ==================================================
# LOAD EMBEDDING MODEL
# ==================================================
print("\nLoading model...")
model = SentenceTransformer(
MODEL_NAME,
device=device
)
# BGE-M3 supports long context
model.max_seq_length = 8192
print(
"Model loaded"
)
# ==================================================
# CREATE EMBEDDINGS
# ==================================================
print("\nCreating embeddings...")
embeddings = model.encode(
texts,
batch_size=BATCH_SIZE,
show_progress_bar=True,
convert_to_numpy=True,
normalize_embeddings=True
)
print(
"Embedding shape:",
embeddings.shape
)
# ==================================================
# SAVE RAW EMBEDDINGS
# ==================================================
embedding_file = os.path.join(
OUTPUT_DIR,
"hbl_embeddings.npy"
)
np.save(
embedding_file,
embeddings
)
print(
"Saved:",
embedding_file
)
# ==================================================
# BUILD FAISS INDEX
# ==================================================
print("\nBuilding FAISS index...")
dimension = embeddings.shape[1]
print(
"Vector dimension:",
dimension
)
# Inner product on normalized vectors = cosine similarity
index = faiss.IndexFlatIP(
dimension
)
index.add(
embeddings.astype(
np.float32
)
)
print(
"FAISS vectors:",
index.ntotal
)
# ==================================================
# SAVE FAISS
# ==================================================
faiss_file = os.path.join(
OUTPUT_DIR,
"hbl_faiss.index"
)
faiss.write_index(
index,
faiss_file
)
print(
"Saved:",
faiss_file
)
# ==================================================
# SAVE METADATA
# ==================================================
metadata_file = os.path.join(
OUTPUT_DIR,
"hbl_metadata.pkl"
)
with open(
metadata_file,
"wb"
) as f:
pickle.dump(
chunks,
f
)
print(
"Saved:",
metadata_file
)
# ==================================================
# TEST SEARCH
# ==================================================
print("\nTesting retrieval...")
query = "What are the charges for foreign currency remittance?"
query_embedding = model.encode(
[query],
normalize_embeddings=True
)
scores, ids = index.search(
query_embedding.astype(np.float32),
5
)
for rank, idx in enumerate(ids[0]):
print("\nRank:", rank+1)
print(
"Score:",
scores[0][rank]
)
print(
chunks[idx]["text"][:300]
)
print("\nDONE")
print("="*60)