From 74154d51b0d0ed1002e51d7b5bdce20bdec46ee1 Mon Sep 17 00:00:00 2001 From: Radionic Date: Tue, 19 Sep 2023 16:50:06 +0800 Subject: [PATCH] fix: segment not updated --- js/ImageEditor.js | 22 ++++++++++++++-------- js/SideBar.js | 3 ++- js/index.js | 2 ++ 3 files changed, 18 insertions(+), 9 deletions(-) diff --git a/js/ImageEditor.js b/js/ImageEditor.js index 3223d6f..05aebd2 100644 --- a/js/ImageEditor.js +++ b/js/ImageEditor.js @@ -71,7 +71,7 @@ function handleImageSize(image) { return { height: h, width: w, samScale }; } -function getClicks() { +export function getClicks() { return imagePrompts.val.map((point) => ({ x: point.x, y: point.y, @@ -79,13 +79,19 @@ function getClicks() { })); } -function drawSegment(clicks) { - runONNX(clicks, embeddings.val).then((mask) => { - const canvas = document.getElementById("mask-canvas"); - const ctx = canvas.getContext("2d"); +export function drawSegment(clicks) { + const canvas = document.getElementById("mask-canvas"); + const ctx = canvas.getContext("2d"); + if (clicks.length === 0) { ctx.clearRect(0, 0, canvas.width, canvas.height); - ctx.drawImage(mask, 0, 0); - }); + return; + } + if (embeddings.val) { + runONNX(clicks, embeddings.val).then((mask) => { + ctx.clearRect(0, 0, canvas.width, canvas.height); + ctx.drawImage(mask, 0, 0); + }); + } } initModel(); @@ -96,7 +102,7 @@ export function ImageEditor() { if (showImageEditor.val && e.code === "Tab") { e.preventDefault(); realTimeSegment = !realTimeSegment; - if (!realTimeSegment && embeddings.val) { + if (!realTimeSegment) { drawSegment(getClicks()); } } diff --git a/js/SideBar.js b/js/SideBar.js index db57057..6708cec 100644 --- a/js/SideBar.js +++ b/js/SideBar.js @@ -1,4 +1,4 @@ -import { updateImagePrompts } from "./ImageEditor.js"; +import { drawSegment, getClicks, updateImagePrompts } from "./ImageEditor.js"; import { imagePrompts, selectedLayer, @@ -49,6 +49,7 @@ export function SideBar() { onclick: () => { selectedLayer.val = key; imagePrompts.val = imagePromptsMulti.val[key]; + drawSegment(getClicks()); }, }, key, diff --git a/js/index.js b/js/index.js index 8c77cf2..431c2dd 100644 --- a/js/index.js +++ b/js/index.js @@ -16,6 +16,7 @@ import { api } from "./api.js"; import { Container } from "./Container.js"; import { loadNpyTensor } from "./onnx.js"; import "https://code.iconify.design/3/3.1.0/iconify.min.js"; +import { drawSegment, getClicks } from "./ImageEditor.js"; /** @type {import( '../../../web/types/litegraph.js').LGraphGroup} */ const recomputeInsideNodesOps = LGraphGroup.prototype.recomputeInsideNodes; @@ -267,6 +268,7 @@ function showMyImageEditor(node) { ); loadNpyTensor(embeedingUrl).then((tensor) => { embeddings.val = tensor; + drawSegment(getClicks()); }); targetNode.val = node; }