feat: real-time sam inference

This commit is contained in:
Radionic
2023-09-18 17:24:31 +08:00
parent eb3e1660d5
commit e79a3fa969
6 changed files with 230 additions and 12 deletions
+41 -9
View File
@@ -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
View File
@@ -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
View File
@@ -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);
}
};
+109
View File
@@ -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
View File
@@ -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();
+5
View File
@@ -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 = {