ai-study-assistant / core /llm_engine.py
MonishRaman's picture
Upload 31 files
af25a2a verified
Raw
History Blame Contribute Delete
10 kB
"""
LLM Engine - Handles API calls to OpenAI and Hugging Face
"""
from typing import Dict, Optional, Tuple
import json
import time
from config import (
OPENAI_API_KEY,
LLM_PROVIDER,
OPENAI_MODEL,
HF_MODEL,
TEMPERATURE,
TOP_P,
)
from core.utils import log_event
class LLMEngine:
"""Base class for LLM interactions."""
def __init__(self, provider: str = LLM_PROVIDER):
"""
Initialize LLM Engine.
Args:
provider: "openai" or "huggingface"
"""
self.provider = provider.lower()
self.model = OPENAI_MODEL if self.provider == "openai" else HF_MODEL
if self.provider == "openai":
self.engine = OpenAIEngine()
elif self.provider == "huggingface":
self.engine = HuggingFaceEngine()
else:
raise ValueError(f"Unknown provider: {provider}")
def generate(self, prompt: str, max_tokens: int = 1000) -> Tuple[bool, str]:
"""
Delegate generation to the selected engine.
"""
return self.engine.generate(prompt, max_tokens)
def stream_generate(self, prompt: str, max_tokens: int = 1000):
"""
Stream generation for real-time responses.
"""
yield from self.engine.stream_generate(prompt, max_tokens)
class OpenAIEngine:
"""OpenAI API integration."""
def __init__(self):
"""Initialize OpenAI engine."""
try:
import openai
if not OPENAI_API_KEY:
raise ValueError(
"OPENAI_API_KEY is missing. Set it in your .env file."
)
openai.api_key = OPENAI_API_KEY
self.client = openai.OpenAI(api_key=OPENAI_API_KEY)
self.model = OPENAI_MODEL
except ImportError:
raise ImportError("openai package not installed. Install with: pip install openai")
def generate(self, prompt: str, max_tokens: int = 1000) -> Tuple[bool, str]:
"""
Generate text using OpenAI API.
Args:
prompt: Input prompt
max_tokens: Max tokens in response
Returns:
Tuple of (success, response_text)
"""
try:
log_event("API_CALL", f"OpenAI - Model: {self.model}, Tokens: {max_tokens}")
response = self.client.chat.completions.create(
model=self.model,
messages=[
{
"role": "system",
"content": "You are a helpful AI study assistant. Provide clear, accurate, and educational responses."
},
{
"role": "user",
"content": prompt
}
],
max_tokens=max_tokens,
temperature=TEMPERATURE,
top_p=TOP_P,
)
result = response.choices[0].message.content
log_event("API_SUCCESS", "OpenAI response received")
return True, result
except Exception as e:
error_msg = str(e)
log_event("API_ERROR", f"OpenAI: {error_msg}")
if "insufficient_quota" in error_msg or "Error code: 429" in error_msg:
return (
False,
"OpenAI API Error: Insufficient quota (429). "
"Add billing/credits in your OpenAI account, or switch provider by setting "
"LLM_PROVIDER=huggingface in .env to use the local model."
)
return False, f"OpenAI API Error: {error_msg}"
def stream_generate(self, prompt: str, max_tokens: int = 1000):
"""
Stream responses from OpenAI API.
Args:
prompt: Input prompt
max_tokens: Max tokens
Yields:
Response chunks
"""
try:
response = self.client.chat.completions.create(
model=self.model,
messages=[
{"role": "system", "content": "You are a helpful AI study assistant."},
{"role": "user", "content": prompt}
],
max_tokens=max_tokens,
temperature=TEMPERATURE,
stream=True,
)
for chunk in response:
if chunk.choices[0].delta.content:
yield chunk.choices[0].delta.content
except Exception as e:
yield f"Error: {str(e)}"
class HuggingFaceEngine:
"""Local Hugging Face model (no API)."""
_shared_tokenizer = None
_shared_model_obj = None
_shared_model_name = None
_shared_device = None
def __init__(self):
"""Initialize local Hugging Face pipeline."""
try:
import torch
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
self.model = "google/flan-t5-base" if HF_MODEL == "local" else HF_MODEL
if (
HuggingFaceEngine._shared_model_obj is None
or HuggingFaceEngine._shared_tokenizer is None
or HuggingFaceEngine._shared_model_name != self.model
):
device = "cuda" if torch.cuda.is_available() else "cpu"
tokenizer = AutoTokenizer.from_pretrained(self.model)
model_obj = AutoModelForSeq2SeqLM.from_pretrained(self.model).to(device)
HuggingFaceEngine._shared_tokenizer = tokenizer
HuggingFaceEngine._shared_model_obj = model_obj
HuggingFaceEngine._shared_model_name = self.model
HuggingFaceEngine._shared_device = device
self.device = HuggingFaceEngine._shared_device
self.tokenizer = HuggingFaceEngine._shared_tokenizer
self.model_obj = HuggingFaceEngine._shared_model_obj
except ImportError:
raise ImportError(
"transformers/torch/sentencepiece not installed. "
"Install with: pip install transformers torch sentencepiece"
)
def generate(self, prompt: str, max_tokens: int = 1000) -> Tuple[bool, str]:
try:
log_event("API_CALL", f"Local HF - {self.model}")
inputs = self.tokenizer(
prompt,
return_tensors="pt",
truncation=True,
max_length=1024,
).to(self.device)
generation_max_length = 200
generation_min_length = 80
import torch
with torch.no_grad():
outputs = self.model_obj.generate(
**inputs,
max_length=generation_max_length,
min_length=generation_min_length,
do_sample=True,
temperature=0.6,
top_p=0.9,
repetition_penalty=1.3,
no_repeat_ngram_size=3,
)
decoded_text = self.tokenizer.decode(outputs[0], skip_special_tokens=True)
response = [{"generated_text": decoded_text}]
result = response[0]["generated_text"]
if not isinstance(result, str):
result = str(result)
result = result.strip()
if not result:
log_event("API_ERROR", "LocalHF: Empty response generated")
return False, "Local Model Error: Empty response generated"
log_event("API_SUCCESS", "Local model response received")
return True, result
except Exception as e:
error_msg = str(e)
log_event("API_ERROR", f"LocalHF: {error_msg}")
print("FULL ERROR:", error_msg)
return False, f"Local Model Error: {error_msg}"
def stream_generate(self, prompt: str, max_tokens: int = 1000):
"""Yield single complete response for compatibility with streaming API."""
success, response = self.generate(prompt, max_tokens)
yield response
class SimpleCache:
"""Simple in-memory cache for responses."""
def __init__(self, max_size: int = 100):
"""Initialize cache."""
self.cache = {}
self.max_size = max_size
def get(self, key: str) -> Optional[str]:
"""Get cached response."""
return self.cache.get(key)
def set(self, key: str, value: str) -> None:
"""Cache a response."""
if len(self.cache) >= self.max_size:
# Remove oldest entry (FIFO)
self.cache.pop(next(iter(self.cache)))
self.cache[key] = value
def clear(self) -> None:
"""Clear cache."""
self.cache.clear()
# Global cache instance
_cache = SimpleCache()
def get_or_generate(prompt: str, engine: LLMEngine, max_tokens: int = 1000) -> Tuple[bool, str]:
"""
Get cached response or generate new one.
Args:
prompt: Input prompt
engine: LLM engine to use
max_tokens: Max tokens
Returns:
Tuple of (success, response)
"""
# Create cache key from prompt
cache_key = hash(prompt) % ((2 ** 63) - 1)
# Check cache
cached = _cache.get(str(cache_key))
if cached:
log_event("CACHE_HIT", "Using cached response")
return True, cached
# Generate if not cached
success, response = engine.generate(prompt, max_tokens)
if success:
_cache.set(str(cache_key), response)
return success, response