nlp-project / src /demo.py
ervua's picture
Deploy Turkish Legal RAG App
6dfa658
Raw
History Blame Contribute Delete
19.8 kB
"""
Demo script for full baseline Turkish legal RAG pipeline.
Run:
python src/demo.py
- USE_LLM = False: stable extractive answers (default).
- USE_LLM = True: try local FLAN-T5, fall back to extractive if output is poor.
- USE_RERANKER = False: optional lightweight lexical rerank after hybrid (RRF).
"""
from __future__ import annotations
import json
import os
import re
import traceback
from pathlib import Path
from cross_encoder_rerank import maybe_cross_encoder_rerank
from data_loader import build_real_corpus
from generator import LocalGenerator
from ingest import build_chunked_corpus, load_corpus, load_jsonl
from metrics import evaluate_answers, evaluate_groundedness, evaluate_retrieval, gold_doc_ids_from_eval_item
from reranker import simple_lexical_rerank
from retriever import BM25Retriever, DenseRetriever, reciprocal_rank_fusion
EMBED_MODEL_NAME = "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2"
GEN_MODEL_NAME = "google/flan-t5-small"
# Optional local fine-tuned artifacts (set env vars or place models under ./models)
DENSE_MODEL_PATH = os.environ.get("DENSE_MODEL_PATH", "").strip() or EMBED_MODEL_NAME
GEN_MODEL_PATH = os.environ.get("GEN_MODEL_PATH", "").strip() or GEN_MODEL_NAME
CROSS_ENCODER_PATH = os.environ.get("CROSS_ENCODER_PATH", "").strip()
DISPLAY_TOP_K = 3
SNIPPET_CHARS = 165
# Presentation-ready defaults (no LLM / no reranker unless you enable below).
USE_LLM = False
USE_RERANKER = False
USE_CROSS_ENCODER = bool(CROSS_ENCODER_PATH)
RETRIEVE_K = 10
EVAL_K = 10
def truncate_snippet(text: str, max_chars: int = SNIPPET_CHARS) -> str:
t = (text or "").replace("\n", " ").strip()
if len(t) <= max_chars:
return t
cut = t[:max_chars].rsplit(" ", 1)[0]
return cut + "…"
def print_ranked_list(header: str, results: list) -> None:
print(f"\n{header}")
for i, r in enumerate(results[:DISPLAY_TOP_K], start=1):
title = r.get("title", "") or ""
t_short = title[:75] + ("…" if len(title) > 75 else "")
snippet = truncate_snippet(r.get("text", ""))
print(f"{i}. [{r['source_id']}] {t_short} | skor={r['score']:.4f}")
print(f" {snippet}")
def extractive_answer_from_context(top_context: str, max_sentences: int = 2, max_chars: int = 420) -> str:
t = (top_context or "").strip()
if not t:
return "Bağlam yetersiz: ilgili metin bulunamadı."
parts = re.split(r"(?<=[.!?…])\s+", t, maxsplit=max_sentences)
out = " ".join(parts[:max_sentences]).strip()
if not out:
out = t
if len(out) > max_chars:
out = out[:max_chars].rsplit(" ", 1)[0] + "…"
return out
def is_poor_llm_output(text: str) -> bool:
t = (text or "").strip()
if len(t) < 15:
return True
low = t.lower()
if "insufficient" in low and "context" in low and len(t) < 120:
return True
return False
def final_answer(
question: str,
hybrid_results: list,
generator: LocalGenerator,
use_llm: bool,
) -> str:
contexts = [r["text"] for r in hybrid_results]
if not contexts:
return "Bağlam yetersiz: ilgili metin bulunamadı."
extractive = extractive_answer_from_context(contexts[0])
if not use_llm:
return extractive
try:
llm_out = generator.generate_answer(question, contexts)
except Exception:
return extractive
if is_poor_llm_output(llm_out):
return extractive
return llm_out.strip()
def hybrid_search(bm25: BM25Retriever, dense: DenseRetriever, question: str, k: int) -> list:
b = bm25.search(question, k=k)
d = dense.search(question, k=k)
fused = reciprocal_rank_fusion(b, d, top_n=k)
if USE_RERANKER:
fused = simple_lexical_rerank(fused, question)
if USE_CROSS_ENCODER:
fused = maybe_cross_encoder_rerank(fused, question, CROSS_ENCODER_PATH)
return fused
def hybrid_search_with_toggle(
bm25: BM25Retriever,
dense: DenseRetriever,
question: str,
k: int,
use_reranker: bool,
use_cross_encoder: bool,
) -> list:
b = bm25.search(question, k=k)
d = dense.search(question, k=k)
fused = reciprocal_rank_fusion(b, d, top_n=k)
if use_reranker:
fused = simple_lexical_rerank(fused, question)
if use_cross_encoder:
fused = maybe_cross_encoder_rerank(fused, question, CROSS_ENCODER_PATH)
return fused
def ensure_working_corpus(project_root: Path) -> Path:
"""
Prefer a held-out-safe retrieval index:
data/corpus_index.jsonl (HF ``test`` split excluded)
Fallback:
data/real_corpus.jsonl """
index_path = project_root / "data" / "corpus_index.jsonl"
real_path = project_root / "data" / "real_corpus.jsonl"
kaggle_dir = project_root / "data" / "kaggle_export"
if index_path.exists():
return index_path
try:
build_real_corpus(
output_path=index_path,
kaggle_dir=kaggle_dir if kaggle_dir.exists() else None,
include_hf=True,
hf_exclude_splits={"test"},
)
return index_path
except Exception as exc:
print(f"[Uyarı] Index corpus oluşturulamadı ({exc}). real_corpus denenecek.")
try:
build_real_corpus(
output_path=real_path,
kaggle_dir=kaggle_dir if kaggle_dir.exists() else None,
include_hf=True,
)
except Exception as exc:
print(f"[Uyarı] Birleşik corpus oluşturulamadı ({exc}).")
print(" Mevcut real_corpus.jsonl veya corpus.jsonl kullanılacak.")
return real_path
def filter_eval_by_corpus(eval_items: list, corpus_ids: set) -> list:
"""Skip items whose gold source_id(s) are not in the loaded corpus."""
kept = []
for item in eval_items:
gids = gold_doc_ids_from_eval_item(item)
if not gids:
continue
if not any(g in corpus_ids for g in gids):
continue
kept.append(item)
return kept
def load_eval_set(root: Path) -> list:
"""
Prefer external eval file if exists:
data/eval_qa_150.jsonl
Fallback to built-in mini set.
"""
eval_path = root / "data" / "eval_qa_150.jsonl"
if eval_path.exists():
try:
items = load_jsonl(eval_path)
if items:
return items
except Exception:
pass
return UNIFIED_EVAL_SET
# En az 10 soru: her biri soru + altın cevap + geçerli kaynak id (HF_* veya KG_*).
# KG_* maddeleri, Kaggle dışa aktarımı yoksa corpus'ta olmayacağı için güvenle atlanır.
UNIFIED_EVAL_SET: list = [
{
"question": "Türkiye'nin devlet şekli nedir?",
"gold_answer": "Anayasa madde 1'e göre, türkiye'nin devlet şekli cumhuriyettir.",
"source_id": "HF_train_0",
},
{
"question": "Türkiye'nin resmi dili nedir?",
"gold_answer": "Anayasa madde 3'e göre, türkiye'nin resmi dili türkçedir.",
"source_id": "HF_train_20",
},
{
"question": "Türkiye'nin başkenti neresidir?",
"gold_answer": "Anayasa madde 3, türkiye'nin başkentini ankara olarak belirler.",
"source_id": "HF_train_21",
},
{
"question": "Anayasa madde 36'ya göre adil yargılanma hakkı nasıl tanımlanır?",
"gold_answer": (
"Anayasa madde 36'ya göre, herkes meşru vasıta ve yollardan faydalanmak suretiyle "
"yargı mercileri önünde davacı veya davalı olarak iddia ve savunma ile adil yargılanma hakkına sahiptir."
),
"source_id": "HF_train_315",
},
{
"question": "Anayasa madde 38 masumiyet karinesini nasıl düzenler?",
"gold_answer": (
"Anayasa madde 38, masumiyet karinesini, suçluluğu hükmen sabit oluncaya kadar "
"kimsenin suçlu sayılamayacağı şeklinde düzenler."
),
"source_id": "HF_train_335",
},
{
"question": "Hırsızlık suçunun unsurları nelerdir?",
"gold_answer": "Hırsızlık suçunun unsurları, zilyedlik hakkının hukuka aykırı olarak elde edilmesi veya kullanılmasıdır.",
"source_id": "HF_train_5628",
},
{
"question": "Hırsızlık suçu nasıl işlenir?",
"gold_answer": (
"Hırsızlık suçu, bir malı hukuka aykırı olarak zilyedliğini elde etmek veya elinde bulundurmak "
"amacıyla yasada belirtilen şartlarla hareket edilmesidir."
),
"source_id": "HF_train_5616",
},
{
"question": "Anayasa madde 2'ye göre Türkiye nasıl bir devlettir?",
"gold_answer": (
"Anayasa madde 2'ye göre, türkiye cumhuriyeti'nin temel nitelikleri demokratik, laik, sosyal bir hukuk devleti olmasıdır."
),
"source_id": "HF_train_10",
},
{
"question": "Anayasa madde 1 cumhuriyetin hangi tarihte ilan edildiğini belirtir mi?",
"gold_answer": "Anayasa madde 1, türkiye cumhuriyeti'nin 29 ekim 1923 tarihinde ilan edildiğini belirtir.",
"source_id": "HF_train_3",
},
{
"question": "Türkiye'nin bayrağı nasıl tanımlanır?",
"gold_answer": "Anayasa madde 3'te türkiye'nin bayrağı, beyaz ay yıldızlı al bayrak olarak tanımlanır.",
"source_id": "HF_train_22",
},
{
"question": "Türk ceza kanununda kasten öldürme suçu nedir?",
"gold_answer": "Bir kişiyi öldürme kastıyla işlenen suçtur.",
"source_id": "HF_test_7",
},
{
"question": "Olağanüstü hallerde masumiyet karinesinin ihlali yasak mıdır?",
"gold_answer": "Anayasa madde 15, olağanüstü hallerde masumiyet karinesinin ihlal edilmesinin de yasak olduğunu belirtir.",
"source_id": "HF_train_134",
},
{
"question": "Anayasa madde 36'ya göre hak arama hürriyetinin sınırlandırılması nasıl düzenlenir?",
"gold_answer": (
"Anayasa madde 36, hak arama hürriyetinin sınırlandırılmasını, ancak kanunla ve demokratik toplum düzeninin "
"gereklerine uygun olarak yapılabileceğini düzenler."
),
"source_id": "HF_train_316",
},
]
def main() -> None:
root = Path(__file__).resolve().parent.parent
corpus_path = ensure_working_corpus(root)
if corpus_path.exists():
raw_records = load_jsonl(corpus_path)
else:
raw_records = load_corpus(root, prefer_real=True)
n_docs = len(raw_records)
docs = build_chunked_corpus(raw_records, chunk_size=220, overlap=40)
n_chunks = len(docs)
print("\n=== Corpus ===")
print(f"Dosya: {corpus_path}")
print(f"Kayıt sayısı (doküman): {n_docs}")
print(f"Parça sayısı (chunk): {n_chunks}")
if not docs:
print("[Hata] Ön işleme sonrası doküman yok. Veri dosyalarını kontrol edin.")
return
corpus_ids = {str(r.get("id", "")).strip() for r in raw_records if r.get("id")}
bm25 = BM25Retriever(docs)
dense = DenseRetriever(docs, model_name=DENSE_MODEL_PATH)
generator = LocalGenerator(model_name=GEN_MODEL_PATH)
sample_question = "Suçluluğu hükmen sabit olmadan kişi suçlu kabul edilir mi?"
bm25_results = bm25.search(sample_question, k=RETRIEVE_K)
dense_results = dense.search(sample_question, k=RETRIEVE_K)
hybrid_results = hybrid_search(bm25, dense, sample_question, RETRIEVE_K)
answer_text = final_answer(sample_question, hybrid_results, generator, USE_LLM)
print("\n=== Ayarlar ===")
print(f"Dense model: {DENSE_MODEL_PATH}")
print(f"Generator model: {GEN_MODEL_PATH}")
print(f"Cross-encoder: {'Açık (' + CROSS_ENCODER_PATH + ')' if USE_CROSS_ENCODER else 'Kapalı'}")
print(f"Yerel LLM (FLAN-T5): {'Açık' if USE_LLM else 'Kapalı (öntanımlı, çıkarımsal yanıt)'}")
print(f"Hibrit sonrası yeniden sıralama: {'Açık' if USE_RERANKER else 'Kapalı'}")
print("\n=== Soru ===")
print(sample_question)
print_ranked_list("=== BM25 Sonuçları ===", bm25_results)
print_ranked_list("=== Dense Sonuçları ===", dense_results)
print_ranked_list("=== Hybrid Sonuçları ===", hybrid_results)
print("\n=== Nihai Yanıt ===")
print(truncate_snippet(answer_text, max_chars=480))
full_eval_items = load_eval_set(root)
retrieval_eval_items = filter_eval_by_corpus(full_eval_items, corpus_ids)
skipped = len(full_eval_items) - len(retrieval_eval_items)
if skipped:
print(f"\n[Değerlendirme] Corpus'ta bulunmayan altın id nedeniyle atlanan örnek: {skipped}")
# Held-out mode: if all external eval ids are excluded from corpus_index, keep answer eval on full set
# and use the built-in aligned set for retrieval-side ranking metrics.
if not retrieval_eval_items and full_eval_items:
fallback_retrieval_eval = filter_eval_by_corpus(UNIFIED_EVAL_SET, corpus_ids)
if fallback_retrieval_eval:
retrieval_eval_items = fallback_retrieval_eval
print(
"\n[Not] Held-out eval id'leri index corpus'ta yok. "
"Retrieval metrikleri için hizalı mini set (UNIFIED_EVAL_SET) kullanıldı."
)
def bm25_fn(q, k):
return bm25.search(q, k=k)
def dense_fn(q, k):
return dense.search(q, k=k)
def hybrid_fn(q, k):
return hybrid_search(bm25, dense, q, k)
bm25_metrics = evaluate_retrieval(retrieval_eval_items, bm25_fn, max_k=EVAL_K)
dense_metrics = evaluate_retrieval(retrieval_eval_items, dense_fn, max_k=EVAL_K)
hybrid_metrics = evaluate_retrieval(retrieval_eval_items, hybrid_fn, max_k=EVAL_K)
print("\n=== Retrieval Değerlendirme ===")
print(
f"(Recall@5 / Recall@10 / MRR / nDCG@10, üst-{EVAL_K}; "
f"geçerli örnek: {len(retrieval_eval_items)})"
)
print("BM25 ->", bm25_metrics)
print("Dense ->", dense_metrics)
print("Hybrid->", hybrid_metrics)
answer_eval_items = [x for x in full_eval_items if str(x.get("gold_answer", "")).strip()]
def answer_fn(question: str) -> str:
h = hybrid_search(bm25, dense, question, RETRIEVE_K)
return final_answer(question, h, generator, USE_LLM)
answer_metrics = evaluate_answers(answer_eval_items, answer_fn)
print("\n=== Cevap Değerlendirme ===")
print(f"(EM / Token F1 / BLEU1 / ROUGE-L; geçerli örnek: {len(answer_eval_items)})")
print(answer_metrics)
def answer_and_context_fn(question: str) -> dict:
h = hybrid_search(bm25, dense, question, RETRIEVE_K)
return {
"answer": final_answer(question, h, generator, USE_LLM),
"contexts": [x.get("text", "") for x in h],
"retrieved_source_ids": [x.get("source_id", "") for x in h],
}
grounded_metrics = evaluate_groundedness(answer_eval_items, answer_and_context_fn)
print("\n=== Groundedness Değerlendirme ===")
print(f"(Faithfulness / CitationAccuracy; geçerli örnek: {len(answer_eval_items)})")
print(grounded_metrics)
def run_variant(name: str, use_reranker: bool, use_cross_encoder: bool, use_llm: bool) -> dict:
def retrieval_fn(q, k):
return hybrid_search_with_toggle(
bm25,
dense,
q,
k,
use_reranker=use_reranker,
use_cross_encoder=use_cross_encoder,
)
def local_answer_fn(q: str) -> str:
h_local = hybrid_search_with_toggle(
bm25,
dense,
q,
RETRIEVE_K,
use_reranker=use_reranker,
use_cross_encoder=use_cross_encoder,
)
return final_answer(q, h_local, generator, use_llm)
ret_m = evaluate_retrieval(retrieval_eval_items, retrieval_fn, max_k=EVAL_K)
ans_m = evaluate_answers(answer_eval_items, local_answer_fn)
return {
"name": name,
"Recall@5": ret_m["Recall@5"],
"Recall@10": ret_m["Recall@10"],
"MRR": ret_m["MRR"],
"nDCG@10": ret_m["nDCG@10"],
"TokenF1": ans_m["TokenF1"],
}
print("\n=== Ablation (Otomatik Karşılaştırma) ===")
variants = [
run_variant("Baseline(Hybrid, Extractive)", use_reranker=False, use_cross_encoder=False, use_llm=False),
run_variant("+LexicalReranker", use_reranker=True, use_cross_encoder=False, use_llm=False),
run_variant(
"+CrossEncoderReranker",
use_reranker=False,
use_cross_encoder=bool(CROSS_ENCODER_PATH),
use_llm=False,
),
run_variant(
"+Lexical+CrossEncoder",
use_reranker=True,
use_cross_encoder=bool(CROSS_ENCODER_PATH),
use_llm=False,
),
run_variant("+LLM", use_reranker=False, use_cross_encoder=False, use_llm=True),
run_variant("+Lexical+LLM", use_reranker=True, use_cross_encoder=False, use_llm=True),
run_variant(
"+CrossEncoder+LLM",
use_reranker=False,
use_cross_encoder=bool(CROSS_ENCODER_PATH),
use_llm=True,
),
run_variant(
"+Lexical+CrossEncoder+LLM",
use_reranker=True,
use_cross_encoder=bool(CROSS_ENCODER_PATH),
use_llm=True,
),
]
for v in variants:
print(
f"{v['name']}: "
f"R@5={v['Recall@5']:.3f}, "
f"R@10={v['Recall@10']:.3f}, "
f"MRR={v['MRR']:.3f}, "
f"nDCG@10={v['nDCG@10']:.3f}, "
f"TokenF1={v['TokenF1']:.3f}"
)
results_payload = {
"settings": {
"corpus_path": str(corpus_path),
"EMBED_MODEL_NAME": EMBED_MODEL_NAME,
"DENSE_MODEL_PATH": DENSE_MODEL_PATH,
"GEN_MODEL_NAME": GEN_MODEL_NAME,
"GEN_MODEL_PATH": GEN_MODEL_PATH,
"CROSS_ENCODER_PATH": CROSS_ENCODER_PATH,
"USE_LLM": USE_LLM,
"USE_RERANKER": USE_RERANKER,
"USE_CROSS_ENCODER": USE_CROSS_ENCODER,
"RETRIEVE_K": RETRIEVE_K,
"EVAL_K": EVAL_K,
},
"counts": {
"documents": n_docs,
"chunks": n_chunks,
"eval_total": len(full_eval_items),
"eval_valid": len(answer_eval_items),
"retrieval_eval_valid": len(retrieval_eval_items),
"eval_skipped": skipped,
},
"retrieval_metrics": {
"BM25": bm25_metrics,
"Dense": dense_metrics,
"Hybrid": hybrid_metrics,
},
"answer_metrics": answer_metrics,
"groundedness_metrics": grounded_metrics,
"ablation": variants,
}
out_path = root / "data" / "results_export.json"
with out_path.open("w", encoding="utf-8") as f:
json.dump(results_payload, f, ensure_ascii=False, indent=2)
print(f"\n[Çıktı] Sonuçlar kaydedildi: {out_path}")
print("\n=== Ek örnek sorular (kısa yanıt) ===")
extras = [
"Türkiye'nin milli marşı nedir?",
"Anayasa madde 2'de laik devlet ne demektir?",
"Zilyetliğin korunması nasıl sağlanır?",
]
for q in extras:
try:
h = hybrid_search(bm25, dense, q, RETRIEVE_K)
ans = final_answer(q, h, generator, USE_LLM)
print(f"\nS: {q}")
print(f"C: {truncate_snippet(ans, max_chars=380)}")
except Exception:
print(f"\nS: {q}")
print("C: (yanıt üretilemedi)")
if __name__ == "__main__":
try:
main()
except Exception:
print("[Hata] Demo çalışırken beklenmeyen bir sorun oluştu:")
traceback.print_exc()