Views
No views yet
google/vit-base-patch16-224)| Class | Description |
|---|---|
freshapples | Fresh, ready-to-eat apples |
freshbanana | Fresh, ripe bananas |
freshoranges | Fresh, ripe oranges |
rottenapples | Overripe/rotten apples |
rottenbanana | Overripe/rotten bananas |
rottenoranges | Overripe/rotten oranges |
unripe apple | Unripe apples |
unripe banana | Unripe bananas |
unripe orange | Unripe oranges |
pip install torch torchvision transformers scikit-learn pillow joblib numpy huggingface_hub1
2import json
3import joblib
4from pathlib import Path
5from PIL import Image
6import torch
7import numpy as np
8from huggingface_hub import hf_hub_download, HfApi
9from transformers import AutoImageProcessor, ViTModel
10import warnings
11
12# ----------------- CONFIG -----------------
13REPO_ID = "Meeteshn/vit_fruit_ripeness_classifier"
14DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
15NESTED_FOLDER = "vit_fruit_ripeness_updated" # your repo uses this nested folder
16TOP_K = 5
17# ------------------------------------------
18
19def hf_download_try(repo_id: str, filename: str, nested_folder: str = NESTED_FOLDER):
20 """
21 Try to download `filename` from repo root, then from nested_folder/filename.
22 Returns local path to downloaded file or raises an informative error.
23 """
24 candidates = [filename, f"{nested_folder}/{filename}"]
25 last_exc = None
26 for f in candidates:
27 try:
28 print(f"Trying to download '{f}' from '{repo_id}'...")
29 path = hf_hub_download(repo_id=repo_id, filename=f)
30 print("Downloaded:", path)
31 return path
32 except Exception as e:
33 print(f"Not found at '{f}': {e}")
34 last_exc = e
35 raise RuntimeError(f"Could not download '{filename}' from repo '{repo_id}'. Last error: {last_exc}")
36
37def load_processor_and_backbone(repo_id: str, nested_folder: str = NESTED_FOLDER, device: str = DEVICE):
38 """
39 Try several likely subfolder locations for processor/backbone.
40 Returns (processor, backbone).
41 """
42 # candidate subfolders for processor
43 proc_candidates = [
44 "processor",
45 f"{nested_folder}/processor",
46 "", # no subfolder (root)
47 ]
48 last_exc = None
49 for sub in proc_candidates:
50 try:
51 if sub == "":
52 print(f"Trying AutoImageProcessor.from_pretrained('{repo_id}')")
53 processor = AutoImageProcessor.from_pretrained(repo_id, use_fast=True)
54 else:
55 print(f"Trying AutoImageProcessor.from_pretrained('{repo_id}', subfolder='{sub}')")
56 processor = AutoImageProcessor.from_pretrained(repo_id, subfolder=sub, use_fast=True)
57 # now try backbone with matching guessed subfolder
58 backbone_sub = sub.replace("processor", "vit_backbone") if sub and "processor" in sub else ("vit_backbone" if sub == "" else f"{nested_folder}/vit_backbone")
59 try:
60 print(f"Trying ViTModel.from_pretrained('{repo_id}', subfolder='{backbone_sub}')")
61 backbone = ViTModel.from_pretrained(repo_id, subfolder=backbone_sub)
62 except Exception as e_backbone:
63 # final fallback: try root vit_backbone
64 print(f"Backbone attempt failed for sub='{backbone_sub}': {e_backbone}. Trying root 'vit_backbone'.")
65 backbone = ViTModel.from_pretrained(repo_id, subfolder="vit_backbone")
66 backbone.to(device)
67 backbone.eval()
68 print(f"Loaded processor/backbone from subfolder='{sub or 'root'}'")
69 return processor, backbone
70 except Exception as e:
71 print(f"Processor load failed for sub='{sub}': {e}")
72 last_exc = e
73 # ultimate fallback: official ViT from hub
74 warnings.warn("Could not load processor/backbone from repo; falling back to official 'google/vit-base-patch16-224'.")
75 processor = AutoImageProcessor.from_pretrained("google/vit-base-patch16-224", use_fast=True)
76 backbone = ViTModel.from_pretrained("google/vit-base-patch16-224")
77 backbone.to(device)
78 backbone.eval()
79 return processor, backbone
80
81# ----------------- Load assets (robust) -----------------
82processor, backbone = load_processor_and_backbone(REPO_ID, nested_folder=NESTED_FOLDER, device=DEVICE)
83
84# Download sklearn artifacts (try root then nested)
85scaler_path = hf_download_try(REPO_ID, "scaler.joblib", nested_folder=NESTED_FOLDER)
86clf_path = hf_download_try(REPO_ID, "logistic_model.joblib", nested_folder=NESTED_FOLDER)
87metadata_path = hf_download_try(REPO_ID, "metadata.json", nested_folder=NESTED_FOLDER)
88
89scaler = joblib.load(scaler_path)
90clf = joblib.load(clf_path)
91metadata = json.loads(Path(metadata_path).read_text(encoding="utf-8"))
92classes = metadata["classes"]
93
94# ----------------- Prediction function -----------------
95def predict(image_path: str):
96 """Predict ripeness condition for a single image."""
97 img = Image.open(image_path).convert("RGB")
98 inputs = processor(images=img, return_tensors="pt")
99 pixel_values = inputs["pixel_values"].to(DEVICE)
100
101 with torch.no_grad():
102 out = backbone(pixel_values=pixel_values, return_dict=True)
103 pooled = getattr(out, "pooler_output", None)
104 if pooled is None:
105 pooled = out.last_hidden_state[:, 0, :]
106 feat = pooled.cpu().numpy()
107
108 feat_scaled = scaler.transform(feat)
109 # get probabilities (works for sklearn logistic / classifiers with predict_proba)
110 if hasattr(clf, "predict_proba"):
111 probs = clf.predict_proba(feat_scaled)[0]
112 else:
113 # fallback for classifiers without predict_proba
114 dec = clf.decision_function(feat_scaled)[0]
115 exp = np.exp(dec - np.max(dec))
116 probs = exp / exp.sum()
117
118 idx = int(np.argmax(probs))
119 return classes[idx], float(probs[idx]), {classes[i]: float(probs[i]) for i in range(len(classes))}
120
121# ----------------- Example usage -----------------
122if __name__ == "__main__":
123 sample_image = "my_apple.jpg" # change as needed
124 label, prob, all_probs = predict(sample_image)
125 print(f"Prediction: {label} ({prob*100:.2f}%)")
126 print("\nTop probabilities:")
127 for cls, p in sorted(all_probs.items(), key=lambda x: -x[1])[:TOP_K]:
128 print(f" {cls}: {p*100:.2f}%")
1291from pathlib import Path
2import csv
3
4def batch_predict(folder_path: str, output_csv: str = "predictions.csv"):
5 """Predict ripeness for all images in a folder."""
6 folder = Path(folder_path)
7
8 with open(output_csv, "w", newline="", encoding="utf-8") as f:
9 writer = csv.writer(f)
10 writer.writerow(["filename", "predicted_label", "probability"])
11
12 for img_path in sorted(folder.rglob("*")):
13 if img_path.suffix.lower() not in [".jpg", ".jpeg", ".png", ".bmp"]:
14 continue
15
16 label, prob, _ = predict(str(img_path))
17 writer.writerow([img_path.name, label, f"{prob*100:.2f}%"])
18
19 print(f"Predictions saved to {output_csv}")
20
21# Usage
22batch_predict("path/to/images")Prediction: rottenapples (71.24%)
Top 5 probabilities:
rottenapples: 71.24%
rottenbanana: 12.35%
freshapples: 6.12%
unripe apple: 4.89%
freshoranges: 2.31%vit_fruit_ripeness_updated/
├── processor/ # AutoImageProcessor configuration
├── vit_backbone/ # ViT feature extractor weights
├── logistic_model.joblib # Trained classifier
├── scaler.joblib # Feature scaler
├── metadata.json # Class labels and metadata
└── features_extracted.npz # (Optional) Cached featuresgoogle/vit-base-patch16-2241@misc{vit-fruit-ripeness-classifier,
2 author = {Nagrecha, Meetesh},
3 title = {ViT Fruit Ripeness Classifier},
4 year = {2024},
5 publisher = {Hugging Face},
6 howpublished = {\url{https://huggingface.co/Meeteshn/vit_fruit_ripeness_classifier}}
7}