gemma-7b-chat / model_loader.py
yekkala's picture
Upload 6 files
c391424 verified
Raw
History Blame Contribute Delete
2.87 kB
"""
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)