Views
No views yet
1from peft import PeftModel
2from PIL import Image
3from transformers import AutoImageProcessor, AutoModelForImageClassification
4
5from torchvision.transforms import (
6 CenterCrop,
7 Compose,
8 Normalize,
9 RandomHorizontalFlip,
10 RandomResizedCrop,
11 Resize,
12 ToTensor,
13)
14
15model_name = 'google/vit-large-patch16-224'
16adapter = 'monsoon-nlp/eyegazer-vit-binary'
17
18image_processor = AutoImageProcessor.from_pretrained(model_name)
19
20normalize = Normalize(mean=image_processor.image_mean, std=image_processor.image_std)
21train_transforms = Compose(
22 [
23 RandomResizedCrop(image_processor.size["height"]),
24 RandomHorizontalFlip(),
25 ToTensor(),
26 normalize,
27 ]
28)
29
30val_transforms = Compose(
31 [
32 Resize(image_processor.size["height"]),
33 CenterCrop(image_processor.size["height"]),
34 ToTensor(),
35 normalize,
36 ]
37)
38
39model = AutoModelForImageClassification.from_pretrained(
40 model_name,
41 ignore_mismatched_sizes=True,
42 num_labels=2,
43)
44
45lora_model = PeftModel.from_pretrained(model, adapter)
46
47img = Image.open("sample.png")
48pimg = val_transforms(img.convert("RGB"))
49batch = pimg.unsqueeze(0)
50op = lora_model(batch)
51vals = op.logits.tolist()[0]
52
53if vals[0] > vals[1]:
54 return "Predicted unaffected"
55else:
56 return "Predicted affected to some degree"