diff --git a/js/ImageEditor.js b/js/ImageEditor.js index c301740..66e9e22 100644 --- a/js/ImageEditor.js +++ b/js/ImageEditor.js @@ -1,4 +1,5 @@ import { SideBar } from "./SideBar.js"; +import { initModel, runONNX } from "./onnx.js"; import { showImageEditor, point_label, @@ -9,9 +10,10 @@ import { imageSize, selectedLayer, imagePromptsMulti, + embeddings, } from "./state.js"; import { van } from "./van.js"; -const { button, div, img } = van.tags; +const { button, div, img, canvas } = van.tags; function updateImagePrompts() { if (selectedLayer.val !== "" && selectedLayer.val !== undefined) { @@ -56,11 +58,21 @@ function handlePointClick(e, point) { updateImagePrompts(); } +function handleImageSize(image) { + // Input images to SAM must be resized so the longest side is 1024 + const LONG_SIDE_LENGTH = 1024; + let w = image.naturalWidth; + let h = image.naturalHeight; + const samScale = LONG_SIDE_LENGTH / Math.max(h, w); + return { height: h, width: w, samScale }; +} + export function ImageEditor() { + initModel(); return div( { class: () => - "absolute flex bg-gray-900 bg-opacity-50 top-0 w-full h-full pointer-events-auto " + + "absolute flex bg-gray-900 bg-opacity-50 top-0 w-full h-full pointer-events-auto" + (showImageEditor.val ? "" : "hidden"), }, button( @@ -99,22 +111,23 @@ export function ImageEditor() { ), div( { - class: - "flex items-center justify-center absolute h-full left-0 right-0 bottom-0 top-0 mx-auto my-auto max-h-[1000px] ", + class: "flex items-center justify-center w-full h-full", }, img({ - class: "w-fit h-full", + class: + "fixed top-1/2 left-1/2 transform -translate-x-1/2 -translate-y-1/2", src: imageUrl, onload: (e) => { - imageSize.val = { - width: e.target.naturalWidth, - height: e.target.naturalHeight, - }; + imageSize.val = handleImageSize(e.target); imageContainerSize.val = { width: e.target.offsetWidth, height: e.target.offsetHeight, }; + + const canvas = document.getElementById("mask-canvas"); + canvas.width = e.target.naturalWidth; + canvas.height = e.target.naturalHeight; }, oncontextmenu: (e) => { e.preventDefault(); @@ -124,6 +137,25 @@ export function ImageEditor() { onclick: (e) => { handleClick(e); }, + onmousemove: (e) => { + // console.log('Yo', e); + }, + }), + canvas({ + class: + "fixed top-1/2 left-1/2 transform -translate-x-1/2 -translate-y-1/2", + id: "mask-canvas", + onclick: (e) => { + if (embeddings.val) { + const clicks = [{ x: e.clientX, y: e.clientY, clickType: 1 }]; + runONNX(clicks, embeddings.val).then((mask) => { + const canvas = document.getElementById("mask-canvas"); + const ctx = canvas.getContext("2d"); + ctx.clearRect(0, 0, canvas.width, canvas.height); + ctx.drawImage(mask, 0, 0); + }); + } + }, }), () => { return div( diff --git a/js/index.js b/js/index.js index 86f342e..7d82a8b 100644 --- a/js/index.js +++ b/js/index.js @@ -5,8 +5,7 @@ import { imageUrl, imagePrompts, targetNode, - selectedLayer, - imagePromptsMulti, + embeddings, } from "./state.js"; import { van } from "./van.js"; import { app } from "./app.js"; @@ -252,6 +251,14 @@ function showMyImageEditor(node) { connectedImageFileName )}&type=input&subfolder=${split.length > 1 ? split[0] : ""}` ); + const embeedingUrl = api.apiURL( + `/view?filename=${encodeURIComponent( + "tmp_emb.npy" + )}&type=output&subfolder=${split.length > 1 ? split[0] : ""}` + ); + loadNpyTensor(embeedingUrl).then((tensor) => { + embeddings.val = tensor; + }); targetNode.val = node; } diff --git a/js/onnx.js b/js/onnx.js new file mode 100644 index 0000000..0db1a46 --- /dev/null +++ b/js/onnx.js @@ -0,0 +1,64 @@ +import "https://cdn.jsdelivr.net/npm/onnxruntime-web/dist/ort.min.js"; +import npyjs from "https://esm.sh/npyjs"; +import { imageSize } from "./state.js"; +import { modelData, onnxMaskToImage } from "./onnx_helper.js"; + +ort.env.wasm.wasmPaths = "https://cdn.jsdelivr.net/npm/onnxruntime-web/dist/"; + +// Define image, embedding and model paths +const IMAGE_PATH = "/assets/data/dogs.jpg"; +const IMAGE_EMBEDDING = "/assets/data/dogs_embedding.npy"; +const MODEL_DIR = "http://127.0.0.1:8188/sam_model"; + +export let model = null; + +// Initialize the ONNX model +export const initModel = async () => { + try { + if (MODEL_DIR === undefined) return; + const URL = MODEL_DIR; + model = await ort.InferenceSession.create(URL); + } catch (e) { + console.log(e); + } +}; + +export const loadNpyTensor = async (tensorFile, dType = "float32") => { + let npLoader = new npyjs(); + console.log('tensorFile', tensorFile); + const npArray = await npLoader.load(tensorFile); + console.log('np array', npArray); + const tensor = new ort.Tensor(dType, npArray.data, npArray.shape); + return tensor; +}; + +export const runONNX = async (clicks, tensor) => { + console.log('tensor', tensor); + try { + if ( + model === null || + clicks === null || + tensor === null || + imageSize.val === null + ) + return; + else { + // Preapre the model input in the correct format for SAM. + // The modelData function is from onnxModelAPI.tsx. + const feeds = modelData({ + clicks, + tensor, + modelScale: imageSize.val, + }); + if (feeds === undefined) return; + // Run the SAM ONNX model with the feeds returned from modelData() + const results = await model.run(feeds); + const output = results[model.outputNames[0]]; + // The predicted mask returned from the ONNX model is an array which is + // rendered as an HTML image using onnxMaskToImage() from maskUtils.tsx. + return onnxMaskToImage(output.data, output.dims[2], output.dims[3]); + } + } catch (e) { + console.log(e); + } +}; diff --git a/js/onnx_helper.js b/js/onnx_helper.js new file mode 100644 index 0000000..3ab436a --- /dev/null +++ b/js/onnx_helper.js @@ -0,0 +1,109 @@ +// Copyright (c) Meta Platforms, Inc. and affiliates. +// All rights reserved. + +// This source code is licensed under the license found in the +// LICENSE file in the root directory of this source tree. + +const modelData = ({ clicks, tensor, modelScale }) => { + const imageEmbedding = tensor; + let pointCoords; + let pointLabels; + let pointCoordsTensor; + let pointLabelsTensor; + + // Check there are input click prompts + if (clicks) { + let n = clicks.length; + + // If there is no box input, a single padding point with + // label -1 and coordinates (0.0, 0.0) should be concatenated + // so initialize the array to support (n + 1) points. + pointCoords = new Float32Array(2 * (n + 1)); + pointLabels = new Float32Array(n + 1); + + // Add clicks and scale to what SAM expects + for (let i = 0; i < n; i++) { + pointCoords[2 * i] = clicks[i].x * modelScale.samScale; + pointCoords[2 * i + 1] = clicks[i].y * modelScale.samScale; + pointLabels[i] = clicks[i].clickType; + } + + // Add in the extra point/label when only clicks and no box + // The extra point is at (0, 0) with label -1 + pointCoords[2 * n] = 0.0; + pointCoords[2 * n + 1] = 0.0; + pointLabels[n] = -1.0; + + // Create the tensor + pointCoordsTensor = new ort.Tensor("float32", pointCoords, [1, n + 1, 2]); + pointLabelsTensor = new ort.Tensor("float32", pointLabels, [1, n + 1]); + } + const imageSizeTensor = new ort.Tensor("float32", [ + modelScale.height, + modelScale.width, + ]); + + if (pointCoordsTensor === undefined || pointLabelsTensor === undefined) + return; + + // There is no previous mask, so default to an empty tensor + const maskInput = new ort.Tensor( + "float32", + new Float32Array(256 * 256), + [1, 1, 256, 256] + ); + // There is no previous mask, so default to 0 + const hasMaskInput = new ort.Tensor("float32", [0]); + + return { + image_embeddings: imageEmbedding, + point_coords: pointCoordsTensor, + point_labels: pointLabelsTensor, + orig_im_size: imageSizeTensor, + mask_input: maskInput, + has_mask_input: hasMaskInput, + }; +}; + +// Convert the onnx model mask prediction to ImageData +function arrayToImageData(input, width, height) { + const [r, g, b, a] = [0, 114, 189, 255]; // the masks's blue color + const arr = new Uint8ClampedArray(4 * width * height).fill(0); + for (let i = 0; i < input.length; i++) { + // Threshold the onnx model mask prediction at 0.0 + // This is equivalent to thresholding the mask using predictor.model.mask_threshold + // in python + if (input[i] > 0.0) { + arr[4 * i + 0] = r; + arr[4 * i + 1] = g; + arr[4 * i + 2] = b; + arr[4 * i + 3] = a; + } + } + return new ImageData(arr, height, width); +} + +// Use a Canvas element to produce an image from ImageData +function imageDataToImage(imageData) { + const canvas = imageDataToCanvas(imageData); + const image = new Image(); + image.src = canvas.toDataURL(); + return image; +} + +// Canvas elements can be created from ImageData +function imageDataToCanvas(imageData) { + const canvas = document.createElement("canvas"); + const ctx = canvas.getContext("2d"); + canvas.width = imageData.width; + canvas.height = imageData.height; + ctx?.putImageData(imageData, 0, 0); + return canvas; +} + +// Convert the onnx model mask output to an HTMLImageElement +function onnxMaskToImage(input, width, height) { + return imageDataToImage(arrayToImageData(input, width, height)); +} + +export { modelData, onnxMaskToImage }; diff --git a/js/state.js b/js/state.js index 45a0008..207ac7c 100644 --- a/js/state.js +++ b/js/state.js @@ -36,4 +36,5 @@ export const selectedLayer = van.state(); /** @type {State} */ export const targetNode = van.state(); -export const imageSize = van.state({ width: 0, height: 0 }); +export const imageSize = van.state({ width: 0, height: 0, samScale: 0 }); +export const embeddings = van.state(); diff --git a/sam/sam_node_remote.py b/sam/sam_node_remote.py index 81eaefc..2b653d1 100644 --- a/sam/sam_node_remote.py +++ b/sam/sam_node_remote.py @@ -2,8 +2,12 @@ import requests from PIL import Image import io import numpy as np +import folder_paths class SAM_Embedding: + def __init__(self): + self.output_dir = folder_paths.get_output_directory() + @classmethod def INPUT_TYPES(s): return { @@ -68,6 +72,7 @@ class SAM_Embedding: # print(output) + np.save(f"{self.output_dir}/tmp_emb.npy", output["image_embedding"]) return (output, ) NODE_CLASS_MAPPINGS = {