CoinPrediction / app.py
XaviGM11's picture
Upload app.py
a04f0aa verified
Raw
History Blame Contribute Delete
2.31 kB
from pathlib import Path
from fastai.vision.all import *
import gradio as gr
import torch
from torchvision import transforms
def label_func(o):
return parent_label(o)
def get_trainval_files(path):
return get_image_files(path)
# Cargamos el modelo entrenado (CPU en Spaces gratuitos).
learn = load_learner('model.pkl', cpu=True)
model = learn.model.eval().float() # red neuronal (logits crudos)
labels = list(learn.dls.vocab) # nombres legibles de las clases
preprocess = transforms.Compose([
transforms.Resize(256),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]),
])
def predict(img):
"""Recibe una imagen, devuelve {clase: probabilidad} para que gr.Label la pinte."""
if img is None:
return None
img = img.convert("RGB")
x = preprocess(img).unsqueeze(0) # (1, 3, 224, 224)
with torch.no_grad():
logits = model(x)[0]
probs = torch.softmax(logits, dim=0)
return {labels[i]: float(probs[i]) for i in range(len(labels))}
title = "Clasificador de monedas 🪙"
description = (
"Sube una foto de una moneda y el modelo predecirá de qué moneda se trata, "
"mostrando las clases más probables con su porcentaje."
)
# Imágenes de ejemplo: detectamos automáticamente las que hayas subido al repo,
# ya sea en una carpeta 'examples/' o en la raíz del Space.
IMG_EXTS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
example_imgs = []
for folder in [Path("examples"), Path(".")]:
if folder.exists():
example_imgs += sorted(
str(p) for p in folder.iterdir()
if p.is_file() and p.suffix.lower() in IMG_EXTS
)
example_imgs = list(dict.fromkeys(example_imgs))[:12] # sin duplicados, máx 12
demo = gr.Interface(
fn=predict,
inputs=gr.Image(type="pil", label="Sube una imagen de la moneda"),
outputs=gr.Label(num_top_classes=5, label="Predicción (top 5)"),
title=title,
description=description,
examples=example_imgs or None, # clicar un ejemplo lo carga en el input
cache_examples=False, # no precalcular (evita lentitud/errores en el build)
flagging_mode="never",
)
if __name__ == "__main__":
demo.launch()