| |
|
|
| 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 |
|
|
| |
| import torch |
| import os |
| import sys |
|
|
| 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="/", |
| ) |
|
|
| |
| 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'])) |
|
|
| |
| model = Model() |
|
|
| |
| 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 |
|
|
| |
| result = predict(app.package, image, threshold, num_objects) |
|
|
| |
| 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 |
| |
| |
| result = predict(app.package, image, threshold, num_objects) |
|
|
| |
| 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 |