Text-GPT / model_loader.py
suraj-ml-projects's picture
Update model_loader.py
0e85735 verified
Raw History Blame Contribute Delete
2.15 kB
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()