anu151105/agentic-browser
2
1"""2Model Manager for handling local model loading and inference.3"""4import os5import sys6import torch7from typing import Dict, Any, Optional, List, Union8from transformers import (9 AutoModelForCausalLM,10 AutoTokenizer,11 AutoModel,12 pipeline,13 BitsAndBytesConfig14)15from sentence_transformers import SentenceTransformer16 17# Add parent directories to path18current_dir = os.path.dirname(os.path.abspath(__file__))19src_dir = os.path.dirname(current_dir)20root_dir = os.path.dirname(src_dir)21 22for path in [root_dir, src_dir]:23 if path not in sys.path:24 sys.path.insert(0, path)25 26try:27 from config.model_config import get_model_path, get_model_config, ModelConfig28except ImportError:29 # Fallback import paths30 try:31 from src.config.model_config import get_model_path, get_model_config, ModelConfig32 except ImportError:33 # Define minimal fallback config34 from dataclasses import dataclass35 from typing import Literal36 37 ModelType = Literal["text-generation", "text-embedding", "vision", "multimodal"]38 DeviceType = Literal["auto", "cpu", "cuda"]39 40 @dataclass41 class ModelConfig:42 model_id: str43 model_path: str44 model_type: ModelType45 device: DeviceType = "auto"46 quantize: bool = True47 use_safetensors: bool = True48 trust_remote_code: bool = True49 description: str = ""50 size_gb: float = 0.051 recommended: bool = False52 53 # Fallback model configurations54 DEFAULT_MODELS = {55 "tiny-llama": ModelConfig(56 model_id="TinyLlama/TinyLlama-1.1B-Chat-v1.0",57 model_path="./models/tiny-llama-1.1b-chat",58 model_type="text-generation",59 quantize=False, # Disable quantization for HF spaces60 description="Very small and fast model, good for quick testing",61 size_gb=1.1,62 recommended=True63 ),64 "mistral-7b": ModelConfig(65 model_id="microsoft/DialoGPT-small", # Use smaller model for HF spaces66 model_path="./models/dialogpt-small",67 model_type="text-generation",68 quantize=False,69 description="Small conversational model",70 size_gb=0.5,71 recommended=True72 )73 }74 75 def get_model_config(model_name: str) -> Optional[ModelConfig]:76 return DEFAULT_MODELS.get(model_name)77 78 def get_model_path(model_name: str) -> str:79 config = get_model_config(model_name)80 if not config:81 raise ValueError(f"Unknown model: {model_name}")82 83 # For HF spaces, use the model_id directly84 return config.model_id85 86class ModelManager:87 """Manages loading and using local models."""88 89 def __init__(self, device: str = None):90 self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")91 self.models: Dict[str, Any] = {}92 self.tokenizers: Dict[str, Any] = {}93 self.pipelines: Dict[str, Any] = {}94 95 def load_model(self, model_name: str, **kwargs) -> Any:96 """Load a model by name."""97 if model_name in self.models:98 return self.models[model_name]99 100 config = get_model_config(model_name)101 if not config:102 raise ValueError(f"Unknown model: {model_name}")103 104 model_path = get_model_path(model_name)105 106 # Load model based on type107 if config.model_type == "text-generation":108 model = self._load_text_generation_model(model_name, config, **kwargs)109 elif config.model_type == "text-embedding":110 model = self._load_embedding_model(model_name, config, **kwargs)111 else:112 raise ValueError(f"Unsupported model type: {config.model_type}")113 114 self.models[model_name] = model115 return model116 117 def _load_text_generation_model(self, model_name: str, config: ModelConfig, **kwargs):118 """Load a text generation model."""119 model_path = get_model_path(model_name)120 121 try:122 # Try to load with quantization if supported123 if config.quantize and "gptq" in config.model_id.lower():124 try:125 from auto_gptq import AutoGPTQForCausalLM126 model = AutoGPTQForCausalLM.from_quantized(127 model_path,128 device_map="auto",129 trust_remote_code=config.trust_remote_code,130 use_safetensors=config.use_safetensors,131 **kwargs132 )133 except ImportError:134 # Fallback to regular loading if auto_gptq is not available135 model = AutoModelForCausalLM.from_pretrained(136 model_path,137 device_map="auto" if torch.cuda.is_available() else "cpu",138 trust_remote_code=config.trust_remote_code,139 **kwargs140 )141 else:142 # Load full precision model143 model = AutoModelForCausalLM.from_pretrained(144 model_path,145 device_map="auto" if torch.cuda.is_available() else "cpu",146 trust_remote_code=config.trust_remote_code,147 torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32,148 **kwargs149 )150 except Exception as e:151 print(f"Error loading model {model_name}: {e}")152 # Fallback to CPU loading153 model = AutoModelForCausalLM.from_pretrained(154 model_path,155 device_map="cpu",156 trust_remote_code=config.trust_remote_code,157 torch_dtype=torch.float32,158 **kwargs159 )160 161 # Load tokenizer162 try:163 tokenizer = AutoTokenizer.from_pretrained(164 model_path,165 trust_remote_code=config.trust_remote_code166 )167 except Exception as e:168 print(f"Error loading tokenizer for {model_name}: {e}")169 # Use a basic tokenizer as fallback170 tokenizer = AutoTokenizer.from_pretrained("gpt2")171 172 self.tokenizers[model_name] = tokenizer173 174 return model175 176 def _load_embedding_model(self, model_name: str, config: ModelConfig, **kwargs):177 """Load a text embedding model."""178 model_path = get_model_path(model_name)179 180 # Use sentence-transformers for embedding models181 model = SentenceTransformer(182 model_path,183 device=self.device,184 **kwargs185 )186 187 return model188 189 def generate_text(190 self,191 model_name: str,192 prompt: str,193 max_length: int = 512,194 temperature: float = 0.7,195 **generation_kwargs196 ) -> str:197 """Generate text using the specified model."""198 if model_name not in self.models:199 self.load_model(model_name)200 201 model = self.models[model_name]202 tokenizer = self.tokenizers[model_name]203 204 # Encode the input205 inputs = tokenizer(prompt, return_tensors="pt").to(self.device)206 207 # Generate text208 with torch.no_grad():209 outputs = model.generate(210 **inputs,211 max_length=max_length,212 temperature=temperature,213 do_sample=True,214 pad_token_id=tokenizer.eos_token_id,215 **generation_kwargs216 )217 218 # Decode and return the generated text219 generated_text = tokenizer.decode(outputs[0], skip_special_tokens=True)220 return generated_text221 222 def get_embeddings(223 self,224 model_name: str,225 texts: Union[str, List[str]],226 batch_size: int = 32,227 **kwargs228 ) -> torch.Tensor:229 """Get embeddings for the input texts."""230 if model_name not in self.models:231 self.load_model(model_name)232 233 model = self.models[model_name]234 235 # Get embeddings using sentence-transformers236 if isinstance(texts, str):237 texts = [texts]238 239 embeddings = model.encode(240 texts,241 batch_size=batch_size,242 show_progress_bar=len(texts) > 1,243 convert_to_tensor=True,244 **kwargs245 )246 247 return embeddings248 249# Global model manager instance250model_manager = ModelManager()251 