tahamajs's picture
download
raw
18.7 kB
"""
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
@dataclass
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"""
@self.app.post("/query")
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))
@self.app.get("/stats")
async def get_stats():
"""Get RAG system statistics"""
return self.rag.get_statistics()
@self.app.post("/visualize")
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.