Spaces:
Running
Running
Download model_loader.py from suraj-ml-projects/Text-GPT: direct link, hf CLI and curl.
- Browser
- Download file 2.15 kB
-
https://huggingface.co/spaces/suraj-ml-projects/Text-GPT/resolve/main/model_loader.py
- Command line
-
hf download hf://spaces/suraj-ml-projects/Text-GPT/model_loader.py
-
curl -L -o model_loader.py https://huggingface.co/spaces/suraj-ml-projects/Text-GPT/resolve/main/model_loader.py
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() |