build_vector_index.py 6.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202
  1. #!/usr/bin/env python3
  2. """
  3. Build vector index for RAG v2 hybrid search.
  4. Embeds all chunks via local Ollama bge-m3, saves as numpy .npy + metadata pickle.
  5. Supports checkpoint/resume for large corpora.
  6. Usage:
  7. python build_vector_index.py [--base-dir /data/knowledge/maritime]
  8. """
  9. import json
  10. import os
  11. import pickle
  12. import sys
  13. import time
  14. import urllib.request
  15. from pathlib import Path
  16. OLLAMA_EMBED_URL = os.environ.get('OLLAMA_EMBED_URL', 'http://127.0.0.1:11435/api/embed')
  17. EMBED_MODEL = os.environ.get('EMBED_MODEL', 'bge-m3')
  18. BATCH_SIZE = 32
  19. CHECKPOINT_INTERVAL = 500 # Save every N chunks
  20. def embed_batch(texts, model=None):
  21. """Embed a batch of texts via Ollama"""
  22. if model is None:
  23. model = EMBED_MODEL
  24. body = json.dumps({
  25. "model": model,
  26. "input": texts,
  27. }).encode('utf-8')
  28. req = urllib.request.Request(
  29. OLLAMA_EMBED_URL,
  30. data=body,
  31. headers={"Content-Type": "application/json"},
  32. method="POST"
  33. )
  34. with urllib.request.urlopen(req, timeout=300) as resp:
  35. data = json.loads(resp.read())
  36. return data.get('embeddings', [])
  37. def main():
  38. import argparse
  39. parser = argparse.ArgumentParser()
  40. parser.add_argument('--base-dir', default=os.environ.get('KNOWLEDGE_DIR', '/data/knowledge/maritime'))
  41. args = parser.parse_args()
  42. base = Path(args.base_dir)
  43. index_path = base / 'rag_index.json'
  44. vec_path = base / 'rag_vectors_v2.npy'
  45. meta_path = base / 'rag_vectors_meta_v2.pkl'
  46. checkpoint_path = base / 'rag_vectors_checkpoint.pkl'
  47. print("=" * 60)
  48. print("Building Vector Index (bge-m3 via Ollama)")
  49. print("=" * 60)
  50. # Load chunks from existing index
  51. if not index_path.exists():
  52. print(f"ERROR: {index_path} not found. Run rebuild_index.py first.")
  53. sys.exit(1)
  54. print(f"Loading chunks from {index_path}...")
  55. t0 = time.time()
  56. with open(index_path) as f:
  57. data = json.load(f)
  58. chunks = data['chunks']
  59. print(f" Loaded {len(chunks)} chunks in {time.time()-t0:.1f}s")
  60. # Check for checkpoint (resume support)
  61. all_embeddings = []
  62. all_meta = []
  63. start_idx = 0
  64. if checkpoint_path.exists():
  65. print("Found checkpoint, resuming...")
  66. with open(checkpoint_path, 'rb') as f:
  67. cp = pickle.load(f)
  68. all_embeddings = cp['embeddings']
  69. all_meta = cp['meta']
  70. start_idx = cp['next_idx']
  71. print(f" Resumed from index {start_idx} ({len(all_embeddings)} embeddings done)")
  72. # Test embedding API
  73. print("\nTesting Ollama embedding API...")
  74. test_vecs = embed_batch(["test"])
  75. if not test_vecs:
  76. print("ERROR: Ollama embedding API not available")
  77. sys.exit(1)
  78. dim = len(test_vecs[0])
  79. print(f" OK, embedding dim: {dim}")
  80. # Embed all chunks
  81. total = len(chunks)
  82. remaining = total - start_idx
  83. print(f"\nEmbedding {remaining} chunks (batch_size={BATCH_SIZE})...")
  84. print(f"Estimated time: ~{remaining / BATCH_SIZE * 0.5 / 60:.1f} minutes")
  85. batch_texts = []
  86. batch_indices = []
  87. errors = 0
  88. embed_start = time.time()
  89. for i in range(start_idx, total):
  90. chunk = chunks[i]
  91. text = chunk['text'][:2000] # bge-m3 max ~8192 tokens, truncate for safety
  92. batch_texts.append(text)
  93. batch_indices.append(i)
  94. if len(batch_texts) >= BATCH_SIZE or i == total - 1:
  95. try:
  96. vecs = embed_batch(batch_texts)
  97. if len(vecs) == len(batch_texts):
  98. all_embeddings.extend(vecs)
  99. for j, idx in enumerate(batch_indices):
  100. c = chunks[idx]
  101. all_meta.append({
  102. 'source': c['source'],
  103. 'chunk_id': c['chunk_id'],
  104. 'text': c['text'],
  105. })
  106. else:
  107. # Partial result, embed one by one
  108. for t in batch_texts:
  109. v = embed_batch([t])
  110. if v:
  111. all_embeddings.append(v[0])
  112. else:
  113. all_embeddings.append([0.0] * dim)
  114. errors += 1
  115. for idx in batch_indices:
  116. c = chunks[idx]
  117. all_meta.append({
  118. 'source': c['source'],
  119. 'chunk_id': c['chunk_id'],
  120. 'text': c['text'],
  121. })
  122. except Exception as e:
  123. errors += 1
  124. # Fill with zeros for failed batch
  125. for _ in batch_texts:
  126. all_embeddings.append([0.0] * dim)
  127. for idx in batch_indices:
  128. c = chunks[idx]
  129. all_meta.append({
  130. 'source': c['source'],
  131. 'chunk_id': c['chunk_id'],
  132. 'text': c['text'],
  133. })
  134. if errors <= 3:
  135. print(f" ERROR at batch {i}: {e}")
  136. batch_texts = []
  137. batch_indices = []
  138. # Progress
  139. done = len(all_embeddings)
  140. if done % (CHECKPOINT_INTERVAL) < BATCH_SIZE or i == total - 1:
  141. elapsed = time.time() - embed_start
  142. rate = (done - start_idx) / max(elapsed, 1)
  143. eta = (total - done) / max(rate, 0.01) / 60
  144. print(f" [{done}/{total}] {rate:.1f} chunks/s, ETA: {eta:.1f}min, errors: {errors}")
  145. # Save checkpoint
  146. with open(checkpoint_path, 'wb') as f:
  147. pickle.dump({
  148. 'embeddings': all_embeddings,
  149. 'meta': all_meta,
  150. 'next_idx': i + 1,
  151. }, f, protocol=pickle.HIGHEST_PROTOCOL)
  152. # Save final index
  153. print(f"\nSaving vector index...")
  154. try:
  155. import numpy as np
  156. vectors = np.array(all_embeddings, dtype=np.float32)
  157. np.save(str(vec_path), vectors)
  158. print(f" Vectors: {vec_path} ({vectors.shape}, {vec_path.stat().st_size/1024/1024:.1f} MB)")
  159. except ImportError:
  160. # Fallback: save as pickle
  161. vec_path = vec_path.with_suffix('.pkl')
  162. with open(vec_path, 'wb') as f:
  163. pickle.dump(all_embeddings, f, protocol=pickle.HIGHEST_PROTOCOL)
  164. print(f" Vectors (pickle): {vec_path}")
  165. with open(meta_path, 'wb') as f:
  166. pickle.dump(all_meta, f, protocol=pickle.HIGHEST_PROTOCOL)
  167. print(f" Meta: {meta_path} ({meta_path.stat().st_size/1024/1024:.1f} MB)")
  168. # Cleanup checkpoint
  169. if checkpoint_path.exists():
  170. checkpoint_path.unlink()
  171. print(" Checkpoint removed")
  172. total_time = time.time() - embed_start
  173. print(f"\nDone! {len(all_embeddings)} vectors in {total_time/60:.1f} minutes, {errors} errors")
  174. if __name__ == '__main__':
  175. main()