CoolFace
Apppublic

dmusingu/CXR-IMAGE-Classification

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
app.py118 linesDownload Raw Back to root
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)