cifar / app.py
klibaner's picture
Upload 4 files
3bc6b95 verified
Raw
History Blame Contribute Delete
3.66 kB
from pathlib import Path
import gradio as gr
import torch
import torch.nn as nn
from torchvision import transforms
class SimpleCNN(nn.Module):
# This must match the network trained in the original notebook.
def __init__(self):
super().__init__()
self.features = nn.Sequential(
nn.Conv2d(3, 32, kernel_size=3, padding=1),
nn.ReLU(),
nn.MaxPool2d(2),
nn.Conv2d(32, 64, kernel_size=3, padding=1),
nn.ReLU(),
nn.MaxPool2d(2),
nn.Conv2d(64, 128, kernel_size=3, padding=1),
nn.ReLU(),
nn.MaxPool2d(2),
)
self.classifier = nn.Sequential(
nn.Flatten(),
nn.Linear(128 * 4 * 4, 256),
nn.ReLU(),
nn.Dropout(0.30),
nn.Linear(256, 10),
)
def forward(self, x):
return self.classifier(self.features(x))
MODEL_PATH = Path(__file__).with_name("cifar10_cnn.pt")
DEVICE = torch.device("cpu")
TRANSFORM = transforms.Compose([
transforms.Resize((32, 32)),
transforms.ToTensor(),
transforms.Normalize(
(0.4914, 0.4822, 0.4465),
(0.2470, 0.2435, 0.2616),
),
])
def load_model():
if not MODEL_PATH.exists():
raise FileNotFoundError(
"Файл cifar10_cnn.pt не найден. Загрузите его в корень Space рядом с app.py."
)
try:
checkpoint = torch.load(MODEL_PATH, map_location=DEVICE, weights_only=True)
except TypeError:
checkpoint = torch.load(MODEL_PATH, map_location=DEVICE)
model = SimpleCNN().to(DEVICE)
model.load_state_dict(checkpoint["model_state_dict"])
model.eval()
classes = checkpoint.get(
"classes",
["airplane", "automobile", "bird", "cat", "deer",
"dog", "frog", "horse", "ship", "truck"],
)
return model, classes
MODEL, CLASSES = load_model()
def predict(image):
if image is None:
return {}, "Загрузите изображение."
x = TRANSFORM(image.convert("RGB")).unsqueeze(0).to(DEVICE)
with torch.no_grad():
probabilities = torch.softmax(MODEL(x), dim=1)[0].cpu()
values, indices = torch.topk(probabilities, k=3)
results = {
CLASSES[index.item()]: float(value)
for value, index in zip(values, indices)
}
label = CLASSES[indices[0].item()]
confidence = float(values[0]) * 100
return results, f"Модель считает, что это: **{label}** ({confidence:.1f}%)."
with gr.Blocks() as demo:
gr.Markdown(
"""# 🖼️ Распознавание изображений — CIFAR-10
Загрузите фото. Модель выберет наиболее похожий класс: **самолёт, автомобиль, птица, кот, олень, собака, лягушка, лошадь, корабль или грузовик**.
> Это учебная модель, обученная на маленьких изображениях CIFAR-10. На обычных фото она может ошибаться."""
)
with gr.Row():
image_input = gr.Image(
label="Загрузите изображение",
type="pil",
sources=["upload", "webcam"],
)
result = gr.Label(label="Три наиболее вероятных класса", num_top_classes=3)
explanation = gr.Markdown()
button = gr.Button("Распознать", variant="primary")
button.click(predict, inputs=image_input, outputs=[result, explanation])
if __name__ == "__main__":
demo.launch()