File size: 3,266 Bytes
79aeec8 fa38fbd 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.94 or conf_2.item() < 0.94:
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) |