rag-lite-qa-system / generator.py
Ak47-model-ml's picture
Upload 9 files
c8c05d6 verified
Raw
History Blame Contribute Delete
18.2 kB
"""
generator.py -- Dual-backend generation for the RAG-lite QA system.
-
FLAN-T5:
1. Explicit rules for full sentences, context usage, and question-type-aware instructions.
2. max_new_tokens raised to 220 (from 150) for more complete answers.
3. no_repeat_ngram_size=3 retained to block repetition loops.
4. repetition_penalty=1.3 retained.
5. do_sample=True with temperature=0.3 added -- slight randomness
improves reasoning vs pure greedy decoding which always picks the
shortest high-probability token sequence.
6. Answer expansion layer: if FLAN returns fewer than 8 words, the
system reruns generation with an explicit expansion prompt to
produce a full sentence rather than a fragment.
7. Post-processing: capitalize first letter, strip trailing whitespace.
8. Context passed to FLAN is capped at top 3 chunks (not 7) to reduce
noise. FLAN is a small model and performs worse with long noisy inputs.
Groq improvements:
9. "Answer ONLY the given question" retained to prevent over-answering.
10. max_tokens=512 retained.
11. temperature=0.1 retained.
Output cleaning:
12. clean_answer() strips separator characters (===, ---) from any
generated output before returning it to app.py.
"""
import os
import re
import torch
from dotenv import load_dotenv
load_dotenv()
LOCAL_MODEL_NAME = "google/flan-t5-base"
GROQ_MODEL_NAME = os.environ.get("GROQ_MODEL", "llama-3.1-8b-instant")
# Cap context at top 3 chunks for FLAN. More chunks add noise that confuses
# the small model. Groq handles longer context fine so it uses all chunks.
FLAN_MAX_CHUNKS = 3
_local_pipeline = None
_FAILURE_PHRASES = {"not found", "not mentioned", "not stated", "unknown", "n/a", "none"}
_HEDGING_PHRASES = [
"i cannot", "i don't know", "not available", "not provided",
"cannot be found", "not mentioned", "not stated", "no information",
]
_LIST_TRIGGER_WORDS = {
"types", "type", "kinds", "kind", "categories", "category",
"applications", "application", "uses", "use", "examples", "example",
"techniques", "technique", "methods", "method", "components",
"challenges", "challenge", "features", "feature", "benefits",
"advantages", "disadvantages", "steps", "stages",
}
_COMPARE_TRIGGER_WORDS = {
"difference", "different", "differs", "compare", "comparison",
"versus", "vs", "unlike", "contrast", "distinguish",
}
# GROUNDING INSTRUCTION
GROUNDING_INSTRUCTION = (
"You are a precise and detail-oriented question-answering assistant. "
"Answer the question using only the information in the provided context. "
"Do not use any external knowledge or information from your training data. "
"You may combine information from multiple sentences within the context "
"to form a complete answer. "
"When asked about types, techniques, categories, applications, uses, or examples, "
"list all items explicitly mentioned in the context. Do not stop after one item. "
"Always write complete, grammatically correct sentences. "
"Do not copy incomplete phrases. Expand short phrases into full explanations. "
"Only say the answer is not found if the information is genuinely absent. "
"Do not add facts, inferences, or assumptions beyond what the context states."
)
# OUTPUT CLEANING
def clean_answer(text: str) -> str:
"""
Remove separator characters and normalize whitespace.
Handles both the === / --- dividers from format_multi_answer_output
and any stray markup the model might generate.
Also capitalizes the first letter and strips trailing whitespace.
"""
if not text:
return text
# Remove lines that are purely separator characters
lines = text.split("\n")
cleaned_lines = []
for line in lines:
stripped = line.strip()
# Drop lines made entirely of = or - characters (dividers)
if stripped and re.fullmatch(r"[=\-]{3,}", stripped):
continue
cleaned_lines.append(line)
text = "\n".join(cleaned_lines)
# Collapse 3+ consecutive newlines to 2
text = re.sub(r"\n{3,}", "\n\n", text)
text = text.strip()
# Capitalize first character
if text and text[0].islower():
text = text[0].upper() + text[1:]
return text
# LOCAL MODEL LOADER
def _get_local_pipeline():
global _local_pipeline
if _local_pipeline is None:
print(f"[INFO] Loading local model: {LOCAL_MODEL_NAME}")
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM, pipeline
device = 0 if torch.cuda.is_available() else -1
tokenizer = AutoTokenizer.from_pretrained(LOCAL_MODEL_NAME)
model = AutoModelForSeq2SeqLM.from_pretrained(LOCAL_MODEL_NAME)
_local_pipeline = pipeline(
"text2text-generation",
model=model,
tokenizer=tokenizer,
device=device,
)
print(f"[INFO] Local model loaded on {'GPU' if device == 0 else 'CPU'}")
return _local_pipeline
# CHUNK MERGING
def merge_chunks(chunks: list[dict], max_chunks: int = None) -> str:
"""
Merge retrieved chunks into one continuous context string.
If max_chunks is set, only the first N chunks (by similarity rank,
which is the order they arrive in) are used.
"""
if not chunks:
return ""
working = chunks[:max_chunks] if max_chunks else chunks
if len(working) == 1:
return working[0]["text"].strip()
ordered = sorted(working, key=lambda c: c.get("start_word", 0))
merged_words = ordered[0]["text"].split()
for chunk in ordered[1:]:
chunk_words = chunk["text"].split()
overlap_len = 0
max_check = min(len(merged_words), len(chunk_words), 40)
for n in range(max_check, 0, -1):
if merged_words[-n:] == chunk_words[:n]:
overlap_len = n
break
new_words = chunk_words[overlap_len:]
if new_words:
merged_words.extend(new_words)
return " ".join(merged_words)
# ------------------------------------------------------------------
# CONFIDENCE SCORING
# ------------------------------------------------------------------
def _score_local_confidence(answer: str, context: str) -> float:
"""Heuristic confidence for FLAN-T5 based on length and context overlap."""
if not answer or answer.lower().strip() in _FAILURE_PHRASES:
return 0.0
words = answer.lower().split()
word_count = len(words)
context_words = set(context.lower().split())
length_score = min(word_count / 20.0, 1.0)
overlap = sum(1 for w in words if w in context_words)
overlap_score = overlap / word_count if word_count > 0 else 0.0
return round(min(0.6 * length_score + 0.4 * overlap_score, 1.0), 3)
def _score_groq_confidence(answer: str) -> float:
"""Heuristic confidence for Groq based on hedging, length, sentence count."""
if not answer:
return 0.0
answer_lower = answer.lower()
hedge_penalty = sum(0.25 for p in _HEDGING_PHRASES if p in answer_lower)
hedge_score = max(0.0, 1.0 - hedge_penalty)
length_score = min(len(answer.split()) / 30.0, 1.0)
sent_count = answer.count(".") + answer.count("?") + answer.count("!")
sentence_score = min(sent_count / 3.0, 1.0)
return round(min(0.5 * hedge_score + 0.3 * length_score + 0.2 * sentence_score, 1.0), 3)
# ------------------------------------------------------------------
# PROMPT BUILDERS
# ------------------------------------------------------------------
def _question_type_instruction(question: str) -> str:
"""
Return a targeted instruction line based on what kind of answer is expected.
Three types handled:
- List/enumeration: "List all items clearly, separated by commas."
- Comparison: "Clearly explain both concepts and highlight the difference."
- Default: "Provide a clear and complete answer in 1-2 full sentences."
"""
q_lower = question.lower()
words = [w.rstrip("?.,;:!") for w in q_lower.split()]
if any(w in _LIST_TRIGGER_WORDS for w in words):
return "List all relevant items explicitly mentioned in the context, separated by commas."
if any(w in _COMPARE_TRIGGER_WORDS for w in words):
return "Clearly explain both concepts and highlight how they differ, using the context."
return "Provide a clear and complete answer in 1-2 full sentences using the context."
def _build_local_prompt(context: str, question: str) -> str:
"""
Prompt for flan-t5-base.
Design decisions:
- Context comes FIRST so FLAN sees it before the instructions. Small
seq2seq models attend more strongly to early tokens.
- Instructions are short, numbered, and concrete. Abstract rules like
"be thorough" do not work on FLAN. Specific rules like "do not give
fragment answers" do.
- No numbered format example (removed -- caused the 1...2...3 loop).
- Question-type instruction at the end as the final directive before Answer:.
"""
type_instruction = _question_type_instruction(question)
return (
f"Context:\n{context.strip()}\n\n"
f"Instructions:\n"
f"You are a precise and detail-oriented question-answering assistant.\n"
f"Rules:\n"
f"- Answer ONLY using the provided context above.\n"
f"- Always write a complete, meaningful sentence. Do NOT give fragment answers.\n"
f"- Do NOT copy incomplete phrases. Expand them into full explanations.\n"
f"- Include important details from the context.\n"
f"- Do NOT repeat the question in your answer.\n"
f"- Do NOT add external knowledge or hallucinate.\n"
f"- Use as much relevant information from the context as possible.\n"
f"- {type_instruction}\n\n"
f"Question: {question.strip()}\n\n"
f"Answer:"
)
def _build_expansion_prompt(short_answer: str, context: str, question: str) -> str:
"""
Expansion prompt used when FLAN returns a very short answer (under 8 words).
Asks the model to expand the fragment into a complete sentence.
"""
return (
f"Context:\n{context.strip()}\n\n"
f"Short answer: {short_answer.strip()}\n\n"
f"Task: Expand the short answer above into a complete, clear sentence "
f"using information from the context. "
f"The question was: {question.strip()}\n\n"
f"Do NOT add external knowledge. Use only the context.\n\n"
f"Complete answer:"
)
def _build_groq_messages(context: str, question: str) -> tuple[str, str]:
"""System and user messages for the Groq API."""
type_instruction = _question_type_instruction(question)
user = (
f"Context:\n{context.strip()}\n\n"
f"Question: {question.strip()}\n\n"
f"Important instructions:\n"
f"- Answer ONLY the given question. Do not include information for other questions.\n"
f"- {type_instruction}\n"
f"- Write complete sentences. Do not give fragment answers."
)
return GROUNDING_INSTRUCTION, user
# ------------------------------------------------------------------
# GENERATION: LOCAL FLAN-T5
# ------------------------------------------------------------------
def _run_flan(prompt: str) -> str:
"""
Run a single FLAN-T5 generation call and return the raw output string.
Centralizes all generation parameters in one place.
"""
gen = _get_local_pipeline()
result = gen(
prompt,
max_new_tokens=220,
do_sample=True,
temperature=0.3,
repetition_penalty=1.3,
no_repeat_ngram_size=3,
truncation=True,
)
return result[0]["generated_text"].strip()
def _generate_local(context: str, question: str) -> tuple[str, bool, float]:
"""
Generate answer with flan-t5-base.
Steps:
1. Run primary generation with the full structured prompt.
2. If the answer is fewer than 8 words (a fragment), run a second
generation pass with an expansion prompt to produce a full sentence.
3. Post-process: clean separators, capitalize first letter.
4. Score confidence.
"""
print(f"[INFO] Using backend: local ({LOCAL_MODEL_NAME})")
try:
prompt = _build_local_prompt(context, question)
answer = _run_flan(prompt)
if not answer or len(answer) < 3:
print("[WARN] Local model returned empty answer")
return "", False, 0.0
if answer.lower().strip() == question.lower().strip():
print("[WARN] Local model echoed the question")
return "", False, 0.0
if answer.lower().strip() in _FAILURE_PHRASES:
print(f"[WARN] Local model returned failure phrase: {answer}")
return "", False, 0.0
# Expansion pass: if answer is too short, try to expand it
if len(answer.split()) < 8:
print(f"[INFO] Answer too short ({len(answer.split())} words), running expansion")
expansion_prompt = _build_expansion_prompt(answer, context, question)
expanded = _run_flan(expansion_prompt)
# Only use expanded if it is actually longer and not a failure phrase
if (expanded
and len(expanded.split()) > len(answer.split())
and expanded.lower().strip() not in _FAILURE_PHRASES):
answer = expanded
print(f"[INFO] Expansion succeeded: {len(answer.split())} words")
# Post-processing: clean and capitalize
answer = clean_answer(answer)
confidence = _score_local_confidence(answer, context)
print(f"[INFO] Local answer | confidence: {confidence:.3f} | words: {len(answer.split())}")
return answer, True, confidence
except Exception as e:
print(f"[ERROR] Local generation failed: {e}")
return "", False, 0.0
# ------------------------------------------------------------------
# GENERATION: GROQ (with streaming)
# ------------------------------------------------------------------
def stream_groq_answer(context: str, question: str, api_key: str):
"""
Generator that streams Groq response tokens progressively.
Yields the cumulative answer string after each token chunk.
"""
try:
from groq import Groq
except ImportError:
yield "Groq package not installed. Run: pip install groq"
return
if not api_key:
yield "Groq API key not found. Set GROQ_API_KEY in your .env file."
return
try:
client = Groq(api_key=api_key, timeout=30.0)
system_msg, user_msg = _build_groq_messages(context, question)
stream = client.chat.completions.create(
model=GROQ_MODEL_NAME,
messages=[
{"role": "system", "content": system_msg},
{"role": "user", "content": user_msg},
],
temperature=0.1,
max_tokens=512,
stream=True,
)
accumulated = ""
for chunk in stream:
delta = chunk.choices[0].delta.content
if delta:
accumulated += delta
yield accumulated
except Exception as e:
err = str(e)
print(f"[ERROR] Groq streaming failed: {err}")
yield f"Groq error: {err}"
def _generate_groq(context: str, question: str, api_key: str) -> tuple[str, bool, float]:
"""Collect streamed Groq response and return complete answer with confidence."""
print(f"[INFO] Using backend: groq ({GROQ_MODEL_NAME})")
full_answer = ""
for partial in stream_groq_answer(context, question, api_key):
full_answer = partial
if not full_answer or len(full_answer) < 5:
print("[WARN] Groq returned empty answer")
return "", False, 0.0
if any(full_answer.startswith(p) for p in ("Groq error:", "Groq package", "Groq API key")):
print(f"[WARN] Groq backend error: {full_answer}")
return full_answer, False, 0.0
full_answer = clean_answer(full_answer)
confidence = _score_groq_confidence(full_answer)
print(f"[INFO] Groq answer | confidence: {confidence:.3f} | words: {len(full_answer.split())}")
return full_answer, True, confidence
# ------------------------------------------------------------------
# PUBLIC INTERFACE
# ------------------------------------------------------------------
def generate_answer(
context_chunks: list[dict],
question: str,
mode: str = "local",
groq_api_key: str = "",
) -> tuple[str, bool, float]:
"""
Main entry point called by app.py after retrieval.
For local mode: uses only top FLAN_MAX_CHUNKS (3) to reduce noise.
For groq mode: uses all retrieved chunks for maximum context.
Returns (answer_text, success_flag, confidence_score).
"""
if mode == "local":
# Limit to top 3 chunks for FLAN -- more chunks add noise
merged = merge_chunks(context_chunks, max_chunks=FLAN_MAX_CHUNKS)
else:
merged = merge_chunks(context_chunks)
if not merged.strip():
print("[WARN] merge_chunks returned empty string")
return "", False, 0.0
resolved_key = (groq_api_key or "").strip() or os.environ.get("GROQ_API_KEY", "")
if mode == "groq":
return _generate_groq(merged, question, resolved_key)
return _generate_local(merged, question)
def fallback_answer(best_chunk: str) -> str:
"""Return best chunk directly when generation fails entirely."""
print("[WARN] Falling back to raw chunk output")
preview = best_chunk.strip()[:500]
return f"Most relevant passage from the provided text:\n\n\"{preview}\""