Download rag_module.py from FannyFa/FannyFa-Model-V1: direct link, hf CLI and curl.
- Browser
- Download file 38.7 kB
-
https://huggingface.co/FannyFa/FannyFa-Model-V1/resolve/main/rag_module.py
- Command line
-
hf download hf://FannyFa/FannyFa-Model-V1/rag_module.py
-
curl -L -o rag_module.py https://huggingface.co/FannyFa/FannyFa-Model-V1/resolve/main/rag_module.py
38.7 kB
| """ | |
| RAG (Retrieval Augmented Generation) Module - v2.1 | |
| Dengan confidence scoring, quality guard, cache, feedback, query expansion, | |
| dan fallback otomatis ke model (chat biasa) kalau KB kosong. | |
| Mode jawaban: | |
| - Direct answer (skor tinggi) -> ambil jawaban asli dari KB, tanpa model | |
| - Generate dengan konteks -> model dengan contoh jawaban dari KB | |
| - Generate tanpa konteks -> model (chat biasa) kalau KB kosong | |
| - Quality fallback -> jawab "belum bisa" kalau output buruk | |
| """ | |
| import re | |
| import json | |
| import time | |
| import hashlib | |
| from datetime import datetime | |
| from pathlib import Path | |
| from typing import List, Dict, Optional, Tuple, Any | |
| from collections import OrderedDict | |
| import torch | |
| from transformers import AutoTokenizer, AutoModelForCausalLM | |
| from semantic_search import HybridSearch, KnowledgeBase, SearchResult | |
| from utils import console, Theme, error_logger, debug_logger, get_device | |
| from models import ModelManager | |
| # --------------------------------------------------------------------------- | |
| # Boilerplate yang sering muncul & harus dipotong dari output model | |
| # --------------------------------------------------------------------------- | |
| _BOILERPLATE_MARKERS = [ | |
| "semoga membantu", | |
| "saya siap membantu lebih lanjut", | |
| "ada lagi yang ingin ditanyakan", | |
| "ada lagi yang bisa dibantu", | |
| "kalau ada pertanyaan lain, silakan", | |
| "apakah ada yang ingin didiskusikan", | |
| "jangan ragu untuk bertanya", | |
| "semoga informasinya jelas", | |
| "bisa ditanyakan lebih detail", | |
| "terima kasih sudah bertanya", | |
| "silakan ajukan pertanyaan yang sopan", | |
| "mari jaga percakapan tetap baik", | |
| "mari fokus pada solusi", | |
| "semoga kita bisa bicara lebih positif", | |
| "saya mohon maaf atas ketidaknyamanan", | |
| "semoga harimu menyenangkan", | |
| ] | |
| # Typo command yang umum → command asli | |
| _COMMAND_ALIASES = { | |
| "xlear": "clear", | |
| "claer": "clear", | |
| "clea": "clear", | |
| "exi": "exit", | |
| "ext": "exit", | |
| "quitt": "exit", | |
| "hlp": "help", | |
| "hep": "help", | |
| } | |
| # Sinonim sederhana untuk query expansion (Indonesia informal) | |
| _SYNONYMS = { | |
| "halo": ["hai", "hello", "hi"], | |
| "hai": ["halo", "hello", "hi"], | |
| "kabar": ["keadaan", "kondisi", "situasi"], | |
| "siapa": ["siapakah", "nama"], | |
| "apa": ["apakah", "apa itu"], | |
| "gimana": ["bagaimana", "caranya"], | |
| "kenapa": ["mengapa", "kok"], | |
| "bantu": ["tolong", "bantuin", "help"], | |
| "belajar": ["pelajari", "study"], | |
| "terima kasih": ["makasih", "thanks", "thank you"], | |
| "maaf": ["sorry", "mohon maaf"], | |
| } | |
| # File untuk persist feedback & session | |
| _FEEDBACK_FILE = Path("/content/data/rag_feedback.json") | |
| _SESSION_FILE_PATTERN = "/content/data/rag_session_{ts}.json" | |
| class RAGPipeline: | |
| """ | |
| RAG Pipeline - Retrieve relevant context, lalu generate response. | |
| Kalau KB kosong, otomatis fallback ke mode chat biasa (pakai model). | |
| """ | |
| def __init__( | |
| self, | |
| knowledge_base: Optional[KnowledgeBase] = None, | |
| embedding_model: str = "all-MiniLM-L6-v2", | |
| search_method: str = "hybrid", | |
| ): | |
| self.kb = knowledge_base or KnowledgeBase(embedding_model=embedding_model) | |
| self.search_method = search_method | |
| self.tokenizer = None | |
| self.model = None | |
| self.device = get_device() | |
| self.model_manager = ModelManager() | |
| self.search_config = { | |
| "top_k": 5, | |
| "chunk_size": 400, | |
| "min_score": 0.2, | |
| "expand_query": True, | |
| } | |
| self.generation_config = { | |
| "max_new_tokens": 100, | |
| "temperature": 0.7, | |
| "top_p": 0.9, | |
| "top_k": 40, | |
| "repetition_penalty": 1.4, | |
| "no_repeat_ngram_size": 3, | |
| "max_sentences": 3, | |
| "max_chars": 500, | |
| "context_window_turns": 3, | |
| # Mode direct answer | |
| "prefer_direct_answer": True, | |
| "direct_answer_threshold": 0.5, | |
| "min_context_score": 0.30, | |
| "verbose": False, | |
| # Quality guard | |
| "min_answer_length": 3, | |
| "enable_cache": True, | |
| "cache_max_size": 100, | |
| # PENTING: kalau False, RAG akan pakai model tanpa KB | |
| "honesty_fallback": False, | |
| } | |
| self._session_history: List[Tuple[str, str]] = [] | |
| self._response_cache: "OrderedDict[str, Dict[str, Any]]" = OrderedDict() | |
| self._feedback: List[Dict[str, Any]] = [] | |
| self._stats = { | |
| "total_queries": 0, | |
| "direct_hits": 0, | |
| "generated": 0, | |
| "fallback": 0, | |
| "cache_hits": 0, | |
| "no_match": 0, | |
| "positive_feedback": 0, | |
| "negative_feedback": 0, | |
| } | |
| self._load_feedback() | |
| def set_model(self, tokenizer, model) -> None: | |
| """Set tokenizer dan model untuk generation""" | |
| self.tokenizer = tokenizer | |
| self.model = model | |
| # ================================================================== | |
| # Query expansion | |
| # ================================================================== | |
| def _expand_query(self, query: str) -> str: | |
| """Tambah sinonim untuk query retrieval yang lebih baik.""" | |
| if not self.search_config.get("expand_query", True): | |
| return query | |
| words = query.lower().split() | |
| expanded_words = list(words) | |
| for word in words: | |
| clean = word.strip("?!.,;:") | |
| if clean in _SYNONYMS: | |
| expanded_words.append(_SYNONYMS[clean][0]) | |
| return " ".join(expanded_words) | |
| # ================================================================== | |
| # Cache | |
| # ================================================================== | |
| def _query_key(query: str) -> str: | |
| return hashlib.md5(query.strip().lower().encode()).hexdigest() | |
| def _cache_get(self, query: str) -> Optional[Dict[str, Any]]: | |
| if not self.generation_config.get("enable_cache", True): | |
| return None | |
| key = self._query_key(query) | |
| if key in self._response_cache: | |
| self._stats["cache_hits"] += 1 | |
| self._response_cache.move_to_end(key) | |
| return self._response_cache[key] | |
| return None | |
| def _cache_set(self, query: str, result: Dict[str, Any]) -> None: | |
| if not self.generation_config.get("enable_cache", True): | |
| return | |
| key = self._query_key(query) | |
| max_size = self.generation_config.get("cache_max_size", 100) | |
| self._response_cache[key] = result | |
| while len(self._response_cache) > max_size: | |
| self._response_cache.popitem(last=False) | |
| def _cache_clear(self) -> int: | |
| n = len(self._response_cache) | |
| self._response_cache.clear() | |
| return n | |
| # ================================================================== | |
| # Retrieval & parsing | |
| # ================================================================== | |
| def retrieve(self, query: str, top_k: Optional[int] = None) -> List[Dict[str, Any]]: | |
| """Retrieve relevant documents dari knowledge base.""" | |
| top_k = top_k or self.search_config["top_k"] | |
| expanded = self._expand_query(query) | |
| queries = [query] if expanded == query else [query, expanded] | |
| seen = set() | |
| results: List[Dict[str, Any]] = [] | |
| for q in queries: | |
| try: | |
| res = self.kb.search(q, top_k=top_k, method=self.search_method) | |
| except Exception as e: | |
| error_logger.error(f"KB search error: {e}") | |
| continue | |
| for r in res: | |
| key = r.get("text", "")[:80] | |
| if key in seen: | |
| continue | |
| seen.add(key) | |
| results.append(r) | |
| results = [r for r in results if r.get("score", 0) >= self.search_config["min_score"]] | |
| results.sort(key=lambda r: r.get("score", 0), reverse=True) | |
| return results[:top_k] | |
| def _parse_retrieved_qa(text: str) -> Optional[Tuple[str, str]]: | |
| """Parse 'User: X AI: Y' dari teks KB.""" | |
| if not text or "AI:" not in text: | |
| return None | |
| idx = text.rfind("AI:") | |
| if idx == -1: | |
| return None | |
| before = text[:idx] | |
| after = text[idx + 3:] | |
| user_part = before | |
| first_user = user_part.find("User:") | |
| if first_user != -1: | |
| user_part = user_part[first_user + 5:] | |
| for marker in ["User:", "AI:"]: | |
| p = after.find(marker) | |
| if p != -1: | |
| after = after[:p] | |
| user_part = user_part.strip() | |
| ai_part = after.strip() | |
| if not user_part or not ai_part: | |
| return None | |
| return user_part, ai_part | |
| def _extract_qa_pairs(self, docs: List[Dict[str, Any]]) -> List[Dict[str, Any]]: | |
| pairs = [] | |
| for doc in docs: | |
| parsed = self._parse_retrieved_qa(doc.get("text", "")) | |
| if parsed: | |
| q, a = parsed | |
| pairs.append({"q": q, "a": a, "score": doc["score"]}) | |
| return pairs | |
| # ================================================================== | |
| # Confidence | |
| # ================================================================== | |
| def _confidence_label(score: float) -> Tuple[str, str]: | |
| if score >= 0.6: | |
| return "HIGH", "green" | |
| elif score >= 0.4: | |
| return "MEDIUM", "yellow" | |
| elif score >= 0.25: | |
| return "LOW", "orange1" | |
| else: | |
| return "VERY LOW", "red" | |
| # ================================================================== | |
| # Quality guard | |
| # ================================================================== | |
| def _is_answer_quality_ok(self, answer: str) -> bool: | |
| if not answer: | |
| return False | |
| min_words = self.generation_config.get("min_answer_length", 3) | |
| words = answer.split() | |
| if len(words) < min_words: | |
| return False | |
| unique = set(w.lower() for w in words) | |
| if len(unique) < max(2, len(words) // 4): | |
| return False | |
| alnum = sum(1 for c in answer if c.isalnum() or c.isspace()) | |
| if len(answer) > 0 and alnum / len(answer) < 0.6: | |
| return False | |
| return True | |
| # ================================================================== | |
| # Generation | |
| # ================================================================== | |
| def generate_with_context( | |
| self, | |
| query: str, | |
| retrieved_docs: List[Dict[str, Any]], | |
| system_prompt: str = None, | |
| ) -> Tuple[str, str, float]: | |
| """ | |
| Return (response, mode, confidence_score) | |
| Prioritas: | |
| 1. Direct answer (skor tinggi dari KB) | |
| 2. Generate dengan contoh dari KB | |
| 3. Generate tanpa context (chat biasa) ← kalau KB kosong | |
| """ | |
| if not self.tokenizer or not self.model: | |
| qa_pairs = self._extract_qa_pairs(retrieved_docs) | |
| if qa_pairs and qa_pairs[0]["score"] >= self.generation_config["direct_answer_threshold"]: | |
| return qa_pairs[0]["a"], "direct", qa_pairs[0]["score"] | |
| return "Model tidak tersedia", "no_model", 0.0 | |
| qa_pairs = self._extract_qa_pairs(retrieved_docs) | |
| confidence = qa_pairs[0]["score"] if qa_pairs else 0.0 | |
| # ---------- STRATEGI 1: DIRECT ANSWER dari KB ---------- | |
| if self.generation_config["prefer_direct_answer"] and qa_pairs: | |
| best = qa_pairs[0] | |
| if best["score"] >= self.generation_config["direct_answer_threshold"]: | |
| cleaned = self._clean_response(best["a"]) | |
| if self._is_answer_quality_ok(cleaned): | |
| return cleaned, "direct", best["score"] | |
| # ---------- STRATEGI 2: GENERATE dengan contoh KB ---------- | |
| examples = [ | |
| p for p in qa_pairs | |
| if p["score"] >= self.generation_config["min_context_score"] | |
| ][:3] | |
| if examples: | |
| mode = "generate_with_context" | |
| else: | |
| mode = "generate_no_context" # fallback chat biasa | |
| # Bangun prompt | |
| prompt_parts = [] | |
| if examples: | |
| example_block = "\n\n".join( | |
| f"User: {e['q']}\nAI: {e['a']}" for e in examples | |
| ) | |
| prompt_parts.append(f"Contoh percakapan:\n{example_block}") | |
| if self._session_history: | |
| n = self.generation_config["context_window_turns"] | |
| recent = self._session_history[-n:] | |
| hist_block = "\n\n".join( | |
| f"User: {q[:80]}\nAI: {a[:120]}" for q, a in recent | |
| ) | |
| prompt_parts.append(f"Percakapan sebelumnya:\n{hist_block}") | |
| prompt_parts.append(f"User: {query}") | |
| prompt = "\n\n".join(prompt_parts) + "\nAI:" | |
| try: | |
| inputs = self.tokenizer.encode( | |
| prompt, | |
| return_tensors="pt", | |
| truncation=True, | |
| max_length=768, | |
| ).to(self.device) | |
| input_length = inputs.shape[1] | |
| model_max_len = getattr(self.model.config, "max_position_embeddings", 1024) | |
| safe_max_new = min( | |
| self.generation_config["max_new_tokens"], | |
| model_max_len - input_length - 10, | |
| ) | |
| if safe_max_new <= 0: | |
| keep = model_max_len - 60 | |
| inputs = inputs[:, -keep:] | |
| input_length = inputs.shape[1] | |
| safe_max_new = 50 | |
| gen_kwargs = { | |
| "max_new_tokens": max(1, safe_max_new), | |
| "temperature": self.generation_config["temperature"], | |
| "top_p": self.generation_config["top_p"], | |
| "top_k": self.generation_config["top_k"], | |
| "repetition_penalty": self.generation_config["repetition_penalty"], | |
| "no_repeat_ngram_size": self.generation_config["no_repeat_ngram_size"], | |
| "do_sample": True, | |
| "pad_token_id": self.tokenizer.eos_token_id, | |
| "eos_token_id": self.tokenizer.eos_token_id, | |
| "use_cache": True, | |
| } | |
| with torch.no_grad(): | |
| outputs = self.model.generate(inputs, **gen_kwargs) | |
| generated = outputs[0, input_length:] | |
| raw = self.tokenizer.decode(generated, skip_special_tokens=True) | |
| cleaned = self._clean_response(raw) | |
| # Quality guard — kalau output aneh, coba contoh KB dulu | |
| if not self._is_answer_quality_ok(cleaned): | |
| for p in qa_pairs: | |
| if p["score"] >= self.generation_config["min_context_score"]: | |
| fallback = self._clean_response(p["a"]) | |
| if self._is_answer_quality_ok(fallback): | |
| return fallback, "direct_fallback", p["score"] | |
| return cleaned or "...", mode, confidence | |
| # Kalau cuma boilerplate + ada contoh KB, pakai KB | |
| if self._is_boilerplate_only(cleaned) and qa_pairs: | |
| for p in qa_pairs: | |
| if p["score"] >= self.generation_config["min_context_score"]: | |
| return p["a"], "direct_fallback", p["score"] | |
| return cleaned, mode, confidence | |
| except torch.cuda.OutOfMemoryError: | |
| console.print(Theme.error(" GPU OOM saat RAG generation")) | |
| try: | |
| self.model = self.model.to("cpu") | |
| torch.cuda.empty_cache() | |
| self.device = "cpu" | |
| return self.generate_with_context(query, retrieved_docs, system_prompt) | |
| except Exception as e: | |
| return f"Error: {str(e)[:100]}", "error", 0.0 | |
| except Exception as e: | |
| error_logger.error(f"RAG generation error: {e}") | |
| return f"Error generating response: {str(e)[:120]}", "error", 0.0 | |
| # ================================================================== | |
| # Response cleaning | |
| # ================================================================== | |
| def _clean_response(self, response: str) -> str: | |
| if not response: | |
| return "" | |
| for marker in ["User:", "AI:", "<|endoftext|>", "<|end|>"]: | |
| pos = response.find(marker) | |
| if pos != -1: | |
| response = response[:pos] | |
| response = ( | |
| response.replace("<|endoftext|>", "") | |
| .replace("<|end|>", "") | |
| .strip() | |
| ) | |
| response = " ".join(response.split()) | |
| response = self._fix_punctuation(response) | |
| response = self._strip_boilerplate(response) | |
| response = self._first_n_sentences( | |
| response, self.generation_config["max_sentences"] | |
| ) | |
| max_chars = self.generation_config["max_chars"] | |
| if len(response) > max_chars: | |
| response = response[:max_chars].rsplit(" ", 1)[0] + "..." | |
| return response.strip() | |
| def _fix_punctuation(text: str) -> str: | |
| text = re.sub(r"[;:,]+\.", ".", text) | |
| text = re.sub(r"\.\.+", ".", text) | |
| text = re.sub(r"[,]{2,}", ",", text) | |
| text = re.sub(r"[;]{2,}", ";", text) | |
| text = re.sub(r"\s+([.,!?;:])", r"\1", text) | |
| text = re.sub(r"([.,!?;:])([^\s\d])", r"\1 \2", text) | |
| text = re.sub(r"\s+\.", ".", text) | |
| return text.strip() | |
| def _strip_boilerplate(text: str) -> str: | |
| if not text: | |
| return text | |
| lower = text.lower() | |
| cutoff = len(text) | |
| min_pos = int(len(text) * 0.3) | |
| for marker in _BOILERPLATE_MARKERS: | |
| pos = lower.find(marker, min_pos) | |
| if pos != -1 and pos < cutoff: | |
| cutoff = pos | |
| if cutoff < len(text): | |
| text = text[:cutoff].rstrip(" ,;.") | |
| return text | |
| def _first_n_sentences(text: str, n: int) -> str: | |
| if not text: | |
| return text | |
| sentences = re.split(r"(?<=[.!?])\s+", text) | |
| if len(sentences) <= n: | |
| return text | |
| trimmed = " ".join(sentences[:n]).strip() | |
| if trimmed and trimmed[-1] not in ".!?": | |
| trimmed += "." | |
| return trimmed | |
| def _is_boilerplate_only(self, text: str) -> bool: | |
| if not text: | |
| return True | |
| lower = text.lower() | |
| boilerplate_count = sum(1 for m in _BOILERPLATE_MARKERS if m in lower) | |
| words = len(text.split()) | |
| if words < 15 and boilerplate_count >= 1: | |
| return True | |
| if words < 8: | |
| return True | |
| return False | |
| # ================================================================== | |
| # Public query API | |
| # ================================================================== | |
| def query( | |
| self, | |
| query: str, | |
| retrieve: bool = True, | |
| top_k: Optional[int] = None, | |
| system_prompt: str = None, | |
| use_cache: bool = True, | |
| ) -> Dict[str, Any]: | |
| self._stats["total_queries"] += 1 | |
| if use_cache: | |
| cached = self._cache_get(query) | |
| if cached is not None: | |
| return {**cached, "from_cache": True} | |
| retrieved_docs = [] | |
| if retrieve: | |
| retrieved_docs = self.retrieve(query, top_k) | |
| response, mode, confidence = self.generate_with_context( | |
| query, retrieved_docs, system_prompt | |
| ) | |
| if mode == "direct": | |
| self._stats["direct_hits"] += 1 | |
| elif mode in ("generate_with_context", "generate_no_context"): | |
| self._stats["generated"] += 1 | |
| elif mode in ("honest_fallback", "quality_fallback"): | |
| self._stats["no_match"] += 1 | |
| self._session_history.append((query, response)) | |
| result = { | |
| "query": query, | |
| "response": response, | |
| "mode": mode, | |
| "confidence": confidence, | |
| "retrieved_documents": retrieved_docs, | |
| "num_sources": len(retrieved_docs), | |
| "from_cache": False, | |
| } | |
| if use_cache: | |
| self._cache_set(query, result) | |
| return result | |
| def query_multi( | |
| self, query: str, top_k: Optional[int] = None | |
| ) -> List[Dict[str, Any]]: | |
| """Ambil beberapa kandidat jawaban dari KB (bukan generate).""" | |
| retrieved_docs = self.retrieve(query, top_k or 3) | |
| qa_pairs = self._extract_qa_pairs(retrieved_docs) | |
| return [ | |
| { | |
| "answer": p["a"], | |
| "question_matched": p["q"], | |
| "score": p["score"], | |
| } | |
| for p in qa_pairs[:3] | |
| ] | |
| # ================================================================== | |
| # Feedback & persistence | |
| # ================================================================== | |
| def _load_feedback(self) -> None: | |
| try: | |
| if _FEEDBACK_FILE.exists(): | |
| with open(_FEEDBACK_FILE) as f: | |
| self._feedback = json.load(f) | |
| pos = sum(1 for fb in self._feedback if fb.get("rating") == "good") | |
| neg = sum(1 for fb in self._feedback if fb.get("rating") == "bad") | |
| self._stats["positive_feedback"] = pos | |
| self._stats["negative_feedback"] = neg | |
| except Exception as e: | |
| debug_logger.debug(f"Feedback load error: {e}") | |
| def _save_feedback(self) -> None: | |
| try: | |
| _FEEDBACK_FILE.parent.mkdir(parents=True, exist_ok=True) | |
| with open(_FEEDBACK_FILE, "w") as f: | |
| json.dump(self._feedback, f, indent=2, ensure_ascii=False) | |
| except Exception as e: | |
| debug_logger.debug(f"Feedback save error: {e}") | |
| def give_feedback(self, rating: str, note: str = "") -> None: | |
| if not self._session_history: | |
| console.print(Theme.warning("Belum ada percakapan untuk diberi feedback")) | |
| return | |
| last_q, last_a = self._session_history[-1] | |
| self._feedback.append({ | |
| "query": last_q, | |
| "answer": last_a, | |
| "rating": rating, | |
| "note": note, | |
| "timestamp": datetime.now().isoformat(), | |
| }) | |
| if rating == "good": | |
| self._stats["positive_feedback"] += 1 | |
| elif rating == "bad": | |
| self._stats["negative_feedback"] += 1 | |
| self._save_feedback() | |
| console.print(Theme.success(f"✓ Feedback '{rating}' disimpan")) | |
| def export_session(self) -> Optional[Path]: | |
| if not self._session_history: | |
| console.print(Theme.warning("Belum ada riwayat")) | |
| return None | |
| ts = datetime.now().strftime("%Y%m%d_%H%M%S") | |
| out = Path(_SESSION_FILE_PATTERN.format(ts=ts)) | |
| out.parent.mkdir(parents=True, exist_ok=True) | |
| try: | |
| with open(out, "w", encoding="utf-8") as f: | |
| json.dump({ | |
| "exported": datetime.now().isoformat(), | |
| "session": [ | |
| {"user": q, "ai": a} for q, a in self._session_history | |
| ], | |
| "stats": self._stats, | |
| }, f, indent=2, ensure_ascii=False) | |
| console.print(Theme.success(f"✓ Session disimpan: {out}")) | |
| return out | |
| except Exception as e: | |
| console.print(Theme.error(f"Gagal export: {e}")) | |
| return None | |
| # ================================================================== | |
| # Interactive | |
| # ================================================================== | |
| def query_interactive(self) -> None: | |
| from rich.prompt import Prompt | |
| from rich.panel import Panel | |
| try: | |
| kb_stats = self.kb.get_stats() | |
| total_docs = kb_stats.get("total_documents", 0) | |
| except Exception: | |
| total_docs = 0 | |
| console.print( | |
| Panel( | |
| "[cyan]RAG Interactive Mode v2.1[/cyan]\n" | |
| "[white]Commands: 'exit' keluar | 'clear'/'reset' | 'help' | " | |
| "'stats' | 'history' | 'topk N' | 'direct on/off' | " | |
| "'threshold X' | 'verbose on/off' | 'cache on/off' | " | |
| "'cache clear' | 'good' | 'bad [note]' | 'export' | 'multi <query>'[/white]", | |
| style="cyan", | |
| ) | |
| ) | |
| if total_docs == 0: | |
| console.print(Theme.info( | |
| " Mode: Chat with model (tanpa Knowledge Base)" | |
| )) | |
| console.print(Theme.dim( | |
| " Jawaban langsung dari model. Mau pakai KB? " | |
| "Menu [09] → Import dari Dataset" | |
| )) | |
| else: | |
| console.print(Theme.dim( | |
| f" Mode: RAG (KB berisi {total_docs} dokumen) | " | |
| f"{'direct+generate' if self.generation_config['prefer_direct_answer'] else 'generate only'} " | |
| f"| cache: {'ON' if self.generation_config['enable_cache'] else 'OFF'}" | |
| )) | |
| while True: | |
| try: | |
| raw = Prompt.ask("[cyan]Query[/cyan]").strip() | |
| if not raw: | |
| continue | |
| lower = raw.lower().strip() | |
| if lower in _COMMAND_ALIASES: | |
| lower = _COMMAND_ALIASES[lower] | |
| # ---------- EXIT ---------- | |
| if lower in ("exit", "quit", "q"): | |
| console.print(Theme.success("Keluar dari RAG")) | |
| break | |
| # ---------- CLEAR / RESET ---------- | |
| if lower in ("clear", "reset"): | |
| self._session_history.clear() | |
| console.print(Theme.success(" Riwayat sesi dihapus")) | |
| continue | |
| # ---------- HELP ---------- | |
| if lower == "help": | |
| console.print(Panel( | |
| "COMMANDS:\n" | |
| " exit - keluar\n" | |
| " clear / reset - hapus riwayat sesi\n" | |
| " history - lihat riwayat sesi\n" | |
| " export - simpan sesi ke file\n" | |
| " help - tampilkan bantuan ini\n" | |
| " stats - statistik KB & config\n" | |
| " topk N - jumlah dokumen retrieval\n" | |
| " direct on/off - mode direct answer\n" | |
| " threshold X - threshold direct answer (0.0-1.0)\n" | |
| " verbose on/off - tampilkan mode & confidence\n" | |
| " cache on/off - aktif/nonaktif respons cache\n" | |
| " cache clear - kosongkan cache\n" | |
| " good - tandai jawaban terakhir benar\n" | |
| " bad [note] - tandai jawaban terakhir salah\n" | |
| " multi <query> - tampilkan 3 kandidat jawaban\n", | |
| title="RAG HELP", style="yellow" | |
| )) | |
| continue | |
| # ---------- HISTORY ---------- | |
| if lower == "history": | |
| if not self._session_history: | |
| console.print(Theme.warning("Belum ada riwayat")) | |
| else: | |
| console.print("\n" + Theme.header(" Session History:")) | |
| for i, (q, a) in enumerate(self._session_history, 1): | |
| console.print(f"\n[cyan][{i}] User:[/cyan] {q}") | |
| console.print(f"[green] AI:[/green] {a}") | |
| continue | |
| # ---------- EXPORT ---------- | |
| if lower == "export": | |
| self.export_session() | |
| continue | |
| # ---------- FEEDBACK ---------- | |
| if lower == "good": | |
| self.give_feedback("good") | |
| continue | |
| if lower.startswith("bad"): | |
| note = raw[3:].strip() if len(raw) > 3 else "" | |
| self.give_feedback("bad", note) | |
| continue | |
| # ---------- STATS ---------- | |
| if lower == "stats": | |
| try: | |
| kb_stats = self.kb.get_stats() | |
| console.print(f"[cyan]KB documents:[/cyan] {kb_stats.get('total_documents', 0)}") | |
| console.print(f"[cyan]Avg length:[/cyan] {kb_stats.get('average_doc_length', 0):.1f} chars") | |
| console.print(f"[cyan]Search method:[/cyan] {self.search_method}") | |
| console.print(f"[cyan]Top K:[/cyan] {self.search_config['top_k']}") | |
| console.print(f"[cyan]Direct threshold:[/cyan] {self.generation_config['direct_answer_threshold']}") | |
| console.print(f"[cyan]Cache size:[/cyan] {len(self._response_cache)}") | |
| console.print() | |
| console.print(Theme.header("RAG Statistics:")) | |
| for k, v in self._stats.items(): | |
| console.print(f" {k}: {v}") | |
| except Exception as e: | |
| console.print(Theme.error(f"Stats error: {e}")) | |
| continue | |
| # ---------- TOPK ---------- | |
| if lower.startswith("topk "): | |
| try: | |
| new_k = int(lower.split()[1]) | |
| if new_k > 0: | |
| self.search_config["top_k"] = new_k | |
| console.print(Theme.success(f" Top K → {new_k}")) | |
| else: | |
| console.print(Theme.warning("Top K harus > 0")) | |
| except (ValueError, IndexError): | |
| console.print(Theme.error("Format: topk <angka>")) | |
| continue | |
| # ---------- DIRECT ---------- | |
| if lower.startswith("direct "): | |
| val = lower.split()[1] if len(lower.split()) > 1 else "" | |
| if val in ("on", "true", "1", "yes"): | |
| self.generation_config["prefer_direct_answer"] = True | |
| console.print(Theme.success(" Direct answer: ON")) | |
| elif val in ("off", "false", "0", "no"): | |
| self.generation_config["prefer_direct_answer"] = False | |
| console.print(Theme.success(" Direct answer: OFF")) | |
| else: | |
| console.print(Theme.error("Format: direct on|off")) | |
| continue | |
| # ---------- THRESHOLD ---------- | |
| if lower.startswith("threshold "): | |
| try: | |
| val = float(lower.split()[1]) | |
| if 0.0 <= val <= 1.0: | |
| self.generation_config["direct_answer_threshold"] = val | |
| console.print(Theme.success(f" Threshold → {val}")) | |
| else: | |
| console.print(Theme.warning("Threshold harus 0.0-1.0")) | |
| except (ValueError, IndexError): | |
| console.print(Theme.error("Format: threshold <0.0-1.0>")) | |
| continue | |
| # ---------- VERBOSE ---------- | |
| if lower.startswith("verbose "): | |
| val = lower.split()[1] if len(lower.split()) > 1 else "" | |
| if val in ("on", "true", "1", "yes"): | |
| self.generation_config["verbose"] = True | |
| console.print(Theme.success(" Verbose: ON")) | |
| elif val in ("off", "false", "0", "no"): | |
| self.generation_config["verbose"] = False | |
| console.print(Theme.success(" Verbose: OFF")) | |
| else: | |
| console.print(Theme.error("Format: verbose on|off")) | |
| continue | |
| # ---------- CACHE ---------- | |
| if lower.startswith("cache "): | |
| cmd = lower.split()[1] | |
| if cmd in ("on", "true", "1", "yes"): | |
| self.generation_config["enable_cache"] = True | |
| console.print(Theme.success(" Cache: ON")) | |
| elif cmd in ("off", "false", "0", "no"): | |
| self.generation_config["enable_cache"] = False | |
| console.print(Theme.success(" Cache: OFF")) | |
| elif cmd == "clear": | |
| n = self._cache_clear() | |
| console.print(Theme.success(f" Cache cleared ({n} entries)")) | |
| else: | |
| console.print(Theme.error("Format: cache on|off|clear")) | |
| continue | |
| # ---------- MULTI ---------- | |
| if lower.startswith("multi "): | |
| q = raw[6:].strip() | |
| if not q: | |
| console.print(Theme.error("Format: multi <query>")) | |
| continue | |
| candidates = self.query_multi(q, top_k=3) | |
| if not candidates: | |
| console.print(Theme.warning("Tidak ada kandidat")) | |
| else: | |
| console.print(Theme.header(f"\n Top {len(candidates)} kandidat:")) | |
| for i, c in enumerate(candidates, 1): | |
| label, color = self._confidence_label(c["score"]) | |
| console.print(f"\n[{color}]#{i} [{label} {c['score']:.3f}][/{color}]") | |
| console.print(f"[dim]Match: {c['question_matched'][:80]}[/dim]") | |
| console.print(c["answer"]) | |
| continue | |
| # ---------- NORMAL QUERY ---------- | |
| console.print(Theme.info("Retrieving documents...")) | |
| result = self.query(raw) | |
| if result["from_cache"]: | |
| console.print(Theme.dim(" (from cache)")) | |
| elif result["retrieved_documents"]: | |
| console.print( | |
| Theme.header(f"\n Retrieved {result['num_sources']} documents:") | |
| ) | |
| for idx, doc in enumerate(result["retrieved_documents"], 1): | |
| score = doc["score"] | |
| label, color = self._confidence_label(score) | |
| snippet = doc["text"][:100].replace("\n", " ") | |
| console.print( | |
| f" [{idx}] [{color}]{label} {score:.3f}[/{color}] {snippet}..." | |
| ) | |
| else: | |
| console.print(Theme.dim(" (Tidak ada dokumen relevan — pakai model langsung)")) | |
| if self.generation_config["verbose"]: | |
| label, color = self._confidence_label(result["confidence"]) | |
| console.print(Theme.dim( | |
| f" Mode: {result['mode']} | Confidence: " | |
| ) + f"[{color}]{label} {result['confidence']:.3f}[/{color}]") | |
| console.print(Theme.header("\n Response:")) | |
| console.print(result["response"]) | |
| console.print() | |
| except KeyboardInterrupt: | |
| console.print("\n" + Theme.warning("RAG session stopped")) | |
| break | |
| except Exception as e: | |
| console.print(Theme.error(f"Error: {e}")) | |
| # ================================================================== | |
| # Misc | |
| # ================================================================== | |
| def update_knowledge_base(self, documents: List[Dict[str, Any]]) -> None: | |
| self.kb.add_documents(documents) | |
| console.print(Theme.success(f"Added {len(documents)} documents to KB")) | |
| def set_search_config(self, **kwargs) -> None: | |
| for key, value in kwargs.items(): | |
| if key in self.search_config: | |
| self.search_config[key] = value | |
| console.print(Theme.info("Search config updated")) | |
| def set_generation_config(self, **kwargs) -> None: | |
| for key, value in kwargs.items(): | |
| if key in self.generation_config: | |
| self.generation_config[key] = value | |
| console.print(Theme.info("Generation config updated")) | |
| def get_kb_stats(self) -> Dict[str, Any]: | |
| return self.kb.get_stats() | |
| # ====================================================================== | |
| # Classes tambahan (backward compat) | |
| # ====================================================================== | |
| class ContextCompressor: | |
| def __init__(self, compression_ratio: float = 0.5): | |
| self.compression_ratio = compression_ratio | |
| def compress_context(self, context: str, query: str = None) -> str: | |
| sentences = context.split(".") | |
| if query: | |
| query_words = set(query.lower().split()) | |
| scores = [ | |
| sum(1 for word in sent.lower().split() if word in query_words) | |
| for sent in sentences | |
| ] | |
| else: | |
| scores = [1] * len(sentences) | |
| num_to_keep = max(1, int(len(sentences) * self.compression_ratio)) | |
| top_indices = sorted( | |
| range(len(scores)), key=lambda i: scores[i], reverse=True | |
| )[:num_to_keep] | |
| top_indices.sort() | |
| return ".".join(sentences[i] for i in top_indices) | |
| class FactChecker: | |
| def __init__(self, model=None, tokenizer=None): | |
| self.model = model | |
| self.tokenizer = tokenizer | |
| def check_consistency( | |
| self, fact: str, context: str, model=None, tokenizer=None | |
| ) -> Dict[str, Any]: | |
| model = model or self.model | |
| tokenizer = tokenizer or self.tokenizer | |
| if not model or not tokenizer: | |
| return {"consistent": True, "confidence": 0.5} | |
| try: | |
| fact_words = set(word.lower() for word in fact.split() if len(word) > 3) | |
| context_words = set( | |
| word.lower() for word in context.split() if len(word) > 3 | |
| ) | |
| overlap = len(fact_words & context_words) | |
| consistency_score = overlap / len(fact_words) if fact_words else 0.5 | |
| return { | |
| "consistent": consistency_score > 0.5, | |
| "confidence": consistency_score, | |
| "overlap_ratio": consistency_score, | |
| } | |
| except Exception as e: | |
| error_logger.error(f"Consistency check error: {e}") | |
| return {"consistent": True, "confidence": 0.5} |