Amitai/Image-Classification
0
1import gradio as gr2from joblib import load3import torch4import clip5from PIL import Image6import numpy as np7import pickle8import pickletools9 10CLF_FILENAME = "lr-model.pkl"11 12clf = load(CLF_FILENAME)13device = "cuda" if torch.cuda.is_available() else "cpu"14model, preprocess = clip.load("ViT-B/32", device)15 16 17def classify_image(img):18 # img comes from gr.Image(type="numpy")19 im = Image.fromarray(img, mode="RGB")20 image_pre_process = [preprocess(im)]21 image_input = torch.tensor(np.stack(image_pre_process)).to(device)22 23 with torch.no_grad():24 image_features = model.encode_image(image_input)25 26 image_data = image_features.cpu().numpy()27 pred = clf.predict(image_data) # e.g., [0] or [1]28 29 outputs = {0: '๐ฑ Biodegradable', 1: '๐ Non-biodegradable'}30 return outputs[int(pred[0] >= 0.5)] # be explicit about indexing31 32 33image = gr.Image(34 type="numpy",35 label="Upload an image",36)37 38iface = gr.Interface(39 fn=classify_image,40 inputs=image,41 outputs=gr.Textbox(label="Prediction"),42 examples=["Pizza.JPG", "poly.JPG"],43)44 45if __name__ == "__main__":46 iface.launch()47 