dmusingu/CXR-IMAGE-Classification
0
1from fastapi import FastAPI, File, UploadFile, HTTPException2from fastapi.responses import JSONResponse3import io4import joblib5import torch6import numpy as np7import torchvision.transforms as transforms8from PIL import Image9import yaml10import traceback11import timm12import logging13from fastapi.logger import logger14 15 16app = FastAPI()17 18device = torch.device("cuda" if torch.cuda.is_available() else "cpu")19 20 21class_mapping = {'tb': 0, 'healthy': 1, 'sick_but_no_tb': 2}22reverse_mapping = {v: k for k, v in class_mapping.items()}23labels = list(class_mapping.keys())24 25def load_model():26 # config = read_params(config_path)27 model = timm.create_model('convnext_base.clip_laiona', pretrained=True, num_classes=3)28 model_state_dict = torch.load('model.pth', map_location=device)29 model.load_state_dict(model_state_dict)30 model.eval()31 return model32 33 34def transform_image(image_bytes):35 my_transforms = transforms.Compose([transforms.Resize(255),36 transforms.CenterCrop(224),37 transforms.ToTensor(),38 transforms.Normalize(39 [0.485, 0.456, 0.406],40 [0.229, 0.224, 0.225])])41 image = Image.open(io.BytesIO(image_bytes)).convert('RGB')42 return my_transforms(image).unsqueeze(0)43 44 45def get_prediction(data):46 tensor = transform_image(data)47 # model = app.package['model']48 with torch.no_grad():49 prediction = model(tensor)50 prediction = reverse_mapping[prediction.argmax().item()]51 return prediction52 53ALLOWED_EXTENSIONS = {'png', 'jpg', 'jpeg'}54 55def allowed_file(filename):56 return '.' in filename and filename.rsplit('.', 1)[1].lower() in ALLOWED_EXTENSIONS57 58 59# @app.get("/predict")60# async def predict(file: UploadFile = File(...)):61# """62# Perform prediction on the uploaded image63# """64 65# logger.info('API predict called')66 67# if not allowed_file(file.filename):68# raise HTTPException(status_code=400, detail="Format not supported")69 70# try:71# img_bytes = await file.read()72# class_name = get_prediction(img_bytes)73# logger.info(f'Prediction: {class_name}')74# return JSONResponse(content={"class_name": class_name})75# except Exception as e:76# logger.error(f'Error: {str(e)}')77# return JSONResponse(content={"error": str(e), "trace": traceback.format_exc()}, status_code=500)78 79 80# # @app.get("/")81# # def greet_json():82# # return {"Hello": "World!"}83 84 85import torch86import requests87from PIL import Image88from torchvision import transforms89 90# model = torch.hub.load('pytorch/vision:v0.6.0', 'resnet18', pretrained=True).eval()91 92 93# Download human-readable labels for ImageNet.94# response = requests.get("https://git.io/JJkYN")95# labels = response.text.split("\n")96 97model = load_model()98 99augs = transforms.Compose([100 transforms.Resize((224, 224)),101 transforms.ToTensor()102])103 104def predict(inp):105 # inp = transforms.Resize((224, 224))(inp).transforms.ToTensor()(inp).unsqueeze(0)106 inp = augs(inp).unsqueeze(0)107 with torch.no_grad():108 prediction = torch.nn.functional.softmax(model(inp)[0], dim=0)109 confidences = {labels[i]: float(prediction[i]) for i in range(3)}110 # prediction = reverse_mapping[prediction]111 return confidences 112 113 114import gradio as gr115 116gr.Interface(fn=predict,117 inputs=gr.Image(type="pil"),118 outputs=gr.Label(num_top_classes=3)).launch(share=True)