startificial/nli-implementation
022
1import os2import torch3from transformers import AutoTokenizer, AutoModelForSequenceClassification4from typing import Dict, List, Any # <-- ADD THIS LINE5 6class EndpointHandler():7 def __init__(self, model_id: str):8 """9 Initializes the handler by loading the model and tokenizer.10 11 Args:12 model_id (str): The Hugging Face model ID (e.g., "MoritzLaurer/DeBERTa-v3-base-mnli")13 This is automatically passed by the Inference Endpoint infrastructure.14 """15 print(f"Loading model '{model_id}'...")16 self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")17 print(f"Using device: {self.device}")18 19 self.tokenizer = AutoTokenizer.from_pretrained(model_id)20 self.model = AutoModelForSequenceClassification.from_pretrained(model_id)21 22 # Move model to the determined device23 self.model.to(self.device)24 # Set model to evaluation mode for consistent inference25 self.model.eval()26 print("Model and tokenizer loaded successfully.")27 28 # --- Determine Label Order ---29 # Preferred: Dynamically get labels from model config30 try:31 # Sort by ID to ensure consistent order if dict isn't ordered32 sorted_labels = sorted(self.model.config.id2label.items())33 self.label_names = [label for _, label in sorted_labels]34 print(f"Using label names from model config: {self.label_names}")35 # Basic validation for NLI task36 if len(self.label_names) != 3:37 print(f"Warning: Expected 3 labels for NLI, but model config has {len(self.label_names)}. Proceeding with model's labels.")38 if not any("entail" in l.lower() for l in self.label_names) or \39 not any("neutral" in l.lower() for l in self.label_names) or \40 not any("contra" in l.lower() for l in self.label_names):41 print(f"Warning: Model labels {self.label_names} might not match standard NLI labels ('entailment', 'neutral', 'contradiction').")42 43 except AttributeError:44 # Fallback: Use the explicitly requested labels if config is missing/malformed45 self.label_names = ["entailment", "neutral", "contradiction"]46 print(f"Warning: Could not read labels from model config. Falling back to default: {self.label_names}")47 print("Ensure this order matches the actual output order of the model!")48 49 print(f"Configured label order for output: {self.label_names}")50 51 52 # Corrected type hints in the signature below53 def __call__(self, data: Dict[str, Any]) -> Dict[str, Any] | List[Dict[str, Any]]:54 """55 Handles inference requests.56 57 Args:58 data (Dict[str, Any]): The input data payload from the request.59 Expected keys: "premise" (str) and "hypothesis" (str).60 Can optionally be nested under "inputs".61 62 Returns:63 Dict[str, Any] | List[Dict[str, Any]]: A dictionary containing error info,64 or a list of dictionaries, each mapping65 a label name to its probability score.66 """67 # --- Input Parsing ---68 inputs = data.get("inputs", data) # Allow for optional "inputs" nesting69 premise = inputs.get("premise")70 hypothesis = inputs.get("hypothesis")71 72 # Basic input validation73 if not premise or not isinstance(premise, str):74 return {"error": "Missing or invalid 'premise' key in input. Expected a string."}75 if not hypothesis or not isinstance(hypothesis, str):76 return {"error": "Missing or invalid 'hypothesis' key in input. Expected a string."}77 78 # --- Tokenization ---79 # Tokenize the premise-hypothesis pair80 try:81 tokenized_inputs = self.tokenizer(82 premise,83 hypothesis,84 return_tensors="pt", # Return PyTorch tensors85 truncation=True, # Truncate if longer than max length86 padding=True, # Pad to the longest sequence in the batch (or max_length)87 max_length=self.tokenizer.model_max_length # Use model's max length88 )89 except Exception as e:90 print(f"Error during tokenization: {e}")91 return {"error": f"Failed to tokenize input: {e}"}92 93 94 # Move tokenized inputs to the same device as the model95 tokenized_inputs = {k: v.to(self.device) for k, v in tokenized_inputs.items()}96 97 # --- Inference ---98 try:99 with torch.no_grad(): # Disable gradient calculations for efficiency100 outputs = self.model(**tokenized_inputs)101 logits = outputs.logits102 103 # Apply Softmax to get probabilities104 probabilities = torch.softmax(logits, dim=-1)105 106 # Move probabilities to CPU and convert to list107 # Squeeze or index [0] if processing single pairs (typical for endpoints)108 scores = probabilities.cpu().numpy()[0].tolist()109 110 # --- Format Output ---111 # Pair labels with their corresponding scores112 result = [{"label": label, "score": score} for label, score in zip(self.label_names, scores)]113 114 return result115 116 except Exception as e:117 print(f"Error during model inference: {e}")118 # Consider logging the full traceback here in a real deployment119 # import traceback120 # traceback.print_exc()121 return {"error": f"Model inference failed: {e}"}