Spaces:
Sleeping
Sleeping
| import os | |
| import io | |
| import torch | |
| import gradio as gr | |
| from typing import List, Tuple, Dict | |
| from sentence_transformers import SentenceTransformer | |
| from transformers import AutoProcessor, Gemma3ForConditionalGeneration | |
| from pypdf import PdfReader | |
| # -------------------------------------------------------------------- | |
| # Konfiguration | |
| # -------------------------------------------------------------------- | |
| EMBED_MODEL_ID = "google/embeddinggemma-300m" | |
| LLM_MODEL_ID = "google/gemma-3-4b-it" | |
| # Globale States (simpler, reicht für einen Space) | |
| EMBED_MODEL: SentenceTransformer = None | |
| LLM_MODEL: Gemma3ForConditionalGeneration = None | |
| LLM_PROCESSOR: AutoProcessor = None | |
| DOCUMENT_STORE: List[Dict] = [] # {content, embedding (tensor), source} | |
| # -------------------------------------------------------------------- | |
| # Model Loading | |
| # -------------------------------------------------------------------- | |
| def get_device() -> torch.device: | |
| if torch.cuda.is_available(): | |
| return torch.device("cuda") | |
| return torch.device("cpu") | |
| def get_embedding_model() -> SentenceTransformer: | |
| global EMBED_MODEL | |
| if EMBED_MODEL is None: | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| # EmbeddingGemma nutzt SentenceTransformers-Wrapper | |
| EMBED_MODEL = SentenceTransformer(EMBED_MODEL_ID, device=device) | |
| return EMBED_MODEL | |
| def get_llm() -> Tuple[Gemma3ForConditionalGeneration, AutoProcessor]: | |
| global LLM_MODEL, LLM_PROCESSOR | |
| if LLM_MODEL is None or LLM_PROCESSOR is None: | |
| device = get_device() | |
| if device.type == "cuda": | |
| # bfloat16 für GPU wie im Model-Card-Beispiel | |
| LLM_MODEL = Gemma3ForConditionalGeneration.from_pretrained( | |
| LLM_MODEL_ID, | |
| torch_dtype=torch.bfloat16, | |
| device_map="auto", | |
| ).eval() | |
| else: | |
| # CPU: float32 | |
| LLM_MODEL = Gemma3ForConditionalGeneration.from_pretrained( | |
| LLM_MODEL_ID, | |
| torch_dtype=torch.float32, | |
| ).to(device).eval() | |
| LLM_PROCESSOR = AutoProcessor.from_pretrained(LLM_MODEL_ID) | |
| return LLM_MODEL, LLM_PROCESSOR | |
| # -------------------------------------------------------------------- | |
| # Datei-Handling & Chunking | |
| # -------------------------------------------------------------------- | |
| def extract_text_from_file(path: str) -> str: | |
| ext = os.path.splitext(path)[1].lower() | |
| if ext in [".txt", ".md", ".markdown"]: | |
| with open(path, "r", encoding="utf-8", errors="ignore") as f: | |
| return f.read() | |
| if ext == ".pdf": | |
| text = [] | |
| reader = PdfReader(path) | |
| for page in reader.pages: | |
| page_text = page.extract_text() | |
| if page_text: | |
| text.append(page_text) | |
| return "\n".join(text) | |
| # Fallback: Versuche als Text zu lesen | |
| try: | |
| with open(path, "r", encoding="utf-8", errors="ignore") as f: | |
| return f.read() | |
| except Exception: | |
| return "" | |
| def chunk_text( | |
| text: str, | |
| chunk_size: int = 800, # ca. "Token"-Näherung über Wörter | |
| chunk_overlap: int = 200, | |
| ) -> List[str]: | |
| words = text.split() | |
| if not words: | |
| return [] | |
| chunks = [] | |
| start = 0 | |
| while start < len(words): | |
| end = min(len(words), start + chunk_size) | |
| chunk = " ".join(words[start:end]).strip() | |
| if chunk: | |
| chunks.append(chunk) | |
| if end == len(words): | |
| break | |
| start = max(0, end - chunk_overlap) | |
| return chunks | |
| # -------------------------------------------------------------------- | |
| # Indexing / RAG | |
| # -------------------------------------------------------------------- | |
| def index_files(file_paths: List[str]) -> str: | |
| """ | |
| Liest hochgeladene Dateien ein, chunked sie und | |
| legt Embeddings in einem einfachen In-Memory-Store ab. | |
| """ | |
| if not file_paths: | |
| return "Keine Dateien übergeben." | |
| embed_model = get_embedding_model() | |
| added_chunks = 0 | |
| for path in file_paths: | |
| if path is None: | |
| continue | |
| text = extract_text_from_file(path) | |
| if not text.strip(): | |
| continue | |
| chunks = chunk_text(text) | |
| if not chunks: | |
| continue | |
| # Embeddings mit EmbeddingGemma – nutzt automatisch passende Prompts | |
| doc_embeddings = embed_model.encode_document( | |
| chunks, | |
| convert_to_tensor=True, | |
| ) # shape: (num_chunks, dim) | |
| for chunk, emb in zip(chunks, doc_embeddings): | |
| DOCUMENT_STORE.append( | |
| { | |
| "content": chunk, | |
| "embedding": emb, # torch.Tensor | |
| "source": os.path.basename(path), | |
| } | |
| ) | |
| added_chunks += 1 | |
| if added_chunks == 0: | |
| return "Keine verwertbaren Text-Chunks gefunden." | |
| return f"Index aktualisiert: {len(DOCUMENT_STORE)} Chunks gespeichert (neu hinzugefügt: {added_chunks})." | |
| def clear_index() -> str: | |
| DOCUMENT_STORE.clear() | |
| return "Index geleert (0 Chunks)." | |
| def retrieve_relevant_chunks( | |
| query: str, | |
| top_k: int = 5, | |
| ) -> List[Dict]: | |
| """ | |
| Einfacher Dense-Retriever auf Basis von EmbeddingGemma-300M. | |
| Nutzt die SentenceTransformers-Similarity-Funktion (inner product / cosine). | |
| """ | |
| if not DOCUMENT_STORE: | |
| return [] | |
| embed_model = get_embedding_model() | |
| # Query-Embedding | |
| q_emb = embed_model.encode_query( | |
| query, | |
| convert_to_tensor=True, | |
| ) # (dim,) | |
| # Stack aller Dokument-Embeddings | |
| doc_embs = torch.stack([d["embedding"] for d in DOCUMENT_STORE]) # (N, dim) | |
| # Similarity via SentenceTransformers-API (richtige Metrik laut Config) | |
| sims = embed_model.similarity(q_emb, doc_embs)[0] # (N,) | |
| top_k = min(top_k, len(DOCUMENT_STORE)) | |
| scores, indices = torch.topk(sims, k=top_k) | |
| results = [] | |
| for score, idx in zip(scores.tolist(), indices.tolist()): | |
| entry = DOCUMENT_STORE[idx] | |
| results.append( | |
| { | |
| "content": entry["content"], | |
| "source": entry["source"], | |
| "score": float(score), | |
| } | |
| ) | |
| return results | |
| # -------------------------------------------------------------------- | |
| # LLM-Generierung (Gemma-3-4B-IT) | |
| # -------------------------------------------------------------------- | |
| def build_system_prompt() -> str: | |
| return ( | |
| "Du bist ein hilfreicher, präziser Assistent. " | |
| "Du beantwortest Fragen basierend auf bereitgestellten Dokumenten. " | |
| "Wenn Informationen nicht im Kontext sind, sag explizit, dass sie nicht im Kontext stehen. " | |
| "Antworte bitte auf Deutsch und zitiere ggf. die relevanten Teile in eigenen Worten." | |
| ) | |
| def build_user_prompt(user_question: str, retrieved_chunks: List[Dict]) -> str: | |
| if not retrieved_chunks: | |
| return ( | |
| f"Benutzerfrage:\n{user_question}\n\n" | |
| "Es wurden keine Dokumente im Kontext gefunden. " | |
| "Beantworte die Frage so gut wie möglich, aber weise darauf hin, " | |
| "dass du keinen Dokumentenkontext hast." | |
| ) | |
| context_str_parts = [] | |
| for i, ch in enumerate(retrieved_chunks, start=1): | |
| context_str_parts.append( | |
| f"[{i}] (Quelle: {ch['source']}, Score: {ch['score']:.3f})\n{ch['content']}" | |
| ) | |
| context_str = "\n\n".join(context_str_parts) | |
| prompt = ( | |
| f"Benutzerfrage:\n{user_question}\n\n" | |
| "Hier sind die relevantesten Kontextpassagen aus den hochgeladenen Dokumenten:\n" | |
| f"{context_str}\n\n" | |
| "Aufgabe:\n" | |
| "- Nutze primär diese Kontexte.\n" | |
| "- Wenn etwas nicht im Kontext steht, sag das klar.\n" | |
| "- Fasse die relevanten Stellen strukturiert zusammen.\n" | |
| "- Antworte auf Deutsch.\n\n" | |
| "Antwort:" | |
| ) | |
| return prompt | |
| def answer_with_rag(question: str, history: List[Tuple[str, str]]) -> str: | |
| model, processor = get_llm() | |
| # 1. Retrieve relevante Chunks | |
| retrieved = retrieve_relevant_chunks(question, top_k=5) | |
| # 2. Prompt bauen | |
| system_prompt = build_system_prompt() | |
| user_prompt = build_user_prompt(question, retrieved) | |
| messages = [ | |
| { | |
| "role": "system", | |
| "content": [{"type": "text", "text": system_prompt}], | |
| }, | |
| { | |
| "role": "user", | |
| "content": [{"type": "text", "text": user_prompt}], | |
| }, | |
| ] | |
| # 3. Chat Template + Generation | |
| inputs = processor.apply_chat_template( | |
| messages, | |
| add_generation_prompt=True, | |
| tokenize=True, | |
| return_dict=True, | |
| return_tensors="pt", | |
| ).to(model.device, dtype=model.dtype) | |
| input_len = inputs["input_ids"].shape[-1] | |
| with torch.inference_mode(): | |
| generated = model.generate( | |
| **inputs, | |
| max_new_tokens=512, | |
| do_sample=True, | |
| temperature=0.7, | |
| top_p=0.9, | |
| ) | |
| gen_tokens = generated[0][input_len:] | |
| answer = processor.decode(gen_tokens, skip_special_tokens=True) | |
| return answer | |
| # -------------------------------------------------------------------- | |
| # Gradio UI | |
| # -------------------------------------------------------------------- | |
| def ingest_files_ui(files) -> str: | |
| if not files: | |
| return "Bitte zuerst eine oder mehrere Dateien hochladen." | |
| paths = [f.name if hasattr(f, "name") else f for f in files] | |
| return index_files(paths) | |
| def chat_fn(message: str, history: List[Tuple[str, str]]): | |
| """ | |
| Callback für gr.Chatbot: nimmt die neue User-Nachricht, macht RAG+LLM | |
| und hängt Antwort an die History an. | |
| """ | |
| if not message or not message.strip(): | |
| return history | |
| answer = answer_with_rag(message, history) | |
| history = history + [(message, answer)] | |
| return history | |
| def build_demo() -> gr.Blocks: | |
| with gr.Blocks(title="Gemma 3 RAG – EmbeddingGemma-300M", theme="soft") as demo: | |
| gr.Markdown( | |
| """ | |
| # 🔍 Gemma 3 RAG Space | |
| **RAG-Pipeline mit `google/embeddinggemma-300m` + `google/gemma-3-4b-it`** | |
| Schritte: | |
| 1. Links Dateien hochladen und indexieren. | |
| 2. Rechts im Chat Fragen zu deinen Dokumenten stellen. | |
| """ | |
| ) | |
| with gr.Row(): | |
| with gr.Column(scale=1): | |
| gr.Markdown("### 📁 Dokumente") | |
| file_uploader = gr.File( | |
| label="Dateien hochladen (.pdf, .txt, .md, ...)", | |
| file_count="multiple", | |
| type="filepath", | |
| ) | |
| index_button = gr.Button("🔄 Index aktualisieren") | |
| clear_index_button = gr.Button("🧹 Index leeren") | |
| index_status = gr.Markdown("Noch keine Dokumente indexiert.") | |
| index_button.click( | |
| fn=index_files, | |
| inputs=file_uploader, | |
| outputs=index_status, | |
| ) | |
| clear_index_button.click( | |
| fn=clear_index, | |
| inputs=None, | |
| outputs=index_status, | |
| ) | |
| with gr.Column(scale=2): | |
| gr.Markdown("### 💬 Chat mit Gemma-3 über deine Dateien") | |
| chatbot = gr.Chatbot( | |
| label="Kontext-sensitiver Chat", | |
| type="tuple", | |
| show_copy_button=True, | |
| ) | |
| msg = gr.Textbox( | |
| label="Deine Frage", | |
| placeholder="Frag etwas zu den hochgeladenen Dokumenten...", | |
| lines=2, | |
| ) | |
| send_btn = gr.Button("Senden") | |
| clear_chat_btn = gr.Button("Chat löschen") | |
| def user_submit(user_message, chat_history): | |
| if not user_message: | |
| return "", chat_history | |
| chat_history = chat_history + [(user_message, None)] | |
| return "", chat_history | |
| # Ablauf: | |
| # 1. User-Text -> History mit (msg, None) | |
| # 2. Chat-Funktion ersetzt None durch generierte Antwort | |
| msg.submit(user_submit, [msg, chatbot], [msg, chatbot]).then( | |
| chat_fn, chatbot, chatbot | |
| ) | |
| send_btn.click(user_submit, [msg, chatbot], [msg, chatbot]).then( | |
| chat_fn, chatbot, chatbot | |
| ) | |
| clear_chat_btn.click( | |
| fn=lambda: [], | |
| inputs=None, | |
| outputs=chatbot, | |
| ) | |
| return demo | |
| demo = build_demo() | |
| if __name__ == "__main__": | |
| # In HF Spaces wird normalerweise automatisch gestartet, | |
| # aber lokal brauchst du das: | |
| demo.launch() | |