SwedenTaxRAGService / src /engine /llm_engine.py
gokceai's picture
test
43f15fb
Raw
History Blame Contribute Delete
8.97 kB
import logging
import re
import threading
import unicodedata
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from src.core.exceptions import InfrastructureError
logger = logging.getLogger(__name__)
WORD_RE = re.compile(r"[a-zA-Z0-9_åäöÅÄÖ]+")
class AnswerGenerator:
def __init__(self, settings):
self.settings = settings
self.model_path = settings.LLM_MODEL_PATH
self.device = settings.resolve_device(settings.LLM_DEVICE)
self.dtype = torch.float16 if self.device == "cuda" else torch.float32
self.use_int8 = settings.LLM_USE_INT8
if self.device != "cuda":
logger.warning(
"No GPU detected - LLM will run on CPU (dtype=%s, slower). Model: %s",
self.dtype,
self.model_path,
)
self.tokenizer = None
self.model = None
self._load_lock = threading.Lock()
self._load_error: Exception | None = None
@property
def is_ready(self) -> bool:
return self.model is not None and self.tokenizer is not None
@property
def has_error(self) -> bool:
return self._load_error is not None
def load(self) -> None:
self._ensure_model_loaded()
def _ensure_model_loaded(self) -> None:
if self.is_ready:
return
with self._load_lock:
if self.is_ready:
return
try:
self.tokenizer = AutoTokenizer.from_pretrained(self.model_path)
if self.tokenizer.pad_token is None:
self.tokenizer.pad_token = self.tokenizer.eos_token
self.model = self._load_model()
if self.device == "cpu":
self.model.to(self.device)
self.model.eval()
logger.info("LLM '%s' loaded on %s.", self.model_path, self.device.upper())
except Exception as e:
self._load_error = e
raise InfrastructureError(f"LLM load failed: {e}") from e
def _load_model(self):
if self.device == "cuda" and self.use_int8:
try:
from transformers import BitsAndBytesConfig # noqa: PLC0415
logger.info("Loading LLM with GPU 8-bit quantization (bitsandbytes).")
return AutoModelForCausalLM.from_pretrained(
self.model_path,
quantization_config=BitsAndBytesConfig(load_in_8bit=True),
low_cpu_mem_usage=True,
device_map="auto",
)
except Exception as exc: # noqa: BLE001
logger.warning("INT8 GPU quantization unavailable; falling back to dtype load: %s", exc)
model = AutoModelForCausalLM.from_pretrained(
self.model_path,
dtype=self.dtype,
low_cpu_mem_usage=True,
device_map="auto" if self.device == "cuda" else None,
)
if self.device == "cpu" and self.use_int8:
try:
logger.info("Applying dynamic INT8 quantization on CPU.")
model = torch.quantization.quantize_dynamic(
model,
{torch.nn.Linear},
dtype=torch.qint8,
)
except Exception as exc: # noqa: BLE001
logger.warning("CPU INT8 quantization failed; using non-quantized model: %s", exc)
return model
def _normalize(self, text: str) -> str:
text = unicodedata.normalize("NFC", text or "")
return text.strip().lower()
def _tokenize(self, text: str) -> list[str]:
return WORD_RE.findall(self._normalize(text))
def _token_set(self, text: str) -> set[str]:
stopwords = {
"the", "a", "an", "and", "or", "to", "of", "in", "on", "for", "with", "is",
"are", "be", "as", "by", "that", "this", "it", "you", "your", "from", "at",
"if", "not", "do", "does", "can", "will", "must", "have", "has", "had",
"i", "we", "they", "them", "their", "our", "about",
}
return {tok for tok in self._tokenize(text) if len(tok) >= 3 and tok not in stopwords}
def _contains_generic_bad_pattern(self, answer: str) -> bool:
answer_l = self._normalize(answer)
bad_patterns = [
"it seems like you're trying to ask something about the weather",
"could you please rephrase your question",
"i don't know about any specific answer to your question",
"there's a lot of irrelevant information here",
"if you'd like to ask a different question",
"the format isn't clear",
]
return any(p in answer_l for p in bad_patterns)
def _domain_mismatch(self, query: str, answer: str) -> bool:
query_l = self._normalize(query)
answer_l = self._normalize(answer)
tax_terms = {
"tax", "vat", "skatteverket", "swedish tax agency", "f-tax",
"income tax", "corporate tax", "excise", "return", "declaration",
}
off_domain_terms = {
"weather", "forecast", "temperature", "rain", "sunny", "humidity",
}
query_is_tax = any(term in query_l for term in tax_terms)
answer_mentions_off_domain = any(term in answer_l for term in off_domain_terms)
answer_mentions_tax = any(term in answer_l for term in tax_terms)
return query_is_tax and answer_mentions_off_domain and not answer_mentions_tax
def _lexical_grounding_too_weak(self, query: str, answer: str, contexts: list[str]) -> bool:
query_tokens = self._token_set(query)
answer_tokens = self._token_set(answer)
context_tokens = self._token_set("\n".join(contexts))
if not answer_tokens or not context_tokens:
return True
answer_vs_context_overlap = len(answer_tokens & context_tokens) / max(len(answer_tokens), 1)
query_vs_answer_overlap = len(query_tokens & answer_tokens) / max(len(query_tokens), 1) if query_tokens else 1.0
logger.info(
"answer_sanity overlap answer_context=%.3f query_answer=%.3f",
answer_vs_context_overlap,
query_vs_answer_overlap,
)
return (
answer_vs_context_overlap < self.settings.MIN_ANSWER_CONTEXT_OVERLAP
or query_vs_answer_overlap < self.settings.MIN_QUERY_ANSWER_OVERLAP
)
def _should_reject_answer(self, query: str, answer: str, contexts: list[str]) -> bool:
answer_clean = (answer or "").strip()
if not answer_clean:
logger.info("answer_sanity rejected: empty")
return True
if len(answer_clean) < self.settings.MIN_ANSWER_CHARS:
logger.info("answer_sanity rejected: too_short chars=%s", len(answer_clean))
return True
if self._contains_generic_bad_pattern(answer_clean):
logger.info("answer_sanity rejected: generic_bad_pattern")
return True
if self._domain_mismatch(query, answer_clean):
logger.info("answer_sanity rejected: domain_mismatch")
return True
if self._lexical_grounding_too_weak(query, answer_clean, contexts):
logger.info("answer_sanity rejected: weak_grounding")
return True
return False
def generate_answer(self, query: str, contexts: list[str]) -> str:
if not contexts:
return self.settings.WARNING_PROMPT
if not self.is_ready:
self._ensure_model_loaded()
trimmed_contexts = contexts[: self.settings.MAX_CONTEXT_CHUNKS]
combined_context = "\n\n".join(trimmed_contexts)
if len(combined_context) > self.settings.MAX_CONTEXT_CHARS:
combined_context = combined_context[: self.settings.MAX_CONTEXT_CHARS]
messages = [
{
"role": "system",
"content": self.settings.SYSTEM_PROMPT_CONTENT.format(
context=combined_context,
),
},
{
"role": "user",
"content": f"QUESTION: {query}",
},
]
inputs = self.tokenizer.apply_chat_template(
messages,
add_generation_prompt=True,
return_tensors="pt",
return_dict=True,
).to(self.device)
with torch.inference_mode():
outputs = self.model.generate(
**inputs,
max_new_tokens=self.settings.LLM_MAX_NEW_TOKENS,
do_sample=False,
pad_token_id=self.tokenizer.eos_token_id,
use_cache=True,
repetition_penalty=1.05,
)
input_length = inputs["input_ids"].shape[-1]
return self.tokenizer.decode(
outputs[0][input_length:], skip_special_tokens=True
).strip()