| 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}") |
| |
| |
| base_model_name = "MBZUAI-Paris/Atlas-Chat-9B" |
| |
| self.tokenizer = AutoTokenizer.from_pretrained( |
| base_model_name, |
| trust_remote_code=True |
| ) |
| |
| |
| if self.tokenizer.pad_token is None: |
| self.tokenizer.pad_token = self.tokenizer.eos_token |
| |
| |
| self.model = AutoModelForCausalLM.from_pretrained( |
| base_model_name, |
| torch_dtype=torch.float16, |
| device_map="auto", |
| trust_remote_code=True, |
| low_cpu_mem_usage=True |
| ) |
| |
| |
| 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""" |
| |
| |
| inputs = data.get("inputs", "") |
| parameters = data.get("parameters", {}) |
| |
| max_new_tokens = parameters.get("max_new_tokens", 300) |
| temperature = parameters.get("temperature", 0.7) |
| |
| |
| messages = [ |
| { |
| "role": "system", |
| "content": "أنت مساعد تجارة إلكترونية جزائري يتحدث الدارجة الجزائرية." |
| }, |
| {"role": "user", "content": inputs} |
| ] |
| |
| |
| prompt = self.tokenizer.apply_chat_template( |
| messages, |
| tokenize=False, |
| add_generation_prompt=True |
| ) |
| |
| |
| tokenized = self.tokenizer(prompt, return_tensors="pt").to(self.model.device) |
| prompt_length = tokenized["input_ids"].shape[1] |
| |
| |
| 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 |
| ) |
| |
| |
| generated_tokens = outputs[0][prompt_length:] |
| response = self.tokenizer.decode(generated_tokens, skip_special_tokens=True) |
| |
| |
| response = response.strip() |
| |
| return {"generated_text": response} |
|
|