Repository navigation
Expand file tree
/
Copy pathcreate_knowledge_base.py
More file actions
114 lines (96 loc) · 3.65 KB
/
Copy pathcreate_knowledge_base.py
File metadata and controls
114 lines (96 loc) · 3.65 KB
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 os
import re
import faiss
import numpy as np
import pickle
import torch
from sentence_transformers import SentenceTransformer
# --- Configuration ---
INPUT_FILE = "scraped_content.txt"
FAISS_INDEX_FILE = "faiss_index.bin"
CHUNKS_FILE = "text_chunks.pkl"
CHUNK_SIZE = 500 # The approximate number of characters in each chunk
CHUNK_OVERLAP = 50 # Number of characters to overlap between chunks
# --- Auto-detect Device (GPU or CPU) ---
# This checks if a CUDA-enabled GPU is available and sets the device accordingly.
device = 'cuda' if torch.cuda.is_available() else 'cpu'
print(f"Using device: {device}")
# --- Load the Local Embedding Model ---
# This will download the model from Hugging Face the first time you run it
# and save it for offline use.
print("Loading the local embedding model (all-MiniLM-L6-v2)...")
embedding_model = SentenceTransformer('all-MiniLM-L6-v2', device=device)
print("Model loaded successfully.")
def clean_text(text):
"""
Cleans the input text by removing multiple newlines and spaces.
"""
text = re.sub(r'\n+', ' ', text)
text = re.sub(r'\s+', ' ', text)
return text.strip()
def chunk_text(text, size=CHUNK_SIZE, overlap=CHUNK_OVERLAP):
"""
Splits the text into overlapping chunks.
"""
chunks = []
start = 0
while start < len(text):
end = start + size
chunks.append(text[start:end])
start += size - overlap
return chunks
def get_embedding(text):
"""
Generates an embedding for a given text using the local Sentence Transformer model.
"""
# The .encode() method converts text into a numerical vector.
return embedding_model.encode(text)
def create_knowledge_base():
"""
Creates a knowledge base by processing a text file, generating local embeddings,
and storing them in a FAISS index.
"""
# --- 1. Load and Process Data ---
try:
with open(INPUT_FILE, 'r', encoding='utf-8') as f:
raw_text = f.read()
except FileNotFoundError:
print(f"Error: The file {INPUT_FILE} was not found.")
print("Please run the scraper.py script first to generate it.")
return
print("Cleaning and chunking text...")
cleaned_text = clean_text(raw_text)
text_chunks = chunk_text(cleaned_text)
text_chunks = [chunk for chunk in text_chunks if chunk.strip()]
if not text_chunks:
print("No text chunks were generated. Check the input file.")
return
print(f"Generated {len(text_chunks)} text chunks.")
# --- 2. Generate Embeddings Locally ---
print("Generating embeddings for text chunks using the local model...")
# We use the model's .encode() method with the full list for efficiency.
embeddings = embedding_model.encode(text_chunks, show_progress_bar=True)
# Convert to NumPy array
embeddings_np = np.array(embeddings).astype('float32')
# --- 3. Create and Save FAISS Index ---
print("Creating FAISS index...")
# Get the dimension of the embeddings from the model's output
d = embeddings_np.shape[1]
index = faiss.IndexFlatL2(d)
index.add(embeddings_np)
try:
faiss.write_index(index, FAISS_INDEX_FILE)
print(f"FAISS index saved to {FAISS_INDEX_FILE}")
except IOError as e:
print(f"Error saving FAISS index: {e}")
return
# --- 4. Save Text Chunks ---
try:
with open(CHUNKS_FILE, 'wb') as f:
pickle.dump(text_chunks, f)
print(f"Text chunks saved to {CHUNKS_FILE}")
except IOError as e:
print(f"Error saving text chunks: {e}")
print("\nKnowledge base creation complete! You can now move to Phase 3.")
if __name__ == "__main__":
create_knowledge_base()