From 5ff4ec9884c52eec33dfd6e8c125d6f6fc13ee4d Mon Sep 17 00:00:00 2001 From: Radionic Date: Wed, 22 Nov 2023 18:08:10 +0800 Subject: [PATCH] feat: remote SAM embedding --- js/index.js | 12 ++++-- routes.py | 103 ++++++++++++++++++++++++++++++++++++---------------- 2 files changed, 81 insertions(+), 34 deletions(-) diff --git a/js/index.js b/js/index.js index d217964..1d7ef47 100644 --- a/js/index.js +++ b/js/index.js @@ -24,10 +24,15 @@ import { api } from "./api.js"; import { Container } from "./Container.js"; import { initModel, loadNpyTensor } from "./onnx.js"; import "https://code.iconify.design/3/3.1.0/iconify.min.js"; -import { autoSegment, drawSegment, getClicks, segmented } from "./LayerEditor.js"; +import { + autoSegment, + drawSegment, + getClicks, + segmented, +} from "./LayerEditor.js"; import { infoDialog } from "./dialog.js"; import { sharedAvatarLink } from "./AvatarPreview.js"; -import { updateImagePrompts } from './LayerEditor.js'; +import { updateImagePrompts } from "./LayerEditor.js"; export const generatedImages = {}; @@ -316,6 +321,7 @@ function showMyImageEditor(node) { isGeneratedImage, embedding_id: id, ckpt, + remote: true, }), }) .then(() => { @@ -352,7 +358,7 @@ function showMyImageEditor(node) { embeddings.val = tensor; if (enableAutoSegment.val && !segmented.val) await autoSegment(); drawSegment(getClicks()); - updateImagePrompts() + updateImagePrompts(); }); targetNode.val = node; }) diff --git a/routes.py b/routes.py index 8a6bd17..038f731 100644 --- a/routes.py +++ b/routes.py @@ -10,9 +10,12 @@ import numpy as np import server import re import base64 +from PIL import Image +import io load_dotenv() + # For speeding up ONNX model, see https://github.com/facebookresearch/segment-anything/tree/main/demo#onnx-multithreading-with-sharedarraybuffer def inject_headers(original_handler): async def _handler(request): @@ -36,11 +39,13 @@ for item in server.PromptServer.instance.routes._items: routes.append(item) server.PromptServer.instance.routes._items = routes + @server.PromptServer.instance.routes.get("/avatar-graph-comfyui/tw-styles.css") async def get_web_styles(request): filename = os.path.join(os.path.dirname(__file__), "js/tw-styles.css") return web.FileResponse(filename) + @server.PromptServer.instance.routes.get("/sam_model") async def get_sam_model(request): model_type = request.rel_url.query.get("type", "vit_h") @@ -77,26 +82,56 @@ async def post_sam_model(request): emb_id = post.get("embedding_id") ckpt = post.get("ckpt") ckpt = folder_paths.get_full_path("sams", ckpt) - model_type = re.findall(r'vit_[lbh]', ckpt)[0] + remote = post.get("remote") + model_type = re.findall(r"vit_[lbh]", ckpt)[0] emb_filename = f"{folder_paths.get_output_directory()}/{emb_id}_{model_type}.npy" + output_json_filename = ( + f"{folder_paths.get_output_directory()}/{emb_id}_{model_type}.json" + ) if not os.path.exists(emb_filename): image = load_image(post.get("image"), is_generated_image) - - sam = sam_model_registry[model_type](checkpoint=ckpt) - predictor = SamPredictor(sam) - - image_np = (image * 255).astype(np.uint8) - predictor.set_image(image_np) - emb = predictor.get_image_embedding().cpu().numpy() - np.save(emb_filename, emb) - with open(f"{folder_paths.get_output_directory()}/{emb_id}_{model_type}.json", "w") as f: - json.dump( - { - "input_size": predictor.input_size, - "original_size": predictor.original_size, + if remote: + # Run embed in remote server + image = Image.fromarray((image * 255).astype(np.uint8)) + buffered = io.BytesIO() + image.save(buffered, format="PNG") + image = base64.b64encode(buffered.getvalue()).decode() + res = requests.post( + "https://avatechgg--sam-embed.modal.run", + headers={ + "Content-type": "application/json", + "Accept": "application/json", }, - f, - ) + data=json.dumps({ + "image": image, + }), + ).json() + emb, input_size, original_size = res["emb"], res["input_size"], res["original_size"] + emb = np.array(emb).astype(np.float32) + np.save(emb_filename, emb) + with open(output_json_filename, "w") as f: + data = { + "input_size": input_size, + "original_size": original_size, + } + json.dump(data, f) + else: + sam = sam_model_registry[model_type](checkpoint=ckpt) + predictor = SamPredictor(sam) + + image_np = (image * 255).astype(np.uint8) + predictor.set_image(image_np) + emb = predictor.get_image_embedding().cpu().numpy() + np.save(emb_filename, emb) + with open(output_json_filename, "w") as f: + json.dump( + { + "input_size": predictor.input_size, + "original_size": predictor.original_size, + }, + f, + ) + print("Finished embedding") return web.json_response({}) @@ -110,6 +145,7 @@ async def post_sam_model(request): # response.raise_for_status() # return web.json_response(response.json()) + @server.PromptServer.instance.routes.get("/get_workflow") async def get_workflow(request): name = request.rel_url.query.get("name", "default") @@ -121,7 +157,7 @@ async def get_workflow(request): # else: if name == "default": name = "Auto_segment_workflow" - + workflows_path = os.path.join(os.path.dirname(__file__), "workflow_templates") workflow = json.load(open(f"{workflows_path}/{name}.json")) return web.json_response(workflow) @@ -140,7 +176,7 @@ async def post_segments(request): f.write(base64.b64decode(value.split(",")[1])) order = list(segments.keys()) - with open(os.path.join(output_dir, "order.json") , "w") as f: + with open(os.path.join(output_dir, "order.json"), "w") as f: json.dump(order, f) return web.json_response({}) @@ -156,13 +192,16 @@ async def post_segments(request): # json.dump(order, f) # return web.json_response({}) + @server.PromptServer.instance.routes.get("/get_webhook") async def get_webhook(request): - url = os.getenv('DISCORD_WEBHOOK_URL') + url = os.getenv("DISCORD_WEBHOOK_URL") return web.json_response(url) + import uuid + @server.PromptServer.instance.routes.post("/create_avatar_from_image") async def post_input_file(request): post = await request.read() @@ -173,20 +212,22 @@ async def post_input_file(request): try: queue_id = uuid.uuid4() - - output_dir = os.path.join(folder_paths.base_path, "input","create_avatar_endpoint") + + output_dir = os.path.join( + folder_paths.base_path, "input", "create_avatar_endpoint" + ) os.makedirs(output_dir, exist_ok=True) - - filename = os.path.join(output_dir,str(queue_id) + ".png") + + filename = os.path.join(output_dir, str(queue_id) + ".png") with open(filename, "wb") as f: f.write(post) - - return web.json_response({ - "redirect_url": "https://ai-assistant.avatech.ai?queue-id=" + str(queue_id) - }) + + return web.json_response( + { + "redirect_url": "https://ai-assistant.avatech.ai?queue-id=" + + str(queue_id) + } + ) except Exception as e: print(e) - return web.json_response({ - "error": e - }) - + return web.json_response({"error": e})