capofwesh20 commited on
Commit
018defe
·
1 Parent(s): 539831d

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +20 -0
app.py CHANGED
@@ -1,3 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  def classify(im):
2
  features = feature_extractor(im, return_tensors='pt')
3
  logits = model(features["pixel_values"])[-1]
 
1
+ import datasets
2
+ from transformers import AutoFeatureExtractor, AutoModelForImageClassification
3
+
4
+ dataset = datasets.load_dataset("beans")
5
+
6
+ extractor = AutoFeatureExtractor.from_pretrained("saved_model_files")
7
+ model = AutoModelForImageClassification.from_pretrained("saved_model_files")
8
+
9
+ labels = dataset['train'].features['labels'].names
10
+
11
+ def classify(im):
12
+ features = feature_extractor(im, return_tensors='pt')
13
+ logits = model(features["pixel_values"])[-1]
14
+ probability = torch.nn.functional.softmax(logits, dim=-1)
15
+ probs = probability[0].detach().numpy()
16
+ confidences = {label: float(probs[i]) for i, label in enumerate(labels)}
17
+ return confidences
18
+
19
+ model_path = 'drive/MyDrive/huggingface/saved_model_files'
20
+
21
  def classify(im):
22
  features = feature_extractor(im, return_tensors='pt')
23
  logits = model(features["pixel_values"])[-1]