File size: 3,077 Bytes
357db13
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
75
76
77
"""Arlo Classify — Choice en Noul.

    from transformers import AutoModel, AutoTokenizer

    model = AutoModel.from_pretrained("NoaberAI/arlo-classify", trust_remote_code=True)
    tokenizer = AutoTokenizer.from_pretrained("NoaberAI/arlo-classify")

    model.choice(tokenizer, "De put op de Kerkstraat is verstopt.",
                 ["verkeer", "riolering", "afval"])
    # [{"label": "riolering", "confidence": 0.91}, ...]

    model.noul(tokenizer, "Er ligt een boom over het fietspad.",
               "Dit is een spoedeisende melding.")
    # {"truth": 0.92}

Het model is een NLI-classifier (entailment / neutral / contradiction). Beide
primitieven leiden hun antwoord af uit diezelfde drie logits, in één forward
pass per vraag — er wordt geen tekst gegenereerd.
"""

import torch
from transformers import DebertaV2ForSequenceClassification

DEFAULT_HYPOTHESIS_TEMPLATE = "Dit gaat over {}."


class ArloClassifyModel(DebertaV2ForSequenceClassification):
    @torch.no_grad()
    def choice(
        self,
        tokenizer,
        text: str,
        options: list[str],
        hypothesis_template: str = DEFAULT_HYPOTHESIS_TEMPLATE,
        multi_label: bool = False,
    ) -> list[dict]:
        """Choice: kies uit een lijst opties. Elke optie wordt als stelling
        tegen de tekst gelegd; de entailment-score bepaalt de uitkomst.

        multi_label=False geeft een kansverdeling over de opties (samen 1),
        multi_label=True scoort elke optie onafhankelijk — gebruik dat als
        meerdere labels tegelijk waar kunnen zijn.
        """
        entailment_id = self.config.label2id.get("entailment", 0)

        hypotheses = [hypothesis_template.format(option) for option in options]
        inputs = tokenizer(
            [text] * len(options), hypotheses,
            return_tensors="pt", truncation=True, padding=True,
        ).to(self.device)

        entail_logits = self(**inputs).logits[:, entailment_id]
        scores = torch.sigmoid(entail_logits) if multi_label else torch.softmax(entail_logits, dim=0)

        results = [
            {"label": option, "confidence": float(score)}
            for option, score in zip(options, scores)
        ]
        return sorted(results, key=lambda r: r["confidence"], reverse=True)

    @torch.no_grad()
    def noul(self, tokenizer, text: str, statement: str) -> dict:
        """Noul: hoe waar is `statement` gegeven `text`? Geeft een waarde
        tussen 0 en 1 — een waarheidsgraad, geen hard ja/nee.

        Gebruikt entailment versus contradiction (de neutral-klasse wordt
        genegeerd), zodat de uitkomst een schone binaire verhouding is.
        """
        entailment_id = self.config.label2id.get("entailment", 0)
        contradiction_id = self.config.label2id.get("contradiction", 2)

        inputs = tokenizer(text, statement, return_tensors="pt", truncation=True).to(self.device)
        logits = self(**inputs).logits[0]

        pair = torch.stack([logits[entailment_id], logits[contradiction_id]])
        return {"truth": float(torch.softmax(pair, dim=0)[0])}