File size: 2,311 Bytes
a04f0aa
df118da
 
260a3da
 
df118da
 
 
 
 
 
 
 
260a3da
 
 
 
 
 
 
 
 
 
 
 
df118da
 
 
 
 
 
260a3da
 
 
 
 
df118da
 
 
 
 
 
 
 
 
a04f0aa
 
 
 
 
 
 
 
 
 
 
 
df118da
 
 
 
 
 
a04f0aa
 
df118da
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
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()