import gradio as gr import torch from torchvision import transforms as T from PIL import Image from io import BytesIO import matplotlib.pyplot as plt from src.train import loss_fn from src.model_config import load_model, set_device from pydantic import GetCoreSchemaHandler from starlette.requests import Request # Inject dummy schema for starlette.Request so Pydantic v2 can handle it def request_pydantic_schema(_: type, handler: GetCoreSchemaHandler): from pydantic_core import core_schema return core_schema.any_schema() Request.__get_pydantic_core_schema__ = request_pydantic_schema device = set_device() model = load_model(device=device) model.eval() transform = T.Compose([ T.Grayscale(num_output_channels=1), T.Resize((28, 28)), T.ToTensor() ]) def predict(image): original = image.convert("L") tensor = transform(original).unsqueeze(0).to(device) original_resized = original.resize((200, 200), Image.NEAREST) with torch.no_grad(): decoded = model(tensor) score = loss_fn(decoded, tensor).item() diff = (decoded - tensor).squeeze().cpu().numpy() ** 2 reconstructed = decoded.squeeze().cpu().numpy() # Anomaly interpretation if score > 0.02: verdict = "🚨 Anomaly detected!" else: verdict = "✅ Everything looks normal." # Create heatmap fig, ax = plt.subplots() ax.imshow(diff, cmap="hot") ax.axis("off") buf1 = BytesIO() plt.savefig(buf1, format="png", bbox_inches="tight", pad_inches=0) plt.close(fig) buf1.seek(0) heatmap = Image.open(buf1) # Reconstructed image fig2, ax2 = plt.subplots() ax2.imshow(reconstructed, cmap="gray") ax2.axis("off") buf2 = BytesIO() plt.savefig(buf2, format="png", bbox_inches="tight", pad_inches=0) plt.close(fig2) buf2.seek(0) reconstructed_img = Image.open(buf2) return f"{score:.5f}", verdict, original_resized, reconstructed_img, heatmap iface = gr.Interface( fn=predict, inputs=gr.Image(type="pil", label="Upload MNIST-style Image"), outputs=[ gr.Textbox(label="Anomaly Score"), gr.Textbox(label="Verdict"), gr.Image(label="Original", height=200, width=200), gr.Image(label="Reconstruction", height=200, width=200), gr.Image(label="Heatmap", height=200, width=200), ], allow_flagging="never", title="🧠 Anomaly Detection", description="Upload a grayscale MNIST image (28x28) to get anomaly score." ) iface.launch(server_name="0.0.0.0", server_port=7860)