SteveaWong's picture
Upload app.py via MAS pipeline
ad4072b verified
Raw
History Blame Contribute Delete
1.6 kB
import fastapi
from fastapi import FastAPI
import torch
import torch.nn as nn
import torchvision.transforms as transforms
from PIL import Image
import io
import base64
app = FastAPI()
class CNN(nn.Module):
def __init__(self):
super(CNN, self).__init__()
self.conv1 = nn.Conv2d(3, 6, 5)
self.pool = nn.MaxPool2d(2, 2)
self.fc1 = nn.Linear(6 * 14 * 14, 120)
self.fc2 = nn.Linear(120, 84)
self.fc3 = nn.Linear(84, 3)
def forward(self, x):
x = self.pool(torch.relu(self.conv1(x)))
x = x.view(-1, 6 * 14 * 14)
x = torch.relu(self.fc1(x))
x = torch.relu(self.fc2(x))
x = self.fc3(x)
return x
model = CNN()
model.load_state_dict(torch.load("best_model.pt", map_location=torch.device("cpu")))
model.eval()
data_transforms = transforms.Compose([
transforms.Resize(32),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])
@app.post("/predict")
async def predict(image_base64: str):
image_bytes = base64.b64decode(image_base64)
image = Image.open(io.BytesIO(image_bytes))
image = data_transforms(image)
image = image.unsqueeze(0)
output = model(image)
_, predicted = torch.max(output, 1)
confidence = torch.nn.functional.softmax(output, dim=1)
return {
"predicted_class": predicted.item(),
"confidence": confidence.tolist()[0]
}
@app.get("/health")
async def health():
return {"status": "ok"}
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8000)