/
IvanMysin
/
Topics
Обзор
Документация
Войти
/
IvanMysin
/
Topics
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
common_agents_notes
src/embedding_manager.py
332 строки
15 KB
ivan
Working on preport scripts
18 фев 2026, 14:27
18 фев 2026, 14:27
5b40c9a
Код
Авторство
О чём код?
import chromadb import numpy as np import pandas as pd from transformers import AutoTokenizer from adapters import AutoAdapterModel import torch import sqlite3 from typing import List, Dict, Tuple, Optional import hashlib from tqdm.auto import tqdm import gc import os from config import DB_CONFIG, MODEL_CONFIG, CHROMA_CONFIG class EmbeddingManager: def __init__(self): self.model_name = MODEL_CONFIG["name"] self.embedding_dim = MODEL_CONFIG["embedding_dimension"] self.batch_size = MODEL_CONFIG["batch_size"] # Создаем директорию для ChromaDB если не существует os.makedirs(CHROMA_CONFIG["path"], exist_ok=True) # Инициализируем модель self.model = self._setup_model() # Инициализируем ChromaDB self.chroma_client = chromadb.PersistentClient(path=str(CHROMA_CONFIG["path"])) self.collection = self._get_or_create_collection() def _setup_model(self) -> AutoAdapterModel: """Настраивает модель для инференса""" print(f"🔄 Загрузка модели {self.model_name}...") self.tokenizer = AutoTokenizer.from_pretrained(self.model_name) model = AutoAdapterModel.from_pretrained(self.model_name) model.load_adapter("allenai/specter2", source="hf", set_active=True) print(f"✅ Модель загружена на {model.device}") return model def _get_or_create_collection(self) -> chromadb.Collection: """Создает или загружает коллекцию в ChromaDB""" try: # Сначала проверяем существующие коллекции existing_collections = self.chroma_client.list_collections() collection_names = [col.name for col in existing_collections] if CHROMA_CONFIG["collection_name"] in collection_names: collection = self.chroma_client.get_collection(CHROMA_CONFIG["collection_name"]) print("✅ Существующая коллекция ChromaDB загружена") return collection else: # Коллекция не существует, создаем новую print("📝 Коллекция не найдена, создаем новую...") collection = self.chroma_client.create_collection( name=CHROMA_CONFIG["collection_name"], metadata={"description": "Research papers embeddings", "model": self.model_name}, distance_function=CHROMA_CONFIG["distance_function"] ) print("✅ Новая коллекция ChromaDB создана") return collection except Exception as e: print(f"⚠️ Ошибка при работе с ChromaDB: {e}") print("🔄 Пытаемся создать коллекцию заново...") # Пытаемся создать коллекцию, игнорируя любые ошибки try: collection = self.chroma_client.create_collection( name=CHROMA_CONFIG["collection_name"], metadata={"description": "Research papers embeddings", "model": self.model_name, "distance": CHROMA_CONFIG["distance_function"]} ) print("✅ Коллекция ChromaDB создана после ошибки") return collection except Exception as e2: print(f"❌ Критическая ошибка при создании коллекции: {e2}") raise def _compute_text_hash(self, text: str) -> str: """Вычисляет хэш текста для отслеживания изменений""" return hashlib.md5(text.encode('utf-8')).hexdigest() def _get_existing_documents(self) -> Dict[str, Dict]: """Возвращает информацию о документах уже в ChromaDB""" try: # Проверяем, есть ли документы в коллекции if self.collection.count() == 0: return {} results = self.collection.get(include=["metadatas"]) existing_docs = {} for i, doc_id in enumerate(results["ids"]): existing_docs[doc_id] = { "text_hash": results["metadatas"][i]["text_hash"], "model_version": results["metadatas"][i].get("model_version", "unknown") } return existing_docs except Exception as e: print(f"⚠️ Ошибка при получении данных из ChromaDB: {e}") return {} def load_documents_from_db( self, date_from: Optional[str] = None, date_to: Optional[str] = None, journals_file: Optional[str] = None ) -> pd.DataFrame: """Загружает документы из SQLite с опциональной фильтрацией Args: date_from: Начальная дата фильтрации (формат: YYYY-MM-DD) date_to: Конечная дата фильтрации (формат: YYYY-MM-DD) journals_file: Путь к файлу со списком журналов (по одному на строку) Returns: pd.DataFrame: Отфильтрованные документы """ print("📥 Загрузка документов из базы данных...") # Загружаем список журналов если указан allowed_journals = None if journals_file and os.path.exists(journals_file): with open(journals_file, 'r', encoding='utf-8') as f: allowed_journals = [line.strip() for line in f if line.strip()] print(f"📋 Загружен список из {len(allowed_journals)} журналов") try: conn = sqlite3.connect(DB_CONFIG["db_path"]) # Строим SQL запрос с условиями conditions = [] params = [] if date_from: conditions.append(f"{DB_CONFIG['Articles']['date_column']} >= ?") params.append(date_from) print(f"📅 Фильтр по дате: от {date_from}") if date_to: conditions.append(f"{DB_CONFIG['Articles']['date_column']} <= ?") params.append(date_to) print(f"📅 Фильтр по дате: до {date_to}") if allowed_journals: # Используем IN clause для списка журналов placeholders = ','.join(['?' for _ in allowed_journals]) conditions.append(f"{DB_CONFIG['Articles']['journal_column']} IN ({placeholders})") params.extend(allowed_journals) print(f"📚 Фильтр по {len(allowed_journals)} журналам") query = f""" SELECT * FROM {DB_CONFIG['table_name']} {('WHERE ' + ' AND '.join(conditions)) if conditions else ''} """ df = pd.read_sql_query(query, conn, params=params) conn.close() # Очистка данных initial_count = len(df) df = df.dropna(subset=[DB_CONFIG['text_column']]) df = df[df[DB_CONFIG['text_column']].str.strip().astype(bool)] print(f"✅ Загружено {len(df)}/{initial_count} документов") if len(df) == 0: print("⚠️ Внимание: после фильтрации не осталось документов!") return df except Exception as e: print(f"❌ Ошибка при загрузке данных: {e}") raise def needs_embedding_update(self, doc_id: str, text: str, existing_docs: Dict) -> bool: """Проверяет, нужно ли обновлять эмбеддинг для документа""" if doc_id not in existing_docs: return True current_hash = self._compute_text_hash(text) stored_hash = existing_docs[doc_id]["text_hash"] stored_model = existing_docs[doc_id].get("model_version", "") # Обновляем если изменился текст или модель return current_hash != stored_hash or stored_model != self.model_name def compute_embeddings_batch(self, texts: List[str]) -> np.ndarray: """Вычисляет эмбеддинги для батча текстов""" with torch.no_grad(): #embeddings = [] # n_docs = 20 # for idx in range(0, len(texts), n_docs): # documents_subset = texts[idx : idx + n_docs] inputs = self.tokenizer(texts, padding=True, truncation=True, return_tensors="pt", return_token_type_ids=False, max_length=512) outputs = self.model(**inputs) # embeddings.append( outputs.last_hidden_state[:, 0, :].detach().numpy() ) # # embeddings = np.concatenate(embeddings) embeddings = outputs.last_hidden_state[:, 0, :].detach().numpy() return embeddings def update_embeddings(self, df: pd.DataFrame, force_update: bool = False) -> int: """Обновляет эмбеддинги в ChromaDB, возвращает количество обновленных документов""" print("🔄 Проверка необходимости обновления эмбеддингов...") existing_docs = self._get_existing_documents() documents_to_update = [] for _, row in df.iterrows(): doc_id = str(row[DB_CONFIG[DB_CONFIG["table_name"]]['id_column']]) text = row[ DB_CONFIG['text_column' ] ] if force_update or self.needs_embedding_update(doc_id, text, existing_docs): documents_to_update.append((doc_id, text, row)) if not documents_to_update: print("✅ Все эмбеддинги актуальны") return 0 print(f"🔄 Вычисление эмбеддингов для {len(documents_to_update)} документов...") # Обрабатываем батчами updated_count = 0 for i in tqdm(range(0, len(documents_to_update), self.batch_size)): batch = documents_to_update[i:i + self.batch_size] batch_ids, batch_texts, batch_rows = zip(*batch) # Вычисляем эмбеддинги batch_embeddings = self.compute_embeddings_batch(batch_texts) # Подготавливаем метаданные metadatas = [] for row in batch_rows: metadata = { "title": row[DB_CONFIG["Articles"]['title_column']], "doi": row[DB_CONFIG["Articles"]['doi_column']] or "", "date": row[DB_CONFIG["Articles"]['date_column']] or "", "text_hash": self._compute_text_hash(row[DB_CONFIG['text_column']]), "model_version": self.model_name } metadatas.append(metadata) # Добавляем в ChromaDB self.collection.upsert( ids=list(batch_ids), embeddings=batch_embeddings.tolist(), metadatas=metadatas, documents=list(batch_texts) ) updated_count += len(batch) # Периодическая очистка памяти if i % 100 == 0: gc.collect() if torch.cuda.is_available(): torch.cuda.empty_cache() print(f"✅ Обновлено {updated_count} эмбеддингов") return updated_count def get_all_embeddings(self, df: pd.DataFrame) -> Tuple[np.ndarray, List[str]]: """Возвращает все эмбеддинги и соответствующие ID документов""" print("📥 Загрузка эмбеддингов из ChromaDB...") # Проверяем, есть ли документы в коллекции if self.collection.count() == 0: print("❌ Коллекция ChromaDB пуста") return np.array([]), [] doc_ids = [str(id) for id in df[DB_CONFIG['Articles']['id_column']].tolist()] try: # Получаем эмбеддинги из ChromaDB results = self.collection.get( ids=doc_ids, include=["embeddings", "metadatas"] ) # Сортируем в порядке исходного DataFrame id_to_embedding = {id: emb for id, emb in zip(results["ids"], results["embeddings"])} id_to_metadata = {id: meta for id, meta in zip(results["ids"], results["metadatas"])} embeddings = [] valid_ids = [] for doc_id in doc_ids: if doc_id in id_to_embedding: embeddings.append(id_to_embedding[doc_id]) valid_ids.append(doc_id) else: print(f"⚠️ Эмбеддинг для документа {doc_id} не найден") if not embeddings: print("❌ Не найдено ни одного эмбеддинга") return np.array([]), [] embeddings_array = np.array(embeddings) print(f"✅ Загружено {len(embeddings_array)} эмбеддингов") return embeddings_array, valid_ids except Exception as e: print(f"❌ Ошибка при загрузке эмбеддингов: {e}") return np.array([]), [] def get_collection_stats(self) -> Dict: """Возвращает статистику коллекции""" try: count = self.collection.count() return { "total_documents": count, "model": self.model_name, "embedding_dimension": self.embedding_dim } except Exception as e: print(f"⚠️ Ошибка при получении статистики: {e}") return {"error": str(e)} def cleanup(self): """Очищает ресурсы""" if hasattr(self, 'model'): del self.model gc.collect() if torch.cuda.is_available(): torch.cuda.empty_cache()