import os import torch from transformers import AutoModelForCausalLM, AutoTokenizer, StoppingCriteria, StoppingCriteriaList from typing import Dict, List, Any # Define a stopping criteria for the model to stop on the "Patient:" token # This helps prevent the model from generating the patient's turn. class StopOnPatientToken(StoppingCriteria): def __init__(self, tokenizer, stop_token_str="Patient:"): super().__init__() # Encode the stop token string, ensuring not to add special tokens around it for this specific check # We are interested in the raw token IDs for "Patient:". # We'll take the first token ID of "Patient:" as the primary stop signal. # This might need adjustment if "Patient:" tokenizes into multiple relevant tokens. 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] # Using the first token of "Patient:" self.stop_token_str = stop_token_str self.tokenizer = tokenizer def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor, **kwargs) -> bool: # Get the last generated token last_token_id = input_ids[0, -1].item() # Check if the last token is the stop token ID if last_token_id == self.stop_token_id: # For more robustness, one could check the last N tokens to see if they form "Patient:" # For simplicity here, we stop if the first token of "Patient:" is generated. # Let's decode the last few tokens to be more sure if input_ids.shape[1] >= 1: # Ensure there's at least one token # Decode the last few tokens (e.g., up to the length of "Patient:") # This part is tricky because "Patient:" might be multiple tokens. # A simpler check is just the first token, but a more robust check would be: 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}") # The 'path' variable will be the root directory of your model in the endpoint's environment. # If empty, it implies the model files are in the current working directory (less common for endpoints). 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() # Set the model to evaluation mode # Ensure pad token is set for open-ended generation 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.") # Initialize stopping criteria 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 generation parameters - can be overridden by user # These are similar to what we used in testing default_params = { "max_new_tokens": 80, "pad_token_id": self.tokenizer.eos_token_id, "eos_token_id": self.tokenizer.eos_token_id, # Explicitly set EOS for generation "no_repeat_ngram_size": 3, "do_sample": True, "top_k": 50, "top_p": 0.92, "temperature": 0.75 } # Override defaults with user-provided parameters gen_params = {**default_params, **parameters} # Tokenize the input prompt # The prompt should ideally contain the conversation history and end with "Therapist:" # e.g., "Patient: I feel sad.\nTherapist:" 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) # Generate response with torch.no_grad(): # Ensure no gradients are computed during inference outputs = self.model.generate( inputs, stopping_criteria=self.stopping_criteria, # Add stopping criteria here **gen_params ) # Decode the generated tokens, excluding the input prompt part # The output includes the input prompt, so we need to slice it off. 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() # Further clean up if "Patient:" was partially generated and then stopped 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}") # It's good practice to return a JSON serializable error return [{"error": str(e), "message": "Inference failed."}]