Spaces:
Sleeping
Sleeping
| import os | |
| import io | |
| import torch | |
| import torch.nn.functional as F | |
| from flask import Flask, render_template, request, jsonify | |
| from torchvision import transforms | |
| from PIL import Image | |
| from resnet_mlp import ResNetMLP | |
| from resnet_kan import ResNetKAN | |
| app = Flask(__name__) | |
| CLASSES = ['airplane', 'automobile', 'bird', 'cat', 'deer', 'dog', 'frog', 'horse', 'ship', 'truck'] | |
| IMAGENET_MEAN = [0.485, 0.456, 0.406] | |
| IMAGENET_STD = [0.229, 0.224, 0.225] | |
| inference_transform = transforms.Compose([ | |
| transforms.Resize(256), | |
| transforms.ToTensor(), | |
| transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD), | |
| ]) | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| def load_model_weights(model, path): | |
| if os.path.exists(path): | |
| model.load_state_dict(torch.load(path, map_location=device)) | |
| model.eval() | |
| return model.to(device) | |
| FREEZE_BACKBONE = True | |
| model_mlp = ResNetMLP( | |
| num_classes=10, | |
| freeze_backbone=FREEZE_BACKBONE, | |
| hidden_dim=512 | |
| ) | |
| model_mlp = load_model_weights(model_mlp, "weights/tesresnet_mlp_cifar10_run1.pth") | |
| model_kan = ResNetKAN( | |
| num_classes=10, | |
| freeze_backbone=FREEZE_BACKBONE | |
| ) | |
| model_kan = load_model_weights(model_kan, "weights/resnet_kan_cifar10_run1.pth") | |
| def home(): | |
| return render_template('resnet.html') | |
| def predict(): | |
| if 'file' not in request.files: | |
| return jsonify({"status": "error", "message": "No file"}), 400 | |
| file = request.files['file'] | |
| img_bytes = file.read() | |
| try: | |
| image = Image.open(io.BytesIO(img_bytes)).convert('RGB') | |
| tensor_img = inference_transform(image).unsqueeze(0).to(device) | |
| with torch.no_grad(): | |
| out_1 = model_mlp(tensor_img) | |
| prob_1 = F.softmax(out_1, dim=1) | |
| conf_1, idx_1 = torch.max(prob_1, 1) | |
| out_2 = model_kan(tensor_img) | |
| prob_2 = F.softmax(out_2, dim=1) | |
| conf_2, idx_2 = torch.max(prob_2, 1) | |
| p1_list = prob_1[0].tolist() | |
| p2_list = prob_2[0].tolist() | |
| all_p1 = [{"class": CLASSES[i], "confidence": round(p1_list[i] * 100, 2)} for i in range(10)] | |
| all_p2 = [{"class": CLASSES[i], "confidence": round(p2_list[i] * 100, 2)} for i in range(10)] | |
| all_p1.sort(key=lambda x: x['confidence'], reverse=True) | |
| all_p2.sort(key=lambda x: x['confidence'], reverse=True) | |
| is_outlier = False | |
| if conf_1.item() < 0.35 and conf_2.item() < 0.35: | |
| is_outlier = True | |
| return jsonify({ | |
| "status": "success", | |
| "is_outlier": is_outlier, | |
| "model_1": { | |
| "class": CLASSES[idx_1.item()], | |
| "confidence": round(conf_1.item() * 100, 2), | |
| "all_probs": all_p1 | |
| }, | |
| "model_2": { | |
| "class": CLASSES[idx_2.item()], | |
| "confidence": round(conf_2.item() * 100, 2), | |
| "all_probs": all_p2 | |
| } | |
| }) | |
| except Exception as e: | |
| return jsonify({"status": "error", "message": str(e)}) | |
| if __name__ == '__main__': | |
| app.run(host='0.0.0.0', port=7860) |