import torch from typing import Dict, Any from transformers import AutoModelForCausalLM, AutoTokenizer from peft import PeftModel class EndpointHandler: def __init__(self, model_dir: str, **kwargs: Any) -> None: """Load base model + LoRA adapter for Atlas-Chat-9B""" print(f"Loading model from {model_dir}") # Load tokenizer from base model base_model_name = "MBZUAI-Paris/Atlas-Chat-9B" self.tokenizer = AutoTokenizer.from_pretrained( base_model_name, trust_remote_code=True ) # Set padding token if self.tokenizer.pad_token is None: self.tokenizer.pad_token = self.tokenizer.eos_token # Load base model in half precision self.model = AutoModelForCausalLM.from_pretrained( base_model_name, torch_dtype=torch.float16, device_map="auto", trust_remote_code=True, low_cpu_mem_usage=True ) # Load LoRA adapter self.model = PeftModel.from_pretrained(self.model, model_dir) self.model.eval() print("Model loaded successfully!") def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]: """Generate response""" # Get input inputs = data.get("inputs", "") parameters = data.get("parameters", {}) max_new_tokens = parameters.get("max_new_tokens", 300) temperature = parameters.get("temperature", 0.7) # Format message with chat template messages = [ { "role": "system", "content": "أنت مساعد تجارة إلكترونية جزائري يتحدث الدارجة الجزائرية." }, {"role": "user", "content": inputs} ] # Use the model's chat template prompt = self.tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True ) # Tokenize tokenized = self.tokenizer(prompt, return_tensors="pt").to(self.model.device) prompt_length = tokenized["input_ids"].shape[1] # Generate with torch.no_grad(): outputs = self.model.generate( **tokenized, max_new_tokens=max_new_tokens, temperature=temperature, do_sample=temperature > 0, top_p=0.95, repetition_penalty=1.1, pad_token_id=self.tokenizer.eos_token_id, eos_token_id=self.tokenizer.eos_token_id ) # Decode only the new tokens generated_tokens = outputs[0][prompt_length:] response = self.tokenizer.decode(generated_tokens, skip_special_tokens=True) # Clean up response response = response.strip() return {"generated_text": response}