""" Alternative model loader with better error handling """ import torch from transformers import AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig import logging from pathlib import Path logger = logging.getLogger(__name__) class SimpleGemmaLoader: """Simple loader for Hugging Face Spaces""" def __init__(self, model_id="google/gemma-7b-it"): self.model_id = model_id self.model = None self.tokenizer = None self.loaded = False def load(self): """Load model with fallbacks""" try: # Try 4-bit quantization first logger.info("Attempting 4-bit quantization load...") bnb_config = BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_quant_type="nf4", bnb_4bit_compute_dtype=torch.float16, ) self.tokenizer = AutoTokenizer.from_pretrained(self.model_id) if self.tokenizer.pad_token is None: self.tokenizer.pad_token = self.tokenizer.eos_token self.model = AutoModelForCausalLM.from_pretrained( self.model_id, quantization_config=bnb_config, device_map="auto", trust_remote_code=True, low_cpu_mem_usage=True ) self.loaded = True logger.info("4-bit quantization successful") except Exception as e: logger.warning(f"4-bit failed: {e}. Trying fp16...") try: # Fallback to fp16 self.model = AutoModelForCausalLM.from_pretrained( self.model_id, torch_dtype=torch.float16, device_map="auto", low_cpu_mem_usage=True ) self.loaded = True logger.info("FP16 load successful") except Exception as e2: logger.error(f"All load attempts failed: {e2}") raise return self.model, self.tokenizer def generate(self, prompt, **kwargs): """Simple generation method""" if not self.loaded: self.load() inputs = self.tokenizer(prompt, return_tensors="pt") inputs = {k: v.to(self.model.device) for k, v in inputs.items()} with torch.no_grad(): outputs = self.model.generate( **inputs, max_new_tokens=kwargs.get('max_new_tokens', 512), temperature=kwargs.get('temperature', 0.7), top_p=kwargs.get('top_p', 0.95), do_sample=True, pad_token_id=self.tokenizer.pad_token_id, ) return self.tokenizer.decode(outputs[0], skip_special_tokens=True)