File size: 3,267 Bytes
79aeec8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
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")

@app.route('/')
def home():
    return render_template('resnet.html')

@app.route('/api/predict', methods=['POST'])
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)