feat: add bounding box to SAM input

This commit is contained in:
Radionic
2023-11-21 15:49:48 +08:00
parent b17e2ea383
commit b36bea3694
4 changed files with 55 additions and 16 deletions
+14 -1
View File
@@ -16,6 +16,7 @@ import {
embeddingID,
alertDialog,
allImagePrompts,
boxesMulti,
} from "./state.js";
import { van } from "./van.js";
import vision from "https://cdn.jsdelivr.net/npm/@mediapipe/tasks-vision@0.10.3";
@@ -227,6 +228,14 @@ export async function autoSegment() {
...positivePoints,
...negativePoints,
];
// Find bounding box of positive points
const box = {
x1: Math.min(...negativePoints.map((x) => x.x)),
y1: Math.min(...negativePoints.map((x) => x.y)),
x2: Math.max(...negativePoints.map((x) => x.x)),
y2: Math.max(...negativePoints.map((x) => x.y)),
};
boxesMulti.val[key] = box;
}
});
imagePrompts.val = imagePromptsMulti.val[selectedLayer.val];
@@ -374,7 +383,11 @@ export async function drawSegment(clicks) {
return;
}
if (embeddings.val) {
const mask = await runONNX(clicks, embeddings.val);
const mask = await runONNX(
clicks,
embeddings.val,
boxesMulti.val[selectedLayer.val]
);
if (mask) {
ctx.clearRect(0, 0, canvas.width, canvas.height);
ctx.drawImage(mask, 0, 0);
+2 -1
View File
@@ -27,7 +27,7 @@ export const loadNpyTensor = async (tensorFile, dType = "float32") => {
return tensor;
};
export const runONNX = async (clicks, tensor) => {
export const runONNX = async (clicks, tensor, box) => {
// console.log('tensor', tensor);
try {
if (
@@ -44,6 +44,7 @@ export const runONNX = async (clicks, tensor) => {
clicks,
tensor,
modelScale: imageSize.val,
box,
});
if (feeds === undefined) return;
// Run the SAM ONNX model with the feeds returned from modelData()
+21 -10
View File
@@ -4,7 +4,7 @@
// 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 modelData = ({ clicks, tensor, modelScale, box }) => {
const imageEmbedding = tensor;
let pointCoords;
let pointLabels;
@@ -18,8 +18,9 @@ const modelData = ({ clicks, tensor, modelScale }) => {
// 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);
const numPoints = box ? n + 3 : n + 1;
pointCoords = new Float32Array(2 * numPoints);
pointLabels = new Float32Array(numPoints);
// Add clicks and scale to what SAM expects
for (let i = 0; i < n; i++) {
@@ -28,15 +29,25 @@ const modelData = ({ clicks, tensor, modelScale }) => {
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;
if (box) {
pointCoords[2 * n] = box.x1;
pointCoords[2 * n + 1] = box.y1;
pointLabels[n] = 2;
pointCoords[2 * n + 2] = box.x2;
pointCoords[2 * n + 3] = box.y2;
pointLabels[n] = 3;
} else {
// 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]);
pointCoordsTensor = new ort.Tensor("float32", pointCoords, [1, numPoints, 2]);
pointLabelsTensor = new ort.Tensor("float32", pointLabels, [1, numPoints]);
}
const imageSizeTensor = new ort.Tensor("float32", [
modelScale.height,
+18 -4
View File
@@ -10,6 +10,12 @@
* @property {number} x - The x coordinate
* @property {number} y - The y coordinate
* @property {>} label - The label
*
* @typedef {Object} Box
* @property {number} x1
* @property {number} y1
* @property {number} x2
* @property {number} y2
*/
import { van } from "./van.js";
@@ -18,13 +24,15 @@ export const iframeSrc = van.state("https://editor.avatech.ai?comfyui=true");
export const showEditor = van.state(false);
// localStorage.getItem("showPreview") == 'true'
export const showPreview = van.state(true);
export const previewUrl = van.state("https://editor.avatech.ai/viewer?avatarId=default&debug=false&width=400&height=400&hideTrigger=true&voiceSelection=true&hideUI=true");
export const previewUrl = van.state(
"https://editor.avatech.ai/viewer?avatarId=default&debug=false&width=400&height=400&hideTrigger=true&voiceSelection=true&hideUI=true"
);
export const previewImg = van.state("");
export const previewImgLoading = van.state(false);
export const enableAutoSegment = van.state(false);
// export const previewUrl = van.state("http://localhost:3006/viewer?avatarId=default&hideUI=true&debug=true&width=300&height=300&showAudioControl=true");
export const isDirty = van.state(false);
export const fileName = van.state('');
export const fileName = van.state("");
export const showImageEditor = van.state(false);
export const showLoading = van.state(false);
export const alertDialog = van.state({
@@ -32,9 +40,9 @@ export const alertDialog = van.state({
time: 0,
});
export const shareLoading = van.state(false);
export const previewModelId = van.state('');
export const previewModelId = van.state("");
export const isGenerateFlow = van.state(false)
export const isGenerateFlow = van.state(false);
export const loadingCaption = van.state("");
export const imageUrl = van.state("");
@@ -44,6 +52,12 @@ export const imageContainerSize = van.state({
height: 0,
});
/** @type {State<Box>} */
export const boxes = van.state();
/** @type {State<Record<string, Box>>} */
export const boxesMulti = van.state({});
/** @type {State<Point[]>} */
export const imagePrompts = van.state([]);