import torch from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig import gc class ModelLoader: def __init__(self): self.text_model = None self.tokenizer = None self.device = "cuda" if torch.cuda.is_available() else "cpu" print(f"Using device: {self.device}") def load_text_model(self): """Safe Phi-3 model loading""" if self.text_model is None: model_id = "microsoft/Phi-3-mini-4k-instruct" # === Proper 4-bit Quantization Config === if self.device == "cuda": quantization_config = BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_compute_dtype=torch.float16, bnb_4bit_quant_type="nf4", bnb_4bit_use_double_quant=True, ) torch_dtype = torch.float16 else: # CPU fallback - no quantization quantization_config = None torch_dtype = torch.float32 print(f"Loading {model_id} with {'4-bit quantization' if quantization_config else 'CPU mode'}...") self.tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True) self.text_model = AutoModelForCausalLM.from_pretrained( model_id, quantization_config=quantization_config, # ← Correct way torch_dtype=torch_dtype, device_map="auto" if self.device == "cuda" else None, trust_remote_code=True, low_cpu_mem_usage=True, ) if self.device == "cpu": self.text_model = self.text_model.to(self.device) return self.text_model, self.tokenizer def clear_memory(self): if self.text_model is not None: del self.text_model self.text_model = None if self.tokenizer is not None: del self.tokenizer self.tokenizer = None gc.collect() if torch.cuda.is_available(): torch.cuda.empty_cache()