"""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)