Spaces:
Running on Zero
Running on Zero
| """LLM handler for Ollama, Groq, and HuggingFace Transformers providers.""" | |
| import logging | |
| import os | |
| from collections.abc import Iterator | |
| from typing import Any | |
| from langchain_core.callbacks import CallbackManagerForLLMRun | |
| from langchain_core.language_models.base import BaseLanguageModel | |
| from langchain_core.language_models.llms import LLM | |
| from pydantic import PrivateAttr | |
| from .config_loader import get_config | |
| logger = logging.getLogger(__name__) | |
| # Available Ollama models (common options) | |
| OLLAMA_MODELS = [ | |
| "llama3.2:3b", | |
| "llama3.2:1b", | |
| "llama3.1:8b", | |
| "phi3:mini", | |
| "gemma2:2b", | |
| "mistral:7b", | |
| "qwen2.5:3b", | |
| ] | |
| DEFAULT_OLLAMA_MODEL = "llama3.2:3b" | |
| # Available Groq models (free tier) | |
| GROQ_MODELS = [ | |
| "openai/gpt-oss-120b", | |
| "llama-3.3-70b-versatile", | |
| "llama-3.1-8b-instant", | |
| "mixtral-8x7b-32768", | |
| "gemma2-9b-it", | |
| ] | |
| DEFAULT_GROQ_MODEL = "openai/gpt-oss-120b" | |
| DEFAULT_HF_MODEL = "Qwen/Qwen3.5-4B" | |
| HF_MODELS = [ | |
| "Qwen/Qwen3.5-4B", | |
| "Qwen/Qwen2.5-3B-Instruct", | |
| "Qwen/Qwen2.5-1.5B-Instruct", | |
| "microsoft/Phi-3.5-mini-instruct", | |
| ] | |
| class HuggingFaceTransformersLLM(LLM): | |
| """LangChain-compatible wrapper around a local HuggingFace Causal LM.""" | |
| model_name: str = DEFAULT_HF_MODEL | |
| temperature: float = 0.7 | |
| max_new_tokens: int = 800 | |
| top_p: float = 0.9 | |
| _tokenizer: Any = PrivateAttr(default=None) | |
| _model: Any = PrivateAttr(default=None) | |
| def _llm_type(self) -> str: | |
| return "huggingface_transformers" | |
| def _ensure_loaded(self) -> None: | |
| """Lazily load tokenizer and model onto the best available device.""" | |
| if self._model is not None and self._tokenizer is not None: | |
| return | |
| import torch | |
| from transformers import AutoModelForCausalLM, AutoTokenizer | |
| logger.info(f"Loading HuggingFace model: {self.model_name}") | |
| self._tokenizer = AutoTokenizer.from_pretrained(self.model_name, trust_remote_code=True) | |
| if self._tokenizer.pad_token is None: | |
| self._tokenizer.pad_token = self._tokenizer.eos_token | |
| dtype = torch.bfloat16 if torch.cuda.is_available() else torch.float32 | |
| device_map = "auto" if torch.cuda.is_available() else None | |
| load_kwargs = { | |
| "trust_remote_code": True, | |
| "device_map": device_map, | |
| } | |
| # Prefer `dtype=` (torch_dtype is deprecated in recent transformers) | |
| model = None | |
| for dtype_key in ("dtype", "torch_dtype"): | |
| try: | |
| model = AutoModelForCausalLM.from_pretrained( | |
| self.model_name, | |
| **{dtype_key: dtype}, | |
| **load_kwargs, | |
| ) | |
| break | |
| except TypeError: | |
| continue | |
| except Exception as causal_err: | |
| logger.warning( | |
| "AutoModelForCausalLM failed for %s (%s); trying image-text class", | |
| self.model_name, | |
| causal_err, | |
| ) | |
| break | |
| if model is None: | |
| try: | |
| from transformers import AutoModelForImageTextToText | |
| model = AutoModelForImageTextToText.from_pretrained( | |
| self.model_name, | |
| dtype=dtype, | |
| **load_kwargs, | |
| ) | |
| except Exception: | |
| model = AutoModelForCausalLM.from_pretrained( | |
| self.model_name, | |
| torch_dtype=dtype, | |
| **load_kwargs, | |
| ) | |
| self._model = model | |
| if device_map is None: | |
| self._model = self._model.to("cpu") | |
| self._model.eval() | |
| logger.info(f"HuggingFace model loaded: {self.model_name}") | |
| def _build_messages(self, prompt: str) -> list[dict[str, str]]: | |
| """Use chat template when available; fall back to raw prompt.""" | |
| return [{"role": "user", "content": prompt}] | |
| def _tokenize_prompt(self, prompt: str) -> tuple[Any, Any]: | |
| """Tokenize prompt into (input_ids, attention_mask) tensors.""" | |
| import torch | |
| messages = self._build_messages(prompt) | |
| attention_mask = None | |
| enable_thinking = bool(get_config().get("llm.enable_thinking", False)) | |
| if hasattr(self._tokenizer, "apply_chat_template"): | |
| template_kwargs: dict[str, Any] = { | |
| "add_generation_prompt": True, | |
| "return_tensors": "pt", | |
| "return_dict": True, | |
| } | |
| # Qwen3 / Qwen3.5 support enable_thinking in chat template | |
| try: | |
| encoded = self._tokenizer.apply_chat_template( | |
| messages, | |
| enable_thinking=enable_thinking, | |
| **template_kwargs, | |
| ) | |
| except TypeError: | |
| try: | |
| encoded = self._tokenizer.apply_chat_template( | |
| messages, | |
| **template_kwargs, | |
| ) | |
| except TypeError: | |
| encoded = self._tokenizer.apply_chat_template( | |
| messages, | |
| add_generation_prompt=True, | |
| return_tensors="pt", | |
| ) | |
| if hasattr(encoded, "input_ids"): | |
| input_ids = encoded["input_ids"] | |
| attention_mask = encoded.get("attention_mask") | |
| elif isinstance(encoded, dict): | |
| input_ids = encoded["input_ids"] | |
| attention_mask = encoded.get("attention_mask") | |
| else: | |
| input_ids = encoded | |
| else: | |
| encoded = self._tokenizer(prompt, return_tensors="pt") | |
| input_ids = encoded["input_ids"] | |
| attention_mask = encoded.get("attention_mask") | |
| if not torch.is_tensor(input_ids): | |
| input_ids = torch.tensor(input_ids) | |
| if attention_mask is None: | |
| attention_mask = torch.ones_like(input_ids) | |
| elif not torch.is_tensor(attention_mask): | |
| attention_mask = torch.tensor(attention_mask) | |
| return input_ids, attention_mask | |
| def _generate_text(self, prompt: str) -> str: | |
| import torch | |
| self._ensure_loaded() | |
| assert self._tokenizer is not None and self._model is not None | |
| input_ids, attention_mask = self._tokenize_prompt(prompt) | |
| model_device = next(self._model.parameters()).device | |
| input_ids = input_ids.to(model_device) | |
| attention_mask = attention_mask.to(model_device) | |
| with torch.inference_mode(): | |
| output_ids = self._model.generate( | |
| input_ids=input_ids, | |
| attention_mask=attention_mask, | |
| max_new_tokens=self.max_new_tokens, | |
| temperature=max(self.temperature, 1e-5), | |
| top_p=self.top_p, | |
| do_sample=self.temperature > 0, | |
| pad_token_id=self._tokenizer.pad_token_id, | |
| eos_token_id=self._tokenizer.eos_token_id, | |
| ) | |
| generated = output_ids[0][input_ids.shape[-1] :] | |
| return self._tokenizer.decode(generated, skip_special_tokens=True).strip() | |
| def _call( | |
| self, | |
| prompt: str, | |
| stop: list[str] | None = None, | |
| run_manager: CallbackManagerForLLMRun | None = None, # noqa: ARG002 | |
| **kwargs: Any, # noqa: ARG002 | |
| ) -> str: | |
| text = self._generate_text(prompt) | |
| if stop: | |
| for token in stop: | |
| if token in text: | |
| text = text.split(token)[0] | |
| return text | |
| def _stream( | |
| self, | |
| prompt: str, | |
| stop: list[str] | None = None, | |
| run_manager: CallbackManagerForLLMRun | None = None, | |
| **kwargs: Any, | |
| ) -> Iterator[Any]: | |
| """Non-token streaming fallback: yield the full generation once.""" | |
| from langchain_core.outputs import GenerationChunk | |
| text = self._call(prompt, stop=stop, run_manager=run_manager, **kwargs) | |
| chunk = GenerationChunk(text=text) | |
| if run_manager: | |
| run_manager.on_llm_new_token(text, chunk=chunk) | |
| yield chunk | |
| class LLMHandler: | |
| """Handle LLM initialization and configuration.""" | |
| def __init__( | |
| self, | |
| provider_override: str | None = None, | |
| api_key_override: str | None = None, | |
| model_override: str | None = None, | |
| ): | |
| """Initialize LLM handler with configuration. | |
| Args: | |
| provider_override: Override provider from config (ollama/groq/transformers) | |
| api_key_override: Override API key from environment | |
| model_override: Override model from config | |
| """ | |
| self.config = get_config() | |
| self._api_key_override = api_key_override | |
| self._model_override = model_override | |
| self.llm: BaseLanguageModel | None = None | |
| # Auto-detect provider based on overrides / env / config | |
| if provider_override: | |
| self.provider = provider_override | |
| else: | |
| groq_api_key = self._get_groq_api_key() | |
| config_provider = self.config.get("llm.provider", "ollama") | |
| # Prefer explicit transformers (HF Spaces) over Groq auto-detect | |
| if config_provider == "transformers": | |
| self.provider = "transformers" | |
| elif groq_api_key: | |
| self.provider = "groq" | |
| logger.info("GROQ_API_KEY detected, using Groq as default provider") | |
| else: | |
| self.provider = config_provider | |
| def _get_groq_api_key(self) -> str | None: | |
| """Get Groq API key from override or environment.""" | |
| if self._api_key_override: | |
| return self._api_key_override | |
| return self.config.get_env("GROQ_API_KEY") | |
| def get_ollama_llm(self) -> BaseLanguageModel: | |
| """Initialize Ollama LLM. | |
| Returns: | |
| OllamaLLM instance | |
| """ | |
| from langchain_ollama import OllamaLLM | |
| model = self._model_override or self.config.get("llm.model", "llama3.2:3b") | |
| temperature = self.config.get("llm.temperature", 0.7) | |
| base_url = self.config.get_env("OLLAMA_BASE_URL", "http://localhost:11434") | |
| logger.info(f"Initializing Ollama with model: {model}") | |
| try: | |
| llm = OllamaLLM( | |
| model=model, | |
| temperature=temperature, | |
| base_url=base_url, | |
| num_predict=self.config.get("llm.max_tokens", 512), | |
| top_p=self.config.get("llm.top_p", 0.9), | |
| ) | |
| # Test the connection | |
| logger.debug("Testing Ollama connection...") | |
| test_response = llm.invoke("Hi") | |
| logger.debug(f"Ollama test response: {test_response[:50]}...") | |
| return llm | |
| except Exception as e: | |
| # Check for "model not found" or 404 error | |
| error_str = str(e).lower() | |
| if "not found" in error_str or "404" in error_str: | |
| logger.warning(f"Ollama model '{model}' not found.") | |
| # We raise a ValueError with a helpful message that app.py can display nicely | |
| msg = ( | |
| f"Model '{model}' not found in Ollama.\n" | |
| f"Please run this command in your terminal:\n" | |
| f"ollama pull {model}" | |
| ) | |
| raise ValueError(msg) from e | |
| logger.error(f"Error initializing Ollama: {e}") | |
| logger.error( | |
| f"Make sure Ollama is running and the model is pulled. Try: ollama pull {model}" | |
| ) | |
| # Re-raise the original exception if it's not a missing model issue | |
| raise | |
| def get_groq_llm(self) -> BaseLanguageModel: | |
| """Initialize Groq LLM. | |
| Returns: | |
| ChatGroq instance | |
| """ | |
| from langchain_groq import ChatGroq | |
| api_key = self._get_groq_api_key() | |
| if not api_key: | |
| msg = "GROQ_API_KEY not found. Set it in .env or provide via UI." | |
| raise ValueError(msg) | |
| model = self._model_override or self.config.get("llm.groq_model", DEFAULT_GROQ_MODEL) | |
| temperature = self.config.get("llm.temperature", 0.7) | |
| max_tokens = self.config.get("llm.max_tokens", 800) | |
| logger.info(f"Initializing Groq with model: {model}") | |
| try: | |
| llm = ChatGroq( | |
| api_key=api_key, | |
| model=model, | |
| temperature=temperature, | |
| max_tokens=max_tokens, | |
| ) | |
| # Test the connection | |
| logger.debug("Testing Groq connection...") | |
| test_response = llm.invoke("Hi") | |
| logger.debug(f"Groq test response: {str(test_response.content)[:50]}...") | |
| return llm | |
| except Exception as e: | |
| logger.error(f"Error initializing Groq: {e}") | |
| raise | |
| def get_transformers_llm(self) -> HuggingFaceTransformersLLM: | |
| """Initialize HuggingFace Transformers LLM (ZeroGPU / local). | |
| Returns: | |
| HuggingFaceTransformersLLM instance | |
| """ | |
| model = self._model_override or self.config.get("llm.hf_model", DEFAULT_HF_MODEL) | |
| temperature = float(self.config.get("llm.temperature", 0.7)) | |
| max_tokens = int(self.config.get("llm.max_tokens", 800)) | |
| top_p = float(self.config.get("llm.top_p", 0.9)) | |
| logger.info(f"Initializing Transformers LLM with model: {model}") | |
| return HuggingFaceTransformersLLM( | |
| model_name=model, | |
| temperature=temperature, | |
| max_new_tokens=max_tokens, | |
| top_p=top_p, | |
| ) | |
| def get_llm(self) -> BaseLanguageModel: | |
| """Get LLM instance based on configured provider. | |
| Returns: | |
| LLM instance | |
| """ | |
| if self.llm is not None: | |
| return self.llm | |
| if self.provider == "ollama": | |
| self.llm = self.get_ollama_llm() | |
| elif self.provider == "groq": | |
| self.llm = self.get_groq_llm() | |
| elif self.provider == "transformers": | |
| self.llm = self.get_transformers_llm() | |
| else: | |
| msg = f"Unsupported LLM provider: {self.provider}" | |
| raise ValueError(msg) | |
| return self.llm | |
| def get_system_prompt(self) -> str: | |
| """Get formatted system prompt from configuration. | |
| Returns: | |
| Formatted system prompt | |
| """ | |
| template = self.config.get( | |
| "llm.system_prompt", | |
| "You are a helpful AI assistant. Answer questions based on the provided context.", | |
| ) | |
| return self.config.format_template(template) | |
| def get_provider(self) -> str: | |
| """Get the current provider name.""" | |
| return self.provider | |
| def get_model(self) -> str: | |
| """Get the current model name.""" | |
| if self._model_override: | |
| return self._model_override | |
| if self.provider == "groq": | |
| return self.config.get("llm.groq_model", DEFAULT_GROQ_MODEL) | |
| if self.provider == "transformers": | |
| return self.config.get("llm.hf_model", DEFAULT_HF_MODEL) | |
| return self.config.get("llm.model", "llama3.2:3b") | |
| def get_llm_handler( | |
| provider_override: str | None = None, | |
| api_key_override: str | None = None, | |
| model_override: str | None = None, | |
| ) -> LLMHandler: | |
| """Get LLM handler instance with optional overrides.""" | |
| return LLMHandler( | |
| provider_override=provider_override, | |
| api_key_override=api_key_override, | |
| model_override=model_override, | |
| ) | |
| def get_available_groq_models() -> list[str]: | |
| """Get list of available Groq models.""" | |
| return GROQ_MODELS.copy() | |
| def get_default_groq_model() -> str: | |
| """Get default Groq model.""" | |
| return DEFAULT_GROQ_MODEL | |
| def get_available_hf_models() -> list[str]: | |
| """Get list of supported HuggingFace models.""" | |
| return HF_MODELS.copy() | |
| def get_default_hf_model() -> str: | |
| """Get default HuggingFace model for ZeroGPU.""" | |
| return DEFAULT_HF_MODEL | |
| def get_available_ollama_models() -> list[str]: | |
| """Get list of currently pulled Ollama models. | |
| Fetches the list from Ollama API. Falls back to static list if unavailable. | |
| """ | |
| try: | |
| import ollama | |
| config = get_config() | |
| base_url = config.get_env("OLLAMA_BASE_URL", "http://localhost:11434") | |
| # Create client with configured host | |
| client = ollama.Client(host=base_url) | |
| response = client.list() | |
| # Extract model names - response.models is a list of Model objects | |
| models = [] | |
| for model in response.models: | |
| # Each model has a .model attribute with the name | |
| name = getattr(model, "model", None) or getattr(model, "name", None) | |
| if name: | |
| models.append(name) | |
| if models: | |
| logger.debug(f"Found {len(models)} pulled Ollama models: {models}") | |
| return models | |
| # No models pulled - return empty list | |
| logger.warning("No Ollama models found. Please pull a model first.") | |
| return [] | |
| except Exception as e: | |
| logger.warning(f"Could not fetch Ollama models: {e}. Using static list.") | |
| return OLLAMA_MODELS.copy() | |
| def get_default_ollama_model() -> str: | |
| """Get default Ollama model.""" | |
| return DEFAULT_OLLAMA_MODEL | |
| def detect_default_provider() -> str: | |
| """Detect default provider based on environment. | |
| Loads .env file first to ensure environment variables are available. | |
| On Hugging Face Spaces → transformers (ZeroGPU). | |
| Locally → Groq if key present, otherwise Ollama (Streamlit UX). | |
| """ | |
| from pathlib import Path | |
| from dotenv import load_dotenv | |
| # Load .env file to ensure GROQ_API_KEY is available | |
| env_path = Path(".env") | |
| if env_path.exists(): | |
| load_dotenv(env_path) | |
| # SPACE_ID / SYSTEM=spaces are set on Hugging Face Spaces | |
| if os.getenv("SPACE_ID") or os.getenv("SYSTEM") == "spaces": | |
| return "transformers" | |
| if os.getenv("GROQ_API_KEY"): | |
| return "groq" | |
| return "ollama" | |