FannyFa-Model-V1 / rag_module.py
FannyFa's picture
Upload 16 files
3eecd6b verified
Raw History Blame Contribute Delete
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
# ==================================================================
@staticmethod
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]
@staticmethod
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
# ==================================================================
@staticmethod
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()
@staticmethod
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()
@staticmethod
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
@staticmethod
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}