finpy1789's picture
Upload folder using huggingface_hub
ef0b208 verified
Raw
History Blame Contribute Delete
4.37 kB
"""Inference engine: loads a base model (4-bit on GPU) plus its finance LoRA
adapter, and hot-swaps to a different base+adapter when the user changes model.
Only one base model is kept in memory at a time (Space GPUs aren't big enough
for the whole catalog); switching models unloads the previous one.
"""
import gc
import threading
import torch
from peft import PeftModel
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
class InferenceEngine:
def __init__(self):
self._lock = threading.Lock()
self.model = None
self.tokenizer = None
self.current_base = None
self.adapter_loaded = False
self.has_gpu = False # decided at load time (ZeroGPU: CUDA only exists inside the GPU call)
def load(self, base_model, adapter=None):
"""Load base_model (+adapter if it exists). Returns a status string.
Must be called where CUDA is actually usable — on ZeroGPU that means
inside the @spaces.GPU-decorated function.
"""
with self._lock:
if self.current_base == base_model:
return self._status(base_model)
self._unload()
self.has_gpu = torch.cuda.is_available()
kwargs = {"device_map": "auto"} if self.has_gpu else {}
if self.has_gpu:
dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
kwargs["dtype"] = dtype
kwargs["quantization_config"] = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_use_double_quant=True,
bnb_4bit_compute_dtype=dtype,
)
else:
kwargs["dtype"] = torch.float32
self.tokenizer = AutoTokenizer.from_pretrained(base_model)
try:
self.model = AutoModelForCausalLM.from_pretrained(base_model, **kwargs)
except Exception as e: # noqa: BLE001
if "quantization_config" not in kwargs:
raise
# bitsandbytes can lag behind brand-new GPU architectures;
# retry unquantized (bf16 fits up to ~27B on large-VRAM slices).
print(f"[warn] 4-bit load failed ({type(e).__name__}: {e}); retrying without quantization")
kwargs.pop("quantization_config")
self.model = AutoModelForCausalLM.from_pretrained(base_model, **kwargs)
self.adapter_loaded = False
if adapter:
try:
self.model = PeftModel.from_pretrained(self.model, adapter)
self.adapter_loaded = True
except Exception as e: # noqa: BLE001 - adapter repo may not exist yet
print(f"[warn] adapter {adapter} unavailable ({e}); serving base model")
self.model.eval()
self.current_base = base_model
return self._status(base_model)
def _status(self, base_model):
tag = "finance adapter active" if self.adapter_loaded else "base model — adapter pending"
return f"{base_model} ({tag})"
def _unload(self):
if self.model is not None:
del self.model
self.model = None
gc.collect()
if self.has_gpu:
torch.cuda.empty_cache()
@torch.inference_mode()
def chat(self, messages, max_new_tokens=512, temperature=0.7):
"""messages: list of {"role": ..., "content": ...}. Returns the reply text."""
with self._lock:
if self.model is None:
raise RuntimeError("No model loaded")
inputs = self.tokenizer.apply_chat_template(
messages, add_generation_prompt=True, return_tensors="pt", return_dict=True
).to(self.model.device)
out = self.model.generate(
**inputs,
max_new_tokens=max_new_tokens,
temperature=temperature,
do_sample=temperature > 0,
top_p=0.9,
pad_token_id=self.tokenizer.pad_token_id or self.tokenizer.eos_token_id,
)
prompt_len = inputs["input_ids"].shape[1]
return self.tokenizer.decode(out[0][prompt_len:], skip_special_tokens=True)