MaskCut / app.py
makakwastaken's picture
Removed wrong input
f4b57bf
Raw
History Blame Contribute Delete
9.47 kB
#!/usr/bin/env python
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from fastapi.middleware.cors import CORSMiddleware
from fastapi.logger import logger
# General
from config import CONFIG
from pydantic import BaseModel
from PIL import Image
# Connection Manager
from connectionManager import ConnectionManager
# CutLER Model
from model import Model
import base64
from io import BytesIO
from predict import predict
# Stable Diffusion Inpainting Model
# from diffusers import StableDiffusionInpaintPipeline
# About
import torch
import os
import sys
# Server API
import uvicorn
app = FastAPI(
title="AdVisual Model Hosting",
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
if CONFIG['ENV'] == 'development':
app.add_middleware(CORSMiddleware, allow_origins=["*"])
else:
app.add_middleware(CORSMiddleware, allow_origins=["https://advisual.io"])
@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 CutLER model
model = Model(CONFIG['DEVICE'])
# Initialize the stable-diffusion-inpainting model
# pipe = StableDiffusionInpaintPipeline.from_pretrained("stabilityai/stable-diffusion-2-inpainting", safety_checker=None)
# pipe.to(CONFIG['DEVICE'])
# Initialize the connection manager
connectionManager = ConnectionManager()
# add model and other preprocess tools too app state
app.package = {
"model": model,
"connectionManager": connectionManager,
# "pipe": pipe
}
@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')
}
# def resize_image(img, height=512, width=512):
# '''Resize image to `size`'''
# size = (width, height)
# img_resized = img.resize(size, Image.ANTIALIAS)
# return img_resized
# def crop_image(img, d=64):
# '''Make dimensions divisible by `d`'''
# new_size = (img.size[0] - img.size[0] % d,
# img.size[1] - img.size[1] % d)
# bbox = [
# int((img.size[0] - new_size[0])/2),
# int((img.size[1] - new_size[1])/2),
# int((img.size[0] + new_size[0])/2),
# int((img.size[1] + new_size[1])/2),
# ]
# img_cropped = img.crop(bbox)
# return img_cropped
# class InpaintBody(BaseModel):
# image: str
# mask: str
# prompt: str
# @app.post("/inpaint")
# async def do_inpaint(body: InpaintBody):
# """
# Perform inpainting on input data
# """
# logger.info('API inpaint called')
# image_data = body.image
# mask_data = body.mask
# prompt = body.prompt
# # Extract base64 from mask and convert to PIL.Image
# if (',' in image_data):
# image = Image.open(BytesIO(base64.b64decode(image_data.split(',')[1])))
# else:
# image = Image.open(BytesIO(base64.b64decode(image_data)))
# # Extract base64 from mask and convert to PIL.Image
# if (',' in mask_data):
# mask = Image.open(BytesIO(base64.b64decode(mask_data.split(',')[1])))
# else:
# mask = Image.open(BytesIO(base64.b64decode(mask_data)))
# # Resize image and mask to 512x512
# image = crop_image(resize_image(image, 512, 512))
# mask = crop_image(resize_image(image, 512, 512))
# pipe = app.package.get('pipe')
# result = pipe(prompt=prompt, image=image, mask_image=mask, num_inference_steps=10, num_images_per_prompt=1)
# images = result['images']
# return images
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-inpaint")
# async def inpaint_websocket_endpoint(websocket: WebSocket):
# connectionManager = app.package.get('connectionManager')
# await connectionManager.connect(websocket)
# await connectionManager.send_json({"ok": True, "status": "CONNECTED"}, websocket)
# while True:
# try:
# data: ImageBody = await connectionManager.receive_json(websocket)
# if (data is None):
# # Wait for data
# if not connectionManager.isConnected(websocket):
# break
# if connectionManager.shouldDisconnect(websocket):
# await websocket.close()
# connectionManager.disconnect(websocket)
# break
# continue
# image_data: str = data.get('image')
# mask_data: str = data.get('mask')
# prompt: str = data.get('prompt')
# await connectionManager.send_json({"ok": True, "status": "STARTED"}, websocket)
# # Extract base64 from mask and convert to PIL.Image
# if (',' in image_data):
# image = Image.open(BytesIO(base64.b64decode(image_data.split(',')[1])))
# else:
# image = Image.open(BytesIO(base64.b64decode(image_data)))
# # Extract base64 from mask and convert to PIL.Image
# if (',' in mask_data):
# mask = Image.open(BytesIO(base64.b64decode(mask_data.split(',')[1])))
# else:
# mask = Image.open(BytesIO(base64.b64decode(mask_data)))
# # Resize image and mask to 512x512
# image = crop_image(resize_image(image, 512, 512))
# mask = crop_image(resize_image(image, 512, 512))
# pipe = app.package.get('pipe')
# result = pipe(prompt=prompt, image=image, mask_image=mask, num_inference_steps=20, num_images_per_prompt=1)
# images = result['images']
# # Convert the result to base64 and send the json back
# result_array = []
# for image in images:
# buffered = BytesIO()
# image.save(buffered, format="JPEG")
# img_str = 'data:image/jpeg;base64,' + base64.b64encode(buffered.getvalue()).decode("utf-8")
# result_array.append(img_str)
# await connectionManager.send_json({"ok": True, "status": "FINISHED", "result": result_array}, websocket)
# await websocket.close()
# connectionManager.disconnect(websocket)
# except WebSocketDisconnect:
# connectionManager.disconnect(websocket)
# break
@app.websocket("/ws")
async def websocket_endpoint(websocket: WebSocket):
connectionManager = app.package.get('connectionManager')
await connectionManager.connect(websocket)
await connectionManager.send_json({"ok": True, "status": "CONNECTED"}, websocket)
while True:
try:
data: ImageBody = await connectionManager.receive_json(websocket)
if (data is None):
# Wait for data
if not connectionManager.isConnected(websocket):
break
if connectionManager.shouldDisconnect(websocket):
await websocket.close()
connectionManager.disconnect(websocket)
break
continue
image: str = data.get('image')
threshold: float = data.get('threshold') or 0.15
num_objects: int = data.get('num_objects') or 1
await connectionManager.send_json({"ok": True, "status": "STARTED"}, websocket)
# 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 connectionManager.send_json({"ok": True, "status": "FINISHED", "result": img_str}, websocket)
await websocket.close()
connectionManager.disconnect(websocket)
except WebSocketDisconnect:
connectionManager.disconnect(websocket)
break
if __name__ == '__main__':
# server api
uvicorn.run("app:app", host="0.0.0.0", port=7860, reload=True)