Team Ai
Apppublic

anu151105/agentic-browser

sourceHugging Facemitupdated 1y agoView on Hugging Face
2likes
model_manager.py251 linesDownload Raw Back to models
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