Team Ai
Modelpublic

startificial/nli-implementation

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes22downloads
handler.py121 linesDownload Raw Back to root
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}"}