| import os |
| import torch |
| from transformers import AutoModelForCausalLM, AutoTokenizer, StoppingCriteria, StoppingCriteriaList |
| from typing import Dict, List, Any |
|
|
| |
| |
| class StopOnPatientToken(StoppingCriteria): |
| def __init__(self, tokenizer, stop_token_str="Patient:"): |
| super().__init__() |
| |
| |
| |
| |
| stop_token_ids_list = tokenizer.encode(stop_token_str, add_special_tokens=False) |
| if not stop_token_ids_list: |
| raise ValueError(f"Stop token string '{stop_token_str}' could not be tokenized.") |
| self.stop_token_id = stop_token_ids_list[0] |
| self.stop_token_str = stop_token_str |
| self.tokenizer = tokenizer |
|
|
| def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor, **kwargs) -> bool: |
| |
| last_token_id = input_ids[0, -1].item() |
| |
| |
| if last_token_id == self.stop_token_id: |
| |
| |
| |
| if input_ids.shape[1] >= 1: |
| |
| |
| |
| decoded_text = self.tokenizer.decode(input_ids[0, -len(self.tokenizer.encode(self.stop_token_str, add_special_tokens=False)):], skip_special_tokens=True) |
| if self.stop_token_str in decoded_text: |
| return True |
| return False |
|
|
| class EndpointHandler: |
| def __init__(self, path: str = ""): |
| """ |
| Initializes the model and tokenizer. |
| Args: |
| path (str): Path to the directory containing model files. |
| On Hugging Face Inference Endpoints, this is automatically set. |
| """ |
| self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| print(f"Loading model on device: {self.device}") |
|
|
| |
| |
| model_path = path if path else "." |
|
|
| try: |
| self.tokenizer = AutoTokenizer.from_pretrained(model_path) |
| self.model = AutoModelForCausalLM.from_pretrained(model_path) |
| self.model.to(self.device) |
| self.model.eval() |
|
|
| |
| if self.tokenizer.pad_token is None: |
| print("Tokenizer pad_token not set. Setting to eos_token.") |
| self.tokenizer.pad_token = self.tokenizer.eos_token |
| self.model.config.pad_token_id = self.model.config.eos_token_id |
| |
| print("Model and tokenizer loaded successfully.") |
| |
| |
| self.stopping_criteria = StoppingCriteriaList([StopOnPatientToken(self.tokenizer)]) |
|
|
| except Exception as e: |
| print(f"Error loading model or tokenizer from path '{model_path}': {e}") |
| raise e |
|
|
| def __call__(self, data: Dict[str, Any]) -> List[Dict[str, Any]]: |
| """ |
| Generates text based on the input prompt and parameters. |
| Args: |
| data (Dict[str, Any]): A dictionary containing: |
| - "inputs" (str): The prompt for the model. |
| - "parameters" (Dict, optional): Generation parameters. |
| Returns: |
| List[Dict[str, Any]]: A list containing a dictionary with "generated_text". |
| """ |
| try: |
| prompt = data.pop("inputs", None) |
| if prompt is None: |
| return [{"error": "No 'inputs' key found in the request data."}] |
|
|
| parameters = data.pop("parameters", {}) |
|
|
| |
| |
| default_params = { |
| "max_new_tokens": 80, |
| "pad_token_id": self.tokenizer.eos_token_id, |
| "eos_token_id": self.tokenizer.eos_token_id, |
| "no_repeat_ngram_size": 3, |
| "do_sample": True, |
| "top_k": 50, |
| "top_p": 0.92, |
| "temperature": 0.75 |
| } |
| |
| gen_params = {**default_params, **parameters} |
|
|
| |
| |
| |
| inputs = self.tokenizer.encode(prompt, return_tensors="pt", truncation=True, max_length=self.model.config.max_position_embeddings - gen_params["max_new_tokens"]) |
| inputs = inputs.to(self.device) |
|
|
| |
| with torch.no_grad(): |
| outputs = self.model.generate( |
| inputs, |
| stopping_criteria=self.stopping_criteria, |
| **gen_params |
| ) |
| |
| |
| |
| generated_sequence = outputs[0] |
| prompt_length = inputs.shape[1] |
| generated_text_tokens = generated_sequence[prompt_length:] |
| |
| response_text = self.tokenizer.decode(generated_text_tokens, skip_special_tokens=True).strip() |
|
|
| |
| if self.stopping_criteria[0].stop_token_str in response_text: |
| response_text = response_text.split(self.stopping_criteria[0].stop_token_str)[0].strip() |
|
|
| return [{"generated_text": response_text}] |
|
|
| except Exception as e: |
| print(f"Error during inference: {e}") |
| |
| return [{"error": str(e), "message": "Inference failed."}] |
|
|
|
|