feat: add bounding box to SAM input
This commit is contained in:
+14
-1
@@ -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
@@ -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
@@ -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
@@ -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([]);
|
||||
|
||||
|
||||
Reference in New Issue
Block a user