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()