tmutton/wcag-accessibility-classifier
014
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%}")