skycope/ai_art_classification
0
1import gradio as gr2from fastai.vision.all import *3import timm4 5 6# load learner7learner = load_learner('export.pkl')8 9im = PILImage.create('dalle.png')10im.thumbnail((224, 224))11im12learner.predict(im)13# get categories14categories = learner.dls.vocab15 16# define function to classify image17def classify_image(inp):18 pred,pred_idx,probs = learner.predict(inp)19 return dict(zip(categories, map(float, probs)))20 21image = gr.inputs.Image(shape=(224,224))22label = gr.outputs.Label()23examples = ['dalle.png', 'midjourney.jpeg', 'stable.jpg']24# make gradio interface25intf = gr.interface.Interface(classify_image, "image", "label", 26 examples=examples)27 28intf.launch(inline=False)