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