Buckets:
| """ | |
| RAG (Retrieval-Augmented Generation) System for Knowledge Graphs | |
| This module implements advanced RAG capabilities for knowledge graph reasoning | |
| """ | |
| import os | |
| import json | |
| import asyncio | |
| from typing import Dict, List, Any, Optional, Tuple | |
| from dataclasses import dataclass | |
| import numpy as np | |
| from datetime import datetime | |
| # Vector databases and embeddings | |
| import chromadb | |
| from chromadb.config import Settings | |
| from sentence_transformers import SentenceTransformer | |
| import faiss | |
| # LangChain components | |
| from langchain.schema import Document | |
| from langchain_community.vectorstores import Chroma, FAISS | |
| from langchain_community.embeddings import HuggingFaceEmbeddings | |
| from langchain.text_splitter import RecursiveCharacterTextSplitter | |
| from langchain.prompts import PromptTemplate | |
| from langchain.chains import RetrievalQA | |
| from langchain_google_genai import ChatGoogleGenerativeAI | |
| # Graph processing | |
| import networkx as nx | |
| from pyvis.network import Network | |
| # FastAPI | |
| from fastapi import FastAPI, HTTPException | |
| from pydantic import BaseModel | |
| import sys | |
| import os | |
| sys.path.append(os.path.join(os.path.dirname(__file__), "..")) | |
| from core.knowledge_graph import KnowledgeGraph, Triple | |
| from core.scientific_kg import AdvancedScientificKG | |
| class RAGResult: | |
| """Result from RAG system""" | |
| query: str | |
| answer: str | |
| retrieved_documents: List[Document] | |
| confidence: float | |
| reasoning_path: List[str] | |
| entities_found: List[str] | |
| relations_found: List[str] | |
| class KnowledgeGraphRAG: | |
| """RAG system specifically designed for Knowledge Graphs""" | |
| def __init__( | |
| self, | |
| kg: KnowledgeGraph, | |
| embedding_model: str = "sentence-transformers/all-MiniLM-L6-v2", | |
| ): | |
| self.kg = kg | |
| self.embedding_model_name = embedding_model | |
| # Initialize embeddings | |
| self.embeddings = HuggingFaceEmbeddings( | |
| model_name=embedding_model, model_kwargs={"device": "cpu"} | |
| ) | |
| # Initialize vector stores | |
| self.triple_store = None | |
| self.entity_store = None | |
| self.relation_store = None | |
| # Initialize LLM with Gemini API | |
| try: | |
| import os | |
| from dotenv import load_dotenv | |
| load_dotenv() | |
| google_api_key = os.getenv("GOOGLE_API_KEY") | |
| if google_api_key: | |
| self.llm = ChatGoogleGenerativeAI( | |
| model="gemini-pro", temperature=0.1, google_api_key=google_api_key | |
| ) | |
| print(f"✅ Gemini LLM initialized for RAG") | |
| else: | |
| raise ValueError("GOOGLE_API_KEY not found in environment") | |
| except Exception as e: | |
| print(f"⚠️ LLM initialization failed: {e}") | |
| print("💡 Using mock LLM for demonstration") | |
| self.llm = None | |
| # Setup vector stores | |
| self._setup_vector_stores() | |
| # Create RAG chains | |
| self._setup_rag_chains() | |
| def _setup_vector_stores(self): | |
| """Setup vector stores for different KG components""" | |
| # 1. Triple store - for semantic search over facts | |
| triple_docs = [] | |
| triple_metadatas = [] | |
| for i, triple in enumerate(self.kg.triples): | |
| # Create multiple representations of each triple | |
| doc_texts = [ | |
| f"{triple.subject} {triple.predicate} {triple.object}", | |
| f"The {triple.subject} is related to {triple.object} through {triple.predicate}", | |
| f"{triple.subject} has relationship {triple.predicate} with {triple.object}", | |
| f"Fact: {triple.subject} -> {triple.predicate} -> {triple.object}", | |
| ] | |
| for doc_text in doc_texts: | |
| triple_docs.append(doc_text) | |
| triple_metadatas.append( | |
| { | |
| "triple_id": i, | |
| "subject": triple.subject, | |
| "predicate": triple.predicate, | |
| "object": triple.object, | |
| "type": "triple", | |
| } | |
| ) | |
| self.triple_store = Chroma.from_texts( | |
| texts=triple_docs, | |
| embedding=self.embeddings, | |
| metadatas=triple_metadatas, | |
| collection_name="kg_triples", | |
| ) | |
| # 2. Entity store - for entity-centric search | |
| entity_docs = [] | |
| entity_metadatas = [] | |
| for entity_id, entity in self.kg.entities.items(): | |
| # Create entity descriptions | |
| entity_text = f"Entity: {entity.name} (Type: {entity.entity_type})" | |
| if entity.attributes: | |
| attr_text = ", ".join( | |
| [f"{k}: {v}" for k, v in entity.attributes.items()] | |
| ) | |
| entity_text += f" Attributes: {attr_text}" | |
| entity_docs.append(entity_text) | |
| entity_metadatas.append( | |
| { | |
| "entity_id": entity_id, | |
| "entity_name": entity.name, | |
| "entity_type": entity.entity_type, | |
| "type": "entity", | |
| } | |
| ) | |
| self.entity_store = Chroma.from_texts( | |
| texts=entity_docs, | |
| embedding=self.embeddings, | |
| metadatas=entity_metadatas, | |
| collection_name="kg_entities", | |
| ) | |
| # 3. Relation store - for relationship patterns | |
| relation_docs = [] | |
| relation_metadatas = [] | |
| for rel_id, relation in self.kg.relations.items(): | |
| rel_text = f"Relation: {relation.name} (Domain: {relation.domain}, Range: {relation.range})" | |
| relation_docs.append(rel_text) | |
| relation_metadatas.append( | |
| { | |
| "relation_id": rel_id, | |
| "relation_name": relation.name, | |
| "domain": relation.domain, | |
| "range": relation.range, | |
| "type": "relation", | |
| } | |
| ) | |
| self.relation_store = Chroma.from_texts( | |
| texts=relation_docs, | |
| embedding=self.embeddings, | |
| metadatas=relation_metadatas, | |
| collection_name="kg_relations", | |
| ) | |
| def _setup_rag_chains(self): | |
| """Setup RAG chains for different types of queries""" | |
| # Template for triple-based queries | |
| triple_template = """You are a knowledge graph expert. Use the following context to answer the question. | |
| Context (Knowledge Graph Facts): | |
| {context} | |
| Question: {question} | |
| Answer based on the knowledge graph facts. If the information is not available, say so clearly. | |
| """ | |
| self.triple_prompt = PromptTemplate( | |
| template=triple_template, input_variables=["context", "question"] | |
| ) | |
| # Template for entity-based queries | |
| entity_template = """You are analyzing entities in a knowledge graph. Use the following entity information to answer the question. | |
| Entity Information: | |
| {context} | |
| Question: {question} | |
| Provide a detailed answer about the entities mentioned. | |
| """ | |
| self.entity_prompt = PromptTemplate( | |
| template=entity_template, input_variables=["context", "question"] | |
| ) | |
| # Template for relation-based queries | |
| relation_template = """You are analyzing relationships in a knowledge graph. Use the following relationship information to answer the question. | |
| Relationship Information: | |
| {context} | |
| Question: {question} | |
| Explain the relationships and their patterns. | |
| """ | |
| self.relation_prompt = PromptTemplate( | |
| template=relation_template, input_variables=["context", "question"] | |
| ) | |
| def _classify_query_type(self, query: str) -> str: | |
| """Classify the type of query""" | |
| query_lower = query.lower() | |
| if any(word in query_lower for word in ["what", "who", "when", "where", "how"]): | |
| return "factual" | |
| elif any( | |
| word in query_lower | |
| for word in ["relationship", "related", "connected", "link"] | |
| ): | |
| return "relational" | |
| elif any( | |
| word in query_lower for word in ["entity", "person", "concept", "thing"] | |
| ): | |
| return "entity" | |
| else: | |
| return "general" | |
| def _retrieve_relevant_documents( | |
| self, query: str, query_type: str, k: int = 5 | |
| ) -> List[Document]: | |
| """Retrieve relevant documents based on query type""" | |
| all_docs = [] | |
| if query_type in ["factual", "general"]: | |
| # Search in triple store | |
| triple_docs = self.triple_store.similarity_search(query, k=k) | |
| all_docs.extend(triple_docs) | |
| if query_type in ["entity", "general"]: | |
| # Search in entity store | |
| entity_docs = self.entity_store.similarity_search(query, k=k) | |
| all_docs.extend(entity_docs) | |
| if query_type in ["relational", "general"]: | |
| # Search in relation store | |
| relation_docs = self.relation_store.similarity_search(query, k=k) | |
| all_docs.extend(relation_docs) | |
| # Remove duplicates and return top k | |
| unique_docs = [] | |
| seen_content = set() | |
| for doc in all_docs: | |
| if doc.page_content not in seen_content: | |
| unique_docs.append(doc) | |
| seen_content.add(doc.page_content) | |
| return unique_docs[:k] | |
| def _extract_entities_from_documents(self, docs: List[Document]) -> List[str]: | |
| """Extract entities mentioned in retrieved documents""" | |
| entities = set() | |
| for doc in docs: | |
| metadata = doc.metadata | |
| if "subject" in metadata: | |
| entities.add(metadata["subject"]) | |
| if "object" in metadata: | |
| entities.add(metadata["object"]) | |
| if "entity_id" in metadata: | |
| entities.add(metadata["entity_id"]) | |
| return list(entities) | |
| def _extract_relations_from_documents(self, docs: List[Document]) -> List[str]: | |
| """Extract relations mentioned in retrieved documents""" | |
| relations = set() | |
| for doc in docs: | |
| metadata = doc.metadata | |
| if "predicate" in metadata: | |
| relations.add(metadata["predicate"]) | |
| if "relation_id" in metadata: | |
| relations.add(metadata["relation_id"]) | |
| return list(relations) | |
| def _generate_reasoning_path(self, query: str, docs: List[Document]) -> List[str]: | |
| """Generate reasoning path from retrieved documents""" | |
| reasoning_steps = [] | |
| reasoning_steps.append(f"Query: {query}") | |
| reasoning_steps.append(f"Retrieved {len(docs)} relevant documents") | |
| for i, doc in enumerate(docs[:3]): # Show top 3 documents | |
| metadata = doc.metadata | |
| if metadata.get("type") == "triple": | |
| reasoning_steps.append( | |
| f"Step {i+1}: Found fact - {metadata['subject']} {metadata['predicate']} {metadata['object']}" | |
| ) | |
| elif metadata.get("type") == "entity": | |
| reasoning_steps.append( | |
| f"Step {i+1}: Found entity - {metadata['entity_name']} ({metadata['entity_type']})" | |
| ) | |
| elif metadata.get("type") == "relation": | |
| reasoning_steps.append( | |
| f"Step {i+1}: Found relation - {metadata['relation_name']}" | |
| ) | |
| return reasoning_steps | |
| async def query(self, question: str, k: int = 5) -> RAGResult: | |
| """Main RAG query method""" | |
| # Classify query type | |
| query_type = self._classify_query_type(question) | |
| # Retrieve relevant documents | |
| retrieved_docs = self._retrieve_relevant_documents(question, query_type, k) | |
| if not retrieved_docs: | |
| return RAGResult( | |
| query=question, | |
| answer="No relevant information found in the knowledge graph.", | |
| retrieved_documents=[], | |
| confidence=0.0, | |
| reasoning_path=["No documents retrieved"], | |
| entities_found=[], | |
| relations_found=[], | |
| ) | |
| # Extract entities and relations | |
| entities_found = self._extract_entities_from_documents(retrieved_docs) | |
| relations_found = self._extract_relations_from_documents(retrieved_docs) | |
| # Generate reasoning path | |
| reasoning_path = self._generate_reasoning_path(question, retrieved_docs) | |
| # Prepare context for LLM | |
| context = "\n".join([doc.page_content for doc in retrieved_docs]) | |
| # Choose appropriate prompt based on query type | |
| if query_type == "entity": | |
| prompt = self.entity_prompt | |
| elif query_type == "relational": | |
| prompt = self.relation_prompt | |
| else: | |
| prompt = self.triple_prompt | |
| # Generate answer using LLM | |
| if self.llm is None: | |
| answer = f"Mock RAG answer for: {question}" | |
| else: | |
| formatted_prompt = prompt.format(context=context, question=question) | |
| response = await self.llm.ainvoke(formatted_prompt) | |
| answer = response.content | |
| # Calculate confidence based on document relevance | |
| confidence = min(0.9, len(retrieved_docs) / k) | |
| return RAGResult( | |
| query=question, | |
| answer=answer, | |
| retrieved_documents=retrieved_docs, | |
| confidence=confidence, | |
| reasoning_path=reasoning_path, | |
| entities_found=entities_found, | |
| relations_found=relations_found, | |
| ) | |
| def visualize_retrieval(self, result: RAGResult, save_path: str = None): | |
| """Visualize the retrieval process""" | |
| # Create network visualization | |
| net = Network( | |
| height="600px", width="100%", bgcolor="#222222", font_color="white" | |
| ) | |
| # Add query node | |
| net.add_node("query", label=f"Query: {result.query}", color="#ff6b6b", size=30) | |
| # Add retrieved documents | |
| for i, doc in enumerate(result.retrieved_documents): | |
| doc_id = f"doc_{i}" | |
| metadata = doc.metadata | |
| if metadata.get("type") == "triple": | |
| label = f"{metadata['subject']}\n{metadata['predicate']}\n{metadata['object']}" | |
| color = "#4ecdc4" | |
| elif metadata.get("type") == "entity": | |
| label = f"Entity: {metadata['entity_name']}" | |
| color = "#45b7d1" | |
| elif metadata.get("type") == "relation": | |
| label = f"Relation: {metadata['relation_name']}" | |
| color = "#96ceb4" | |
| else: | |
| label = f"Document {i}" | |
| color = "#feca57" | |
| net.add_node(doc_id, label=label, color=color, size=20) | |
| net.add_edge("query", doc_id, color="#ffffff") | |
| # Add answer node | |
| net.add_node( | |
| "answer", | |
| label=f"Answer: {result.answer[:100]}...", | |
| color="#ff9ff3", | |
| size=25, | |
| ) | |
| # Connect documents to answer | |
| for i in range(len(result.retrieved_documents)): | |
| net.add_edge(f"doc_{i}", "answer", color="#ffffff", dashes=True) | |
| # Save or show | |
| if save_path: | |
| net.save_graph(save_path) | |
| else: | |
| net.show("rag_visualization.html") | |
| def get_statistics(self) -> Dict[str, Any]: | |
| """Get RAG system statistics""" | |
| return { | |
| "total_triples": len(self.kg.triples), | |
| "total_entities": len(self.kg.entities), | |
| "total_relations": len(self.kg.relations), | |
| "triple_store_size": ( | |
| len(self.triple_store._collection.get()["ids"]) | |
| if self.triple_store | |
| else 0 | |
| ), | |
| "entity_store_size": ( | |
| len(self.entity_store._collection.get()["ids"]) | |
| if self.entity_store | |
| else 0 | |
| ), | |
| "relation_store_size": ( | |
| len(self.relation_store._collection.get()["ids"]) | |
| if self.relation_store | |
| else 0 | |
| ), | |
| "embedding_model": self.embedding_model_name, | |
| } | |
| class RAGAPI: | |
| """FastAPI interface for RAG system""" | |
| def __init__(self, rag_system: KnowledgeGraphRAG): | |
| self.rag = rag_system | |
| self.app = FastAPI(title="Knowledge Graph RAG API") | |
| self._setup_routes() | |
| def _setup_routes(self): | |
| """Setup API routes""" | |
| async def query_rag(question: str, k: int = 5): | |
| """Query the RAG system""" | |
| try: | |
| result = await self.rag.query(question, k) | |
| return { | |
| "query": result.query, | |
| "answer": result.answer, | |
| "confidence": result.confidence, | |
| "reasoning_path": result.reasoning_path, | |
| "entities_found": result.entities_found, | |
| "relations_found": result.relations_found, | |
| "num_documents": len(result.retrieved_documents), | |
| } | |
| except Exception as e: | |
| raise HTTPException(status_code=500, detail=str(e)) | |
| async def get_stats(): | |
| """Get RAG system statistics""" | |
| return self.rag.get_statistics() | |
| async def visualize_query(question: str, save_path: str = None): | |
| """Visualize a query""" | |
| try: | |
| result = await self.rag.query(question) | |
| if save_path: | |
| self.rag.visualize_retrieval(result, save_path) | |
| return {"message": f"Visualization saved to {save_path}"} | |
| else: | |
| self.rag.visualize_retrieval(result) | |
| return {"message": "Visualization opened in browser"} | |
| except Exception as e: | |
| raise HTTPException(status_code=500, detail=str(e)) | |
| def run(self, host: str = "0.0.0.0", port: int = 8001): | |
| """Run the RAG API server""" | |
| import uvicorn | |
| uvicorn.run(self.app, host=host, port=port) | |
| def create_rag_system(kg: KnowledgeGraph) -> Tuple[KnowledgeGraphRAG, RAGAPI]: | |
| """Create RAG system for Knowledge Graph""" | |
| # Create RAG system | |
| rag_system = KnowledgeGraphRAG(kg) | |
| # Create API interface | |
| rag_api = RAGAPI(rag_system) | |
| return rag_system, rag_api | |
Xet Storage Details
- Size:
- 18.7 kB
- Xet hash:
- 20db2c5556b298e8c9c1d692fa3948c08c5999501145f35e308e8d96565da55b
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.