arlo-classify / modeling_arlo_classify.py
joranvnbeekAI's picture
Voeg custom modeling-bestand toe
357db13 verified
Raw History Blame Contribute Delete
3.08 kB
"""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])}