Team Ai
Modelpublic

tmutton/wcag-accessibility-classifier

sourceHugging Faceupdated 11d agoView on Hugging Face
0likes14downloads
predict.py49 linesDownload Raw Back to root
1import torch
2
3from transformers import (
4    AutoModelForSequenceClassification,
5    AutoTokenizer,
6)
7
8MODEL_PATH = "."
9
10tokenizer = AutoTokenizer.from_pretrained(MODEL_PATH)
11
12model = AutoModelForSequenceClassification.from_pretrained(
13    MODEL_PATH
14)
15
16
17def predict(text):
18    inputs = tokenizer(
19        text,
20        return_tensors="pt",
21        truncation=True,
22    )
23
24    with torch.no_grad():
25        outputs = model(**inputs)
26
27    probabilities = torch.softmax(
28        outputs.logits,
29        dim=-1
30    )[0]
31
32    predicted_id = torch.argmax(probabilities).item()
33
34    predicted_label = model.config.id2label[predicted_id]
35    confidence = probabilities[predicted_id].item()
36
37    return predicted_label, confidence
38
39
40while True:
41    text = input("\nDescribe an accessibility issue (or 'quit'): ")
42
43    if text.lower() == "quit":
44        break
45
46    label, confidence = predict(text)
47
48    print(f"\nPrediction: {label}")
49    print(f"Confidence: {confidence:.1%}")