#!/usr/bin/env python from fastapi import FastAPI, WebSocket, WebSocketDisconnect from fastapi.middleware.cors import CORSMiddleware from fastapi.logger import logger from model import Model import base64 from io import BytesIO from pydantic import BaseModel from config import CONFIG from predict import predict # About import torch import os import sys # Server Framework import uvicorn app = FastAPI( title="AdVisual MaskCut Model", description="Description of the ML Model", version="0.0.1", terms_of_service=None, contact=None, license_info=None, docs_url="/", ) # Allow CORS for local debugging app.add_middleware(CORSMiddleware, allow_origins=["*"]) @app.on_event("startup") async def startup_event(): """ Initialize FastAPI and add variables """ logger.info('Running envirnoment: {}'.format(CONFIG['ENV'])) logger.info('PyTorch using device: {}'.format(CONFIG['DEVICE'])) # Initialize the pytorch model model = Model() # add model and other preprocess tools too app state app.package = { "model": model } @app.get("/ping") def ping(): return {"ok": True, "message": "Pong"} @app.get("/about") def show_about(): """ Get deployment information, for debugging """ logger.info('API /about called') def bash(command): output = os.popen(command).read() return output return { "sys.version": sys.version, "torch.__version__": torch.__version__, "torch.cuda.is_available()": torch.cuda.is_available(), "torch.version.cuda": torch.version.cuda, "torch.backends.cudnn.version()": torch.backends.cudnn.version(), "torch.backends.cudnn.enabled": torch.backends.cudnn.enabled, "nvidia-smi": bash('nvidia-smi') } class ImageBody(BaseModel): image: str threshold: float = 0.15 num_objects: int = 1 @app.post("/predict") async def do_predict(body: ImageBody): """ Perform prediction on input data """ logger.info('API predict called') image: str = body.image threshold: float = body.threshold num_objects: int = body.num_objects # Run the algorithm result = predict(app.package, image, threshold, num_objects) # Convert the result to base64 and send the json back buffered = BytesIO() result.save(buffered, format="JPEG") img_str = 'data:image/jpeg;base64,' + base64.b64encode(buffered.getvalue()).decode("utf-8") return {"ok": True, "status": "FINISHED", "result": img_str} @app.websocket("/ws") async def websocket_endpoint(websocket: WebSocket): await websocket.accept() while True: try: data = await websocket.receive_json() image: str = data.get('image') threshold: float = data.get('threshold') or 0.15 num_objects: int = data.get('num_objects') or 1 await websocket.send_json({"ok": True, "status": "STARTED"}) if image == None: await websocket.send_json({ "ok": False, "status": "ERROR", "message": "No image provided" }) break # Run the algorithm result = predict(app.package, image, threshold, num_objects) # Convert the result to base64 and send the json back buffered = BytesIO() result.save(buffered, format="JPEG") img_str = 'data:image/jpeg;base64,' + base64.b64encode(buffered.getvalue()).decode("utf-8") await websocket.send_json({"ok": True, "status": "FINISHED", "result": img_str}) await websocket.close() except WebSocketDisconnect: break if __name__ == '__main__': # server api uvicorn.run("main:app", host="0.0.0.0", port=8000, reload=True)