""" Pipeline d'ingestion optimisé — Green AI Edition ================================================= Aligne sur app.py : mêmes embeddings locaux, même chunking parent-child, même collection ChromaDB (rag_documents_optimized). Optimisations : - sentence-transformers local (zéro appel API Azure) - PyMuPDF get_text("markdown") pour préserver la structure des PDFs - Parent-child chunking (child 380 chars pour le retrieval, parent 1024 pour le LLM) - Extraction parallèle (ThreadPoolExecutor, I/O-bound) - Embedding par lots de 128 (CPU/MPS) - Filtre anti-doublon sur les sources déjà indexées """ import os import sys import logging from pathlib import Path from concurrent.futures import ThreadPoolExecutor, as_completed os.environ.setdefault("ANONYMIZED_TELEMETRY", "False") import fitz import chromadb from chromadb.config import Settings from langchain_text_splitters import RecursiveCharacterTextSplitter from sentence_transformers import SentenceTransformer logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s") logger = logging.getLogger(__name__) # --------------------------------------------------------------------------- # Configuration # --------------------------------------------------------------------------- PROJECT_ROOT = Path(__file__).parent DATA_DIR = Path("/data") if Path("/data").is_dir() else PROJECT_ROOT / "data" TRAIN_DATA_DIR = PROJECT_ROOT / "train_data" CHROMA_PERSIST_DIR = str(DATA_DIR / "chroma_db_optimized") COLLECTION_NAME = "rag_documents_optimized" CHUNK_SIZE = 380 CHUNK_OVERLAP = 80 PARENT_CHUNK_SIZE = 1024 PARENT_CHUNK_OVERLAP = 100 EMBEDDING_BATCH_SIZE = 128 # taille de lot pour l'encodage local CPU/MPS MAX_WORKERS = 8 # threads pour l'extraction parallèle des fichiers INSERT_BATCH = 500 # taille de lot pour l'insertion ChromaDB _SEPARATORS = ["---", "\n## ", "\n### ", "\n\n", "\n", ". ", " ", ""] # --------------------------------------------------------------------------- # Modèle local (chargé une fois) # --------------------------------------------------------------------------- logger.info("Chargement du modèle d'embedding (paraphrase-multilingual-MiniLM-L12-v2)...") _embedding_model = SentenceTransformer("paraphrase-multilingual-MiniLM-L12-v2") logger.info(f"Modèle prêt. Dimension : {_embedding_model.get_embedding_dimension()}") # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- def extract_text_from_pdf(pdf_path: Path) -> str: doc = fitz.open(str(pdf_path)) pages_text = [] for page in doc: try: text = page.get_text("markdown") except Exception: text = page.get_text("text") if text.strip(): pages_text.append(text.strip()) doc.close() return "\n\n---\n\n".join(pages_text) def chunk_text(text: str, source: str) -> list[dict]: """Parent-child chunking — même logique que app.py.""" parent_splitter = RecursiveCharacterTextSplitter( chunk_size=PARENT_CHUNK_SIZE, chunk_overlap=PARENT_CHUNK_OVERLAP, separators=_SEPARATORS, ) child_splitter = RecursiveCharacterTextSplitter( chunk_size=CHUNK_SIZE, chunk_overlap=CHUNK_OVERLAP, separators=_SEPARATORS, ) result = [] for parent_text in parent_splitter.split_text(text): for child_text in child_splitter.split_text(parent_text): result.append({ "text": child_text, "parent_text": parent_text, "source": source, "chunk_index": len(result), }) return result def generate_embeddings(texts: list[str]) -> list[list[float]]: """Embedding local L2-normalisé — zéro appel API.""" return _embedding_model.encode( texts, normalize_embeddings=True, show_progress_bar=False ).tolist() def process_single_file(file_path: Path) -> list[dict]: source = str(file_path.relative_to(TRAIN_DATA_DIR)) try: if file_path.suffix.lower() == ".pdf": text = extract_text_from_pdf(file_path) else: text = file_path.read_text(encoding="utf-8") if not text.strip(): return [] return chunk_text(text, source=source) except Exception as e: logger.warning(f"[Erreur] {source} : {e}") return [] # --------------------------------------------------------------------------- # Pipeline principal # --------------------------------------------------------------------------- def main(): if not TRAIN_DATA_DIR.exists(): logger.error(f"Dossier introuvable : {TRAIN_DATA_DIR}") sys.exit(1) logger.info("=" * 60) logger.info("PIPELINE D'INGESTION OPTIMISÉ (embeddings locaux)") logger.info("=" * 60) chroma_client = chromadb.PersistentClient( path=CHROMA_PERSIST_DIR, settings=Settings(anonymized_telemetry=False), ) collection = chroma_client.get_or_create_collection( name=COLLECTION_NAME, metadata={"hnsw:space": "cosine"}, ) # Détection de mismatch de dimensions (ex : ancienne DB Azure 1536 dims) expected_dim = _embedding_model.get_embedding_dimension() if collection.count() > 0: try: sample = collection.get(limit=1, include=["embeddings"]) existing_dim = len(sample["embeddings"][0]) if existing_dim != expected_dim: logger.warning( f"Dimension mismatch ({existing_dim} → {expected_dim}). " "Recréation de la collection." ) chroma_client.delete_collection(COLLECTION_NAME) collection = chroma_client.create_collection( name=COLLECTION_NAME, metadata={"hnsw:space": "cosine"} ) except Exception as e: logger.warning(f"Vérification dimensions impossible : {e}") # Étape 0 : Filtre anti-doublon logger.info("Étape 0 : Vérification des sources déjà indexées...") existing_sources: set[str] = set() existing_data = collection.get(include=["metadatas"]) if existing_data and existing_data["metadatas"]: for meta in existing_data["metadatas"]: if meta and "source" in meta: existing_sources.add(meta["source"]) logger.info(f"→ {len(existing_sources)} fichiers déjà présents en base.") all_files = sorted( list(TRAIN_DATA_DIR.rglob("*.pdf")) + list(TRAIN_DATA_DIR.rglob("*.txt")) ) files_to_process = [ f for f in all_files if str(f.relative_to(TRAIN_DATA_DIR)) not in existing_sources ] logger.info(f"{len(all_files)} fichiers au total, {len(files_to_process)} nouveaux.") if not files_to_process: logger.info("Base déjà à jour. Rien à faire.") return # Étape 1 : Extraction et chunking parallèle (I/O-bound) logger.info(f"Étape 1 : Extraction parallèle ({MAX_WORKERS} threads)...") all_chunks: list[dict] = [] with ThreadPoolExecutor(max_workers=MAX_WORKERS) as executor: futures = {executor.submit(process_single_file, f): f for f in files_to_process} for idx, future in enumerate(as_completed(futures), start=1): all_chunks.extend(future.result()) if idx % 10 == 0 or idx == len(files_to_process): logger.info(f" [{idx}/{len(files_to_process)}] — {len(all_chunks)} chunks") if not all_chunks: logger.error("Aucun texte extrait. Arrêt.") return logger.info(f"→ {len(all_chunks)} chunks générés.") # Étape 2 : Embedding local par lots logger.info(f"Étape 2 : Embedding local par lots de {EMBEDDING_BATCH_SIZE}...") batches = [ all_chunks[i:i + EMBEDDING_BATCH_SIZE] for i in range(0, len(all_chunks), EMBEDDING_BATCH_SIZE) ] all_embeddings: list[list[float]] = [] for idx, batch in enumerate(batches, start=1): all_embeddings.extend(generate_embeddings([c["text"] for c in batch])) if idx % 10 == 0 or idx == len(batches): logger.info(f" [{idx}/{len(batches)}] — {len(all_embeddings)} embeddings") logger.info(f"→ {len(all_embeddings)} embeddings générés.") # Étape 3 : Insertion ChromaDB par lots logger.info(f"Étape 3 : Insertion dans ChromaDB (lots de {INSERT_BATCH})...") base = collection.count() ids = [f"doc_{base + i}" for i in range(len(all_chunks))] texts = [c["text"] for c in all_chunks] metadatas = [ { "source": c["source"], "chunk_index": c["chunk_index"], "parent_text": c.get("parent_text", c["text"]), } for c in all_chunks ] for i in range(0, len(all_chunks), INSERT_BATCH): collection.add( ids=ids[i:i + INSERT_BATCH], embeddings=all_embeddings[i:i + INSERT_BATCH], documents=texts[i:i + INSERT_BATCH], metadatas=metadatas[i:i + INSERT_BATCH], ) logger.info("=" * 60) logger.info("INGESTION TERMINÉE") logger.info(f" DB : {CHROMA_PERSIST_DIR}") logger.info(f" Collection : {COLLECTION_NAME}") logger.info(f" Chunks ajoutés : {len(all_chunks)}") logger.info(f" Total en DB : {collection.count()}") logger.info("=" * 60) if __name__ == "__main__": main()