LightRAGDemo / app.py
Soha85's picture
Update app.py
9a09168 verified
Raw History Blame Contribute Delete
7.89 kB
import streamlit as st
import os
import tempfile
import asyncio
import nest_asyncio
import shutil
import torch
import gc
import networkx as nx
from pyvis.network import Network
import streamlit.components.v1 as components
from sentence_transformers import SentenceTransformer
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
from lightrag import LightRAG, QueryParam
from lightrag.utils import EmbeddingFunc
from langchain_community.document_loaders import PyPDFLoader
# Allow nested loops for asyncio
nest_asyncio.apply()
# --- Page Config ---
st.set_page_config(page_title="LightRAG Case Study Guide", layout="wide")
st.title("๐Ÿ“š LightRAG Case Study Guide for Students")
st.markdown("""
This application implements the **LightRAG** framework using the `lightrag-hku` library.
It uses a dual-level retrieval system (Low-level & High-level) to maintain relational context.
""")
# --- Sidebar Parameters ---
with st.sidebar:
st.header("RAG Parameters")
chunk_size = st.slider("Chunk Token Size", 100, 1000, 512)
chunk_overlap = st.slider("Chunk Overlap Token Size", 0, 200, 50)
query_mode = st.selectbox("Query Mode", ["hybrid", "naive", "local", "global"], index=0)
st.divider()
st.info("Model: TinyLlama/TinyLlama-1.1B-Chat-v1.0(4-bit Quantized)")
st.success("Optimized for 16GB Memory Limit")
# --- Model Loading ---
@st.cache_resource
def load_models():
# Clear memory before loading
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
st.info("๐Ÿง  Loading Embedding Model (CPU)...")
embedding_model = SentenceTransformer("all-MiniLM-L6-v2", device="cpu")
st.info("โณ Loading LLM (Qwen2.5-3B-Instruct)...")
# Using Qwen2.5-3B-Instruct as a balance between performance and memory
model_name = "TinyLlama/TinyLlama-1.1B-Chat-v1.0"
device = "cuda" if torch.cuda.is_available() else "cpu"
tokenizer = AutoTokenizer.from_pretrained(model_name)
# Always use 4-bit quantization to stay within 16GB, even on CPU if possible
# Note: bitsandbytes 4-bit is primarily for CUDA. For CPU, we'll use float32 with low_cpu_mem_usage.
if device == "cuda":
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_compute_dtype=torch.float16,
bnb_4bit_quant_type="nf4",
bnb_4bit_use_double_quant=True
)
model = AutoModelForCausalLM.from_pretrained(
model_name,
quantization_config=bnb_config,
device_map="auto",
)
else:
# On CPU, we use float32 but ensure low_cpu_mem_usage. 3B model in float32 is ~6GB.
model = AutoModelForCausalLM.from_pretrained(
model_name,
torch_dtype=torch.float32,
low_cpu_mem_usage=True,
).to("cpu")
return embedding_model, tokenizer, model
embedding_model, tokenizer, model = load_models()
# --- LightRAG Functions ---
async def embed_func(texts):
return embedding_model.encode(texts, convert_to_numpy=True)
def sync_llm_call(prompt, kwargs):
system_prompt = kwargs.get("system_prompt") or "You are a helpful assistant."
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": prompt},
]
text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
inputs = tokenizer(text, return_tensors="pt").to(model.device)
if inputs["input_ids"].shape[1] > 1500:
inputs["input_ids"] = inputs["input_ids"][:, :1500]
inputs["attention_mask"] = inputs["attention_mask"][:, :1500]
with torch.no_grad():
output_ids = model.generate(
**inputs,
max_new_tokens=512,
do_sample=False,
temperature=None,
top_p=None,
)
new_tokens = output_ids[0][inputs["input_ids"].shape[1]:]
return tokenizer.decode(new_tokens, skip_special_tokens=True).strip()
async def llm_func(prompt, **kwargs):
safe_kwargs = {k: v for k, v in kwargs.items() if isinstance(v, (str, int, float, bool, type(None)))}
loop = asyncio.get_event_loop()
return await loop.run_in_executor(None, sync_llm_call, prompt, safe_kwargs)
# --- App Logic ---
WORKING_DIR = "./lightrag_storage"
if not os.path.exists(WORKING_DIR):
os.makedirs(WORKING_DIR)
# Initialize RAG instance
rag = LightRAG(
working_dir=WORKING_DIR,
embedding_func=EmbeddingFunc(
embedding_dim=384,
max_token_size=256,
func=embed_func,
),
llm_model_func=llm_func,
chunk_token_size=chunk_size,
chunk_overlap_token_size=chunk_overlap,
)
uploaded_file = st.file_uploader("Upload a PDF document", type="pdf")
if uploaded_file:
# Use a persistent temp file path to avoid unlinking issues across button clicks
tmp_path = os.path.join(tempfile.gettempdir(), "uploaded_doc.pdf")
with open(tmp_path, "wb") as f:
f.write(uploaded_file.getvalue())
if st.button("Index Document"):
with st.status("Processing...", expanded=True) as status:
st.write("๐Ÿ“„ Loading PDF...")
loader = PyPDFLoader(tmp_path)
pages = loader.load()
documents = [p.page_content for p in pages]
st.write(f"โœ‚๏ธ Indexing {len(documents)} pages into LightRAG...")
async def run_indexing():
await rag.initialize_storages()
await rag.ainsert(documents)
# Clear memory after indexing
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
asyncio.run(run_indexing())
st.write("โœ… Indexing complete!")
status.update(label="Ready for Questions!", state="complete", expanded=False)
st.session_state['indexed'] = True
if st.session_state.get('indexed'):
tab1, tab2 = st.tabs(["๐Ÿ’ฌ Q&A Interface", "๐Ÿ•ธ๏ธ Knowledge Graph"])
with tab1:
query = st.text_input("Ask a question:")
if query:
with st.spinner("Generating answer..."):
async def run_query():
return await rag.aquery(query, param=QueryParam(mode=query_mode))
answer = asyncio.run(run_query())
st.markdown("### ๐Ÿค– Answer")
st.write(answer)
with tab2:
st.markdown("### Knowledge Graph")
graph_path = os.path.join(WORKING_DIR, "graph_chunk_entity_relation.graphml")
if os.path.exists(graph_path):
try:
G = nx.read_graphml(graph_path)
if len(G.nodes) > 100:
st.warning("Graph is large. Showing a subgraph of the first 100 nodes.")
nodes = list(G.nodes)[:100]
G = G.subgraph(nodes)
net = Network(height="500px", width="100%", bgcolor="#ffffff", font_color="black")
net.from_nx(G)
with tempfile.NamedTemporaryFile(delete=False, suffix=".html") as tmp_html:
net.save_graph(tmp_html.name)
with open(tmp_html.name, 'r', encoding='utf-8') as f:
components.html(f.read(), height=550)
os.unlink(tmp_html.name)
except Exception as e:
st.error(f"Error loading graph: {e}")
else:
st.info("Knowledge graph file not found. It will be generated during indexing.")
else:
st.info("Please upload a PDF to start.")
if 'indexed' in st.session_state:
st.session_state['indexed'] = False