feat: real-time sam inference
This commit is contained in:
+41
-9
@@ -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(
|
||||
|
||||
+9
-2
@@ -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;
|
||||
}
|
||||
|
||||
|
||||
+64
@@ -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);
|
||||
}
|
||||
};
|
||||
@@ -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 };
|
||||
+2
-1
@@ -36,4 +36,5 @@ export const selectedLayer = van.state();
|
||||
/** @type {State<LGraphNode>} */
|
||||
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();
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
Reference in New Issue
Block a user