dizza01/BioMistral-7B-DARE
029
1import os2import json3import torch4from transformers import AutoModelForCausalLM, AutoTokenizer5from peft import AutoPeftModelForCausalLM6 7DEFAULT_SYSTEM_PROMPT = (8 "You are a QA assistant. "9 "Use only the provided context. "10 "If the answer is not present in the context, say so clearly."11)12 13class EndpointHandler:14 def __init__(self, path: str = ""):15 model_dir = path or "/repository"16 17 self.tokenizer = AutoTokenizer.from_pretrained(18 model_dir,19 trust_remote_code=True,20 )21 22 if self.tokenizer.pad_token_id is None:23 self.tokenizer.pad_token = self.tokenizer.eos_token24 25 dtype = torch.float16 if torch.cuda.is_available() else torch.float3226 27 adapter_config_path = os.path.join(model_dir, "adapter_config.json")28 if os.path.exists(adapter_config_path):29 self.model = AutoPeftModelForCausalLM.from_pretrained(30 model_dir,31 trust_remote_code=True,32 torch_dtype=dtype,33 low_cpu_mem_usage=True,34 device_map="auto" if torch.cuda.is_available() else None,35 )36 else:37 self.model = AutoModelForCausalLM.from_pretrained(38 model_dir,39 trust_remote_code=True,40 torch_dtype=dtype,41 low_cpu_mem_usage=True,42 device_map="auto" if torch.cuda.is_available() else None,43 )44 45 self.model.eval()46 47 def _build_messages(self, inputs):48 if isinstance(inputs, list):49 messages = inputs50 elif isinstance(inputs, dict) and "context" in inputs and "question" in inputs:51 messages = [52 {"role": "system", "content": DEFAULT_SYSTEM_PROMPT},53 {54 "role": "user",55 "content": f"Context:\n{inputs['context']}\n\nQuestion: {inputs['question']}",56 },57 ]58 else:59 messages = [60 {"role": "system", "content": DEFAULT_SYSTEM_PROMPT},61 {"role": "user", "content": str(inputs)},62 ]63 64 has_system = any(message.get("role") == "system" for message in messages)65 if not has_system:66 messages = [{"role": "system", "content": DEFAULT_SYSTEM_PROMPT}] + messages67 68 return messages69 70 def __call__(self, data):71 inputs = data.get("inputs", "")72 params = data.get("parameters", {}) or {}73 74 max_new_tokens = min(int(params.get("max_new_tokens", 128)), 512)75 temperature = float(params.get("temperature", 0.0))76 top_p = float(params.get("top_p", 1.0))77 do_sample = bool(params.get("do_sample", False))78 repetition_penalty = float(params.get("repetition_penalty", 1.0))79 no_repeat_ngram_size = int(params.get("no_repeat_ngram_size", 0))80 debug = bool(params.get("debug", False))81 82 messages = self._build_messages(inputs)83 84 prompt = self.tokenizer.apply_chat_template(85 messages,86 tokenize=False,87 add_generation_prompt=True,88 )89 90 enc = self.tokenizer(91 prompt,92 return_tensors="pt",93 truncation=True,94 max_length=min(getattr(self.tokenizer, "model_max_length", 4096), 4096),95 )96 97 if torch.cuda.is_available():98 enc = {key: value.to(self.model.device) for key, value in enc.items()}99 100 generate_kwargs = dict(101 **enc,102 max_new_tokens=max_new_tokens,103 do_sample=do_sample,104 repetition_penalty=repetition_penalty,105 pad_token_id=self.tokenizer.pad_token_id,106 eos_token_id=self.tokenizer.eos_token_id,107 )108 109 if do_sample:110 generate_kwargs["temperature"] = max(temperature, 1e-5)111 generate_kwargs["top_p"] = top_p112 113 if no_repeat_ngram_size > 0:114 generate_kwargs["no_repeat_ngram_size"] = no_repeat_ngram_size115 116 with torch.no_grad():117 out = self.model.generate(**generate_kwargs)118 119 generated_ids = out[0][enc["input_ids"].shape[-1]:]120 text = self.tokenizer.decode(generated_ids, skip_special_tokens=True).strip()121 122 response = {"generated_text": text}123 if debug:124 response["prompt"] = prompt125 response["messages"] = messages126 return response