Spaces:
Sleeping
Sleeping
Download app.py from Soha85/LightRAGDemo: direct link, hf CLI and curl.
- Browser
- Download file 7.89 kB
-
https://huggingface.co/spaces/Soha85/LightRAGDemo/resolve/main/app.py
- Command line
-
hf download hf://spaces/Soha85/LightRAGDemo/app.py
-
curl -L -o app.py https://huggingface.co/spaces/Soha85/LightRAGDemo/resolve/main/app.py
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 --- | |
| 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 | |