MindMate_v1 / handler.py
pubgmob1024's picture
Create handler.py
e61197f verified
Raw
History Blame Contribute Delete
7.11 kB
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."}]