File size: 1,734 Bytes
6aabd96
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
590f763
 
6aabd96
 
 
 
590f763
 
 
 
 
 
6aabd96
 
 
 
 
 
590f763
6aabd96
 
590f763
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6aabd96
590f763
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
"""
inference_engine.py

Controlled Implicit Inference using MNLI.
Compatible with transformers >=5.x

roberta-large-mnli Label-Meaning mapping is usually:
    LABEL_0	CONTRADICTION
    LABEL_1	NEUTRAL
    LABEL_2	ENTAILMENT
"""

from transformers import pipeline


class InferenceEngine:
    def __init__(self, threshold=0.80):
        """
        threshold: minimum confidence score required
        """
        self.threshold = threshold

        self.classifier = pipeline(
            "text-classification",
            model="roberta-large-mnli"
        )

    def validate_hypotheses(self, premise, hypotheses):
        """
        Validate candidate hypotheses using MNLI.
        Returns:
            List of inference classifications.
        """

        validated = []

        label_mapping = {
            "LABEL_0": "CONTRADICTION",
            "LABEL_1": "NEUTRAL",
            "LABEL_2": "ENTAILMENT"
        }

        for hypothesis in hypotheses:

            result = self.classifier(
                f"{premise} </s></s> {hypothesis}"
            )[0]

            raw_label = result["label"]
            score = result["score"]

            label = label_mapping.get(raw_label, raw_label)

            validated.append({
                "hypothesis": hypothesis,
                "label": label,
                "confidence": round(score, 3),
                "confidence_tier": self._confidence_tier(score)
            })

        return validated
    
    def _confidence_tier(self, score):
        """
        Convert numeric confidence into analyst-friendly tier.
        """

        if score >= 0.95:
            return "High"

        elif score >= 0.80:
            return "Moderate"

        return "Low"