Files
anti-plagiarism/services/worker-gpu/app/faiss_manager.py
jze9 38b000703a fix(gpu): рабочий FAISS-индекс вместо необучаемого IVFFlat
Семантический поиск (уровень 3) не работал вообще: индекс IndexIVFFlat с
nlist=1024 требует ~40000 векторов для обучения, а до обучения векторы
складывались во временный in-memory _flat_index, который:
  - не сохранялся на диск (save() писал только пустой _index) → терялся при
    рестарте;
  - не участвовал в поиске (search() искал только в необученном _index и
    сразу возвращал []).
Итог: в FAISS всегда было 0 векторов, семантика возвращала пусто.

Заменено на IndexIDMap2(IndexFlatIP): без обучения, работает с первого
вектора, doc_id хранится внутри индекса, корректно персистится. На
нормализованных векторах inner product = cosine, порог 0.75 сохраняет смысл.
Добавлена идемпотентность add_vectors (remove_ids перед add).

Проверено: 66 документов → ntotal=66, поиск возвращает релевантные
результаты со score 0.72-0.78, round-trip save/load сохраняет векторы.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-23 19:28:34 +05:00

184 lines
7.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Singleton менеджер FAISS индекса.
Используется плоский индекс IndexFlatIP, обёрнутый в IndexIDMap2, что даёт:
- отсутствие этапа обучения (в отличие от IVFFlat) — индекс работоспособен сразу,
начиная с первого вектора;
- хранение doc_id прямо внутри индекса (add_with_ids) — не нужен отдельный
маппинг на диске, а поиск сразу возвращает doc_id из PostgreSQL;
- корректную персистентность через faiss.write_index / read_index.
Векторы модели нормализованы (normalize_embeddings=True), поэтому inner product
эквивалентен косинусной близости. Для масштаба проекта (сотни тысяч — единицы
миллионов документов) полный перебор по FlatIP по скорости приемлем.
"""
import logging
import os
from pathlib import Path
import numpy as np
from app.config import settings
logger = logging.getLogger(__name__)
class FAISSManager:
"""Singleton для управления FAISS индексом (IndexIDMap2 поверх IndexFlatIP)."""
_index = None
# doc_id -> faiss_id. Для IDMap2 faiss_id == doc_id, но маппинг сохраняем
# для совместимости с вызывающим кодом (plagiarism.embed_documents).
_reverse_map: dict[int, int] = {}
_use_gpu: bool = False
@classmethod
def _new_index(cls):
"""Создать новый пустой индекс нужного типа."""
import faiss
base = faiss.IndexFlatIP(settings.EMBED_DIM)
return faiss.IndexIDMap2(base)
@classmethod
def load_or_create(cls) -> None:
"""Загрузить индекс с диска или создать новый.
Старый несовместимый индекс (например, IVFFlat от предыдущей версии,
который не хранит id_map) безопасно пересоздаётся — полезных векторов в
нём всё равно не было.
"""
import faiss
index_path = settings.FAISS_INDEX_PATH
if os.path.exists(index_path):
try:
loaded = faiss.read_index(index_path)
if hasattr(loaded, "id_map"):
cls._index = loaded
cls._rebuild_reverse_map()
logger.info(
f"Загрузка FAISS индекса из {index_path} "
f"({cls._index.ntotal} векторов)"
)
else:
logger.warning(
"На диске несовместимый FAISS индекс (без id_map) — "
"пересоздаём как IndexIDMap2(IndexFlatIP)"
)
cls._index = cls._new_index()
cls._reverse_map = {}
except Exception as e:
logger.warning(f"Не удалось загрузить FAISS индекс ({e}) — создаём новый")
cls._index = cls._new_index()
cls._reverse_map = {}
else:
logger.info("Создание нового FAISS индекса IndexIDMap2(IndexFlatIP)...")
cls._index = cls._new_index()
cls._reverse_map = {}
cls._use_gpu = False # FlatIP на CPU достаточно быстр для целевого масштаба
@classmethod
def _ensure(cls) -> None:
"""Ленивая инициализация индекса при первом обращении."""
if cls._index is None:
cls.load_or_create()
@classmethod
def _rebuild_reverse_map(cls) -> None:
"""Восстановить _reverse_map из id, хранящихся внутри загруженного индекса."""
import faiss
try:
ids = faiss.vector_to_array(cls._index.id_map)
cls._reverse_map = {int(i): int(i) for i in ids}
except Exception as e:
logger.warning(f"Не удалось восстановить reverse_map из индекса: {e}")
cls._reverse_map = {}
@classmethod
def search(cls, query_vector: np.ndarray, k: int = 20) -> list[tuple[int, float]]:
"""Поиск k ближайших векторов.
Args:
query_vector: Нормализованный вектор запроса, форма (768,)
k: Количество результатов
Returns:
Список кортежей (doc_id, cosine_score), отсортированных по убыванию score
"""
cls._ensure()
if cls._index.ntotal == 0:
return []
try:
query = query_vector.reshape(1, -1).astype(np.float32)
distances, ids = cls._index.search(query, min(k, cls._index.ntotal))
results = []
for idx, dist in zip(ids[0], distances[0]):
if idx == -1:
continue
# Для IDMap2 idx — это уже doc_id из PostgreSQL
results.append((int(idx), float(dist)))
return results
except Exception as e:
logger.error(f"Ошибка поиска FAISS: {e}")
return []
@classmethod
def add_vectors(cls, vectors: np.ndarray, doc_ids: list[int]) -> None:
"""Добавить (или обновить) векторы в индекс.
Идемпотентно по doc_id: при повторном эмбеддинге старый вектор документа
удаляется перед добавлением нового, чтобы не плодить дубли.
Args:
vectors: numpy массив формы (N, 768)
doc_ids: Список doc_id из PostgreSQL
"""
cls._ensure()
if len(doc_ids) == 0:
return
vectors = np.asarray(vectors, dtype=np.float32)
ids = np.asarray(doc_ids, dtype=np.int64)
# Удалить существующие id, чтобы повторный эмбеддинг не создавал дубли
try:
cls._index.remove_ids(ids)
except Exception:
pass
cls._index.add_with_ids(vectors, ids)
for doc_id in doc_ids:
cls._reverse_map[int(doc_id)] = int(doc_id)
logger.info(f"Добавлено {len(doc_ids)} векторов в FAISS. Всего: {cls._index.ntotal}")
@classmethod
def save(cls) -> None:
"""Сохранить индекс на диск."""
cls._ensure()
import faiss
index_path = Path(settings.FAISS_INDEX_PATH)
index_path.parent.mkdir(parents=True, exist_ok=True)
faiss.write_index(cls._index, str(index_path))
logger.info(f"FAISS индекс сохранён: {index_path} ({cls._index.ntotal} векторов)")
@classmethod
def total_vectors(cls) -> int:
"""Количество векторов в индексе."""
if cls._index is None:
return 0
return cls._index.ntotal