Merge branch 'app' into dev

# Conflicts:
#	js/LayerEditor.js
#	js/index.js
#	routes.py
#	sam/sam_multilayer.py
This commit is contained in:
EdwinWong
2023-12-20 18:28:47 +08:00
29 changed files with 3808 additions and 920 deletions
+2
View File
@@ -1,7 +1,9 @@
# Created by https://www.toptal.com/developers/gitignore/api/node,python,react
# Edit at https://www.toptal.com/developers/gitignore?templates=node,python,react
workflow_templates/
js/output.css
*.task
### Node ###
# Logs
+45 -15
View File
@@ -43,7 +43,29 @@ def append_to_sys_path(path):
if path not in sys.path:
sys.path.append(path)
folder_paths.folder_names_and_paths["sams"] = ([os.path.join(folder_paths.models_dir, "sams")], folder_paths.supported_pt_extensions)
folder_paths.folder_names_and_paths["sams"] = (
[os.path.join(folder_paths.models_dir, "sams")],
folder_paths.supported_pt_extensions,
)
def download_model(url, save_path):
response = requests.get(url, stream=True)
response.raise_for_status()
file_size = int(response.headers.get("Content-Length", 0))
chunk_size = 1024
num_bars = int(file_size / chunk_size)
with open(save_path, "wb") as f:
for chunk in tqdm(
response.iter_content(chunk_size=chunk_size),
total=num_bars,
unit="KB",
desc=url.split("/")[-1],
):
f.write(chunk)
def download_sam_model():
model_dir = get_folder_paths("sams")[0]
@@ -56,24 +78,32 @@ def download_sam_model():
if "sam_vit_h_4b8939.pth" not in files:
print("Downloading sam model...")
url = "https://dl.fbaipublicfiles.com/segment_anything/sam_vit_h_4b8939.pth"
response = requests.get(url, stream=True)
response.raise_for_status()
file_size = int(response.headers.get("Content-Length", 0))
chunk_size = 1024
num_bars = int(file_size / chunk_size)
with open(f"{model_dir}/sam_vit_h_4b8939.pth", "wb") as f:
for chunk in tqdm(
response.iter_content(chunk_size=chunk_size),
total=num_bars,
unit="KB",
desc=url.split("/")[-1],
):
f.write(chunk)
download_model(url, f"{model_dir}/sam_vit_h_4b8939.pth")
download_sam_model()
def download_face_and_pose_landmarker():
model_dir = os.path.join(ag_path, "mediapipe_models")
if not os.path.isdir(model_dir):
os.makedirs(model_dir)
model_path = os.path.join(model_dir, "face_landmarker.task")
if not os.path.isfile(model_path):
print("Downloading face landmarker model...")
url = "https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/latest/face_landmarker.task"
download_model(url, model_path)
model_path = os.path.join(model_dir, "pose_landmarker_full.task")
if not os.path.isfile(model_path):
print("Downloading pose landmarker model...")
url = "https://storage.googleapis.com/mediapipe-models/pose_landmarker/pose_landmarker_full/float16/latest/pose_landmarker_full.task"
download_model(url, model_path)
download_face_and_pose_landmarker()
paths = ["blender", "sam"]
files = []
+2 -2
View File
@@ -100,14 +100,14 @@ class ObjectOps:
import global_bpy
bpy = global_bpy.get_bpy()
if props.get("BPY_OBJ") != None:
if props.get("BPY_OBJ") is not None:
bpy.context.view_layer.objects.active = props["BPY_OBJ"]
results = self.blender_process(bpy, **props)
if results is None:
# print(results)
if props.get("BPY_OBJ") != None:
if props.get("BPY_OBJ") is not None:
return (props["BPY_OBJ"], )
else:
return (bpy.context.view_layer.objects.active, )
+3
View File
@@ -14,6 +14,9 @@ def genreate_mesh_from_texture(bpy, image):
contours, _ = cv2.findContours(
gray, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
if len(contours) == 0:
raise Exception("No contours found. Please ensure that the image has the correct segments (e.g. when you click on the mouth, it should display a proper blue area over the mouth region).")
# Get the largest contour
areas = [cv2.contourArea(contour) for contour in contours]
+2 -1
View File
@@ -8,5 +8,6 @@ class GroupOps(blender_node.ObjectOps):
RETURN_TYPES = (blender_node.BPY_OBJS,)
def blender_process(self, bpy, BPY_OBJ, BPY_OBJ2, **props):
return ([BPY_OBJ, BPY_OBJ2],)
prop_values = props.values()
return ([BPY_OBJ, BPY_OBJ2, *prop_values],)
+4 -2
View File
@@ -11,8 +11,10 @@ class Mesh_JoinMesh(blender_node.ObjectOps):
def blender_process(self, bpy, BPY_OBJ, **props):
prop_values = props.values()
for obj in list(prop_values) + [BPY_OBJ]:
obj.select_set(True)
bpy.context.view_layer.objects.active = BPY_OBJ
if obj is not None:
obj.select_set(True)
if bpy.context.view_layer.objects is not None:
bpy.context.view_layer.objects.active = BPY_OBJ
bpy.ops.object.join()
return (BPY_OBJ,)
+24
View File
@@ -0,0 +1,24 @@
import blender_node
from mesh_utils import genreate_mesh_from_texture, assign_texture
class Object_UV_Modifier(blender_node.EditOps):
EXTRA_INPUT_TYPES = {
"scale": ('FLOAT', {'default': 1, "display": "number", "step": 0.01}),
"texture_name": ('STRING', {'default': 'Texture', })
}
CUSTOM_NAME = "UV Modifier"
def blender_process(self, bpy, BPY_OBJ, scale, texture_name):
import bmesh
bm = bmesh.from_edit_mesh(BPY_OBJ.data)
uv_layer = bm.loops.layers.uv.verify()
for f in bm.faces:
# move all of the UVs in this face up one UDIM tile
for l in f.loops:
l[uv_layer].uv = (l[uv_layer].uv[0], 0.998 if l[uv_layer].uv[1] == 1 else l[uv_layer].uv[1])
bmesh.update_edit_mesh(BPY_OBJ.data)
+19 -3
View File
@@ -3,8 +3,25 @@ import { van } from "./van.js";
const { div, span } = van.tags;
export function Alert() {
const color = van.state("bg-orange-100 text-orange-700 border-orange-500");
van.derive(() => {
if (alertDialog.val.time > 0) {
switch (alertDialog.val.type) {
case "error":
color.val = "bg-red-100 text-red-700 border-red-500";
break;
case "success":
color.val = "bg-green-100 text-green-700 border-green-500";
break;
case "info":
color.val = "bg-blue-100 text-blue-700 border-blue-500";
break;
case "warning":
default:
color.val = "bg-orange-100 text-orange-700 border-orange-500";
break;
}
setTimeout(() => {
alertDialog.val = { text: "", time: 0 };
}, alertDialog.val.time);
@@ -14,13 +31,12 @@ export function Alert() {
return div(
{
class: () =>
"absolute bottom-8 flex justify-center w-full " +
"absolute z-[100] bottom-8 flex justify-center w-full " +
(alertDialog.val.text ? "" : "hidden"),
},
div(
{
class:
"bg-orange-100 border-t-4 border-orange-500 rounded-sm text-orange-700 p-2",
class: () => `${color.val} border-t-4 rounded-sm p-2`,
},
() => span(alertDialog.val.text)
)
+28
View File
@@ -0,0 +1,28 @@
import { van } from "./van.js";
const { div, span } = van.tags;
export function AppHeader() {
return div(
{
class: () => "absolute flex justify-between top-0 w-full text-white p-4",
},
div(
{},
span(
{
class:
"block bg-gradient-to-b from-gray-500 to-white text-transparent bg-clip-text text-2xl",
},
"Avatech v1"
),
span(
{
class:
"bg-gradient-to-b from-gray-500 to-white text-transparent bg-clip-text text-lg",
},
"Get your DALLE3 AI Personal Clone"
)
),
span({ class: "text-gray-300" }, "Twitter")
);
}
+802 -7
View File
@@ -1,22 +1,817 @@
import { van } from "./van.js";
const { button, iframe, div, img } = van.tags;
import { showEditor, previewUrl, showPreview } from "./state.js";
import {
imageUrl,
showPreview,
previewUrl,
showEditor,
previewImg,
previewImgLoading,
alertDialog,
isGenerateFlow,
enableAutoSegment
} from "./state.js";
const { button, iframe, div, img, input, label, span, textarea, ul, li } =
van.tags;
import { app } from "./app.js";
import { uploadPreview } from "./index.js";
import { api } from "./api.js";
import { segmented, uploadSegments } from "./LayerEditor.js";
import { initModel } from "./onnx.js";
// import { uploadSegments } from "./LayerEditor.js";
const workflowList = [
"idle_avatar_(trigger)",
"Auto_segment_workflow",
"BronyaZaychik_(ChinaDress)",
"BronyaZaychik_(Default_Silverwing)",
"BronyaZaychik_(Non-official_office_ladysuit)",
"BronyaZaychik_(Official_office_ladysuit)",
"BronyaZaychikLora_withhand",
"SilverWolf_(Default)",
"SilverWolf_(Maid)",
"SilverWolfLora_withhand",
];
function editSegment(stage) {
/** @type {import('../../../web/types/litegraph.js').LGraph}*/
const graph = app.graph;
const imageNodes = graph.findNodesByType("LoadImage");
if (!imageNodes[0].imgs) return;
const nodes = graph.findNodesByType("SAM MultiLayer");
/** @type {any[]}*/
const widgets = nodes[0].widgets;
console.log(nodes[0]);
console.log(nodes[0].widgets);
widgets.find((x) => x.type == "button").callback();
stage.val = 2;
}
// const workflowList = ["Auto_segment_workflow"];
/**
* Load JSON workflow
* @param {string} name - The name of the workflow to load
*/
async function loadJSONWorkflow(name) {
if (name === 'default' || name.toLowerCase().startsWith("auto_segment")) {
enableAutoSegment.val = true
} else {
enableAutoSegment.val = false
}
const json = await (await fetch(`./get_workflow?name=${name}`)).json();
app.loadGraphData(json);
console.log(json);
}
async function updatePositivePrompt(app, prompt) {
const positivePrompt = app.graph
.findNodesByType("CLIPTextEncode")
.find((x) => x.color == "#232");
if (!positivePrompt) {
alertDialog.val = {
text: "Cannot find the CLIPTextEncode node. Please make sure the workflow is correct.",
time: 5000,
};
return;
}
positivePrompt.widgets[0].inputEl.value = prompt;
}
async function updateSeedValue(app, seed) {
const kSampler = app.graph.findNodesByType("KSampler")[0];
if (!kSampler) {
alertDialog.val = {
text: "Cannot find the KSampler node. Please make sure the workflow is correct.",
time: 5000,
};
return;
}
kSampler.widgets[0].value = seed;
kSampler.widgets[1].value = "fixed";
}
async function uploadImage() {
/** @type {import('../../../web/types/litegraph.js').LGraph}*/
const graph = app.graph;
const nodes = graph.findNodesByType("LoadImage");
previewImgLoading.val = true;
console.log(previewImgLoading.val);
/** @type {any[]}*/
const widgets = nodes[0].widgets;
console.log(nodes[0]);
widgets.find((x) => x.type == "button").callback();
while (true) {
await new Promise((resolve) => setTimeout(resolve, 1000));
if (nodes[0]?.imgs) {
if (previewImg.val != "" && previewImg.val == nodes[0].imgs[0].currentSrc)
continue;
previewImgLoading.val = false;
return nodes[0].imgs[0].currentSrc;
}
}
}
const jsonWorkflowLoading = van.state(true);
export const sharedAvatarLink = van.state("");
async function prepareImageFromUrlRedirect(stage) {
await new Promise((resolve) => setTimeout(resolve, 2000));
const queue_id = new URLSearchParams(window.location.search).get("queue-id");
if (queue_id && queue_id != "") {
console.log(queue_id);
stage.val = 1;
const graph = app.graph;
const node = graph.findNodesByType("LoadImage");
const imageName = queue_id + ".png";
console.log(node[0]);
node[0].widgets_values[0] = imageName;
node[0].widgets[0].value = imageName;
node[0].widgets[0]._value = imageName;
graph.change();
previewImg.val = api.apiURL(
`/view?filename=${encodeURIComponent(
imageName
)}&type=input&subfolder=create_avatar_endpoint${app.getPreviewFormatParam()}`
);
console.log(previewImg);
}
const dragndrop = document.getElementById("dnd");
dragndrop.addEventListener("dragenter", (evt) => {
evt.preventDefault();
dragndrop.className =
"h-96 w-full border-2 border-purple-500 text-purple-500 border-dashed rounded-lg flex justify-center items-center";
});
dragndrop.addEventListener("dragleave", (evt) => {
evt.preventDefault();
dragndrop.className =
"h-96 w-full border-2 border-black border-dashed items-center rounded-lg flex justify-center";
});
dragndrop.addEventListener("dragover", (evt) => {
evt.preventDefault();
});
dragndrop.addEventListener("drop", async (evt) => {
evt.preventDefault();
dragndrop.className =
"h-96 w-full border-2 border-black border-dashed items-center rounded-lg flex justify-center";
if (evt.dataTransfer.files.length > 1) return;
if (
evt.dataTransfer.files[0].type != "image/jpeg" &&
evt.dataTransfer.files[0].type != "image/png" &&
evt.dataTransfer.files[0].type != "image/webp"
)
return;
stage.val = 1;
previewImg.val = URL.createObjectURL(evt.dataTransfer.files[0]);
if (Object.entries(evt.dataTransfer.files).length) {
await uploadFile(evt.dataTransfer.files[0], true);
}
});
}
export function AvatarPreview() {
console.log("getting workflow json now");
loadJSONWorkflow("default").then(() => {
console.log("done loading");
jsonWorkflowLoading.val = false;
});
return div(
{ class: () => (!showEditor.val ? "" : "hidden") },
iframe({
const loading = van.state(false);
const shareLoading = van.state("share"); // share, loading, shared
api.addEventListener("execution_start", (evt) => {
loading.val = true;
});
api.addEventListener("executed", (evt) => {
const nodeId = evt.detail.node;
const targetNode = graph._nodes_by_id[nodeId];
if (targetNode.type === "AvatarMainOutput") {
loading.val = false;
}
});
const email = van.state("");
const stage = van.state(0); // 0: upload image, 1: edit segment, 2: generate
// This will wait 2 seconds until the everything is loaded
prepareImageFromUrlRedirect(stage);
const renderSteps = () => {
return div(
{
class: () =>
"flex flex-col bg-white justify-center w-[32rem] max-w-[100%]",
},
div(
{
class: () =>
" bg-gradient-to-b from-black via-[#5F5F5F] via-60% to-white text-transparent bg-clip-text font-gabarito text-4xl",
},
"Avatech v1"
),
div(
{
class: () =>
" bg-gradient-to-b from-black via-[#5F5F5F] via-50% to-white text-transparent bg-clip-text font-gabarito text-2xl",
},
"Get your DALLE3 AI Personal Clone"
),
div(
{
class: () =>
" w-full flex flex-col justify-center items-center gap-4",
},
!isGenerateFlow.val
? div(
{
class: () =>
"flex flex-col justify-center items-center gap-4 w-full",
},
div(
{ class: () => "w-full flex mt-2" },
button(
{
class: () => `btn w-full normal-case`,
onclick: async () => {
// previewImg.val = await uploadImage();
// stage.val = 1;
var input = document.createElement("input");
input.type = "file";
document.body.appendChild(input);
// when the input content changes, do something
input.onchange = async function (e) {
stage.val = 1;
if (Object.entries(e.target.files).length) {
await uploadFile(e.target.files[0], true);
}
previewImg.val = URL.createObjectURL(e.target.files[0]);
// upload files
document.body.removeChild(input);
};
// Trigger file browser
input.click();
},
},
div({ class: "badge badge-neutral" }, "1"),
div("Upload your image"),
span({
class: "iconify text-lg",
"data-icon": "material-symbols:drive-folder-upload",
"data-inline": "false",
}),
() =>
previewImgLoading.val
? span({
class: "loading loading-spinner loading-md",
})
: "",
),
),
() => {
const dnd = div(
{
id: "dnd",
class: () =>
"h-96 w-full border-2 border-black border-dashed items-center rounded-lg flex justify-center text-black",
},
"or drag and drop the image here",
);
const image = img({
class: () => "z-[10] object-contain w-full h-[394px] border",
src: previewImg,
onload: () => {
segmented.val = false;
}
});
if (isMobileDevice()) {
return previewImg.val !== "" ? image : "";
} else {
return previewImg.val === "" ? dnd : image;
}
},
button(
{
class: () =>
"btn w-full normal-case " +
(stage.val < 1 ? "btn-disabled" : ""),
onclick: () => {
enableAutoSegment.val = true;
editSegment(stage)
},
},
div({ class: "badge badge-neutral" }, "2"),
"Edit Segment",
),
button(
{
class: () =>
"btn w-full normal-case " +
(stage.val < 2 ? "btn-disabled" : ""),
onclick: async () => {
// const uploaded = await uploadSegments();
// if (!uploaded) return;
const graph = app.graph;
const imageNodes = graph.findNodesByType("LoadImage");
if (!imageNodes[0].imgs) return;
document.getElementById("queue-button").click();
},
},
div({ class: "badge badge-neutral" }, "3"),
() =>
loading.val
? span({
class: "loading loading-spinner loading-md",
})
: "Make It Alive!",
),
)
: div(
{
class:
"flex flex-col justify-center items-center gap-4 w-full text-black",
},
div(
{
class:
"w-full mt-2 flex flex-col rounded-md left-0 top-0",
},
textarea({
class:
"textarea textarea-bordered border-gray-300 border-b-0 focus:outline-none resize-none rounded-t-md rounded-b-none text-md h-36",
placeholder: "Enter your prompt",
defaultValue:
"1girl, looking at viewer, open mouth, simple background, white background, smile",
id: "positivePromptProxy",
}),
div(
{
class:
"flex flex-row gap-2 border border-gray-300 rounded-b-md text-md items-center",
},
span({ class: "ml-4" }, "Seed"),
div({ class: "divider divider-horizontal m-0" }),
input({
type: "text",
class: "input border-none focus:outline-none w-full p-0",
placeholder: "Seed",
defaultValue: "1234",
id: "seedProxy",
}),
div(
{
onclick: () => {
const random4Digits =
Math.floor(Math.random() * 9000) + 1000;
console.log(
random4Digits,
document.getElementById("seedProxy").value,
);
document.getElementById("seedProxy").value =
random4Digits.toString();
},
},
span({
class: "iconify text-2xl mr-4 hover:cursor-pointer",
"data-icon": "fad:random-1dice",
"data-inline": "false",
}),
),
),
),
button(
{
class: "btn w-full normal-case ",
onclick: async () => {
loading.val = true;
updatePositivePrompt(
app,
document.getElementById("positivePromptProxy").value,
);
updateSeedValue(
app,
document.getElementById("seedProxy").value,
);
const sam = app.graph.findNodesByType("SAM MultiLayer")[0];
if (!sam) {
alertDialog.val = {
text: "Cannot find the SAM node. Please make sure the workflow is correct.",
time: 5000,
};
return;
}
const ckpt = sam.widgets[0].value;
const modelType = ckpt.match(/vit_[lbh]/)?.[0];
await initModel(modelType);
await uploadSegments();
document.getElementById("queue-button").click();
},
},
div({ class: "badge badge-neutral" }, "1"),
() =>
loading.val
? span({ class: "loading loading-spinner loading-md" })
: "Make It Alive!",
),
button(
{
class: () =>
"btn w-full normal-case ",
onclick: () => {
enableAutoSegment.val = false;
editSegment(stage)
},
},
div({ class: "badge badge-neutral" }, "2"),
"Edit Segment",
),
// button(
// {
// class: "btn w-full normal-case",
// onclick: () => {
// /** @type {import('../../../web/types/litegraph.js').LGraph}*/
// const graph = app.graph;
// const nodes = graph.findNodesByType("SAM MultiLayer");
// /** @type {any[]}*/
// const widgets = nodes[0].widgets;
// console.log(nodes[0]);
// console.log(nodes[0].widgets);
// widgets.find((x) => x.type == "button").callback();
// },
// },
// div({ class: "badge badge-neutral" }, "2"),
// "(Optional) Edit Segment",
// ),
),
),
);
};
const renderIFrame = () => {
return iframe({
id: "avatech-viewer-iframe",
title: "avatech-viewer-iframe",
name: "avatech-viewer-iframe",
allow: "cross-origin-isolated",
class: () =>
"w-[320px] h-[370px] absolute right-0 top-0 z-[100] pointer-events-auto flex mt-4 mr-4 rounded-2xl border-none " +
"w-full h-full min-w-[400px] min-h-[400px] z-[100] pointer-events-auto flex border-none overflow-hidden" +
(showPreview.val ? "" : "hidden"),
// src: "https://labs.avatech.ai/viewer/default",
// src: "http://localhost:3000/viewer/default",
src: previewUrl,
})
});
};
const renderShareLink = () => {
return div(
{
class: () =>
"w-full flex flex-col gap-2 justify-center items-center mt-8",
},
div(
{
class: () =>
"w-full flex justify-center font-bold italic text-gray-500",
},
span("We are launching OpenAI Assistant API integration soon!")
),
div(
{ class: () => "w-[24rem] flex justify-center items-center" },
input({
type: "text",
class: () =>
"w-full input input-bordered text-black rounded rounded-l-md rounded-r-none !outline-none",
onchange: (e) => {
email.val = e.target.value;
},
placeholder: "Enter your email",
}),
button(
{
class: () =>
"btn rounded rounded-l-none rounded-r-md no-animation bg-neutral-800 hover:bg-neutral-950 text-white border-none normal-case",
onclick: async () => {
if (shareLoading.val === "share") {
shareLoading.val = "loading";
const url = await (await fetch("./get_webhook")).json();
await uploadPreview();
await fetch(url, {
method: "POST",
headers: {
"Content-Type": "application/json",
},
body: JSON.stringify({
username: "Avabot",
avatar_url:
"https://avatech-avatar-dev1.nyc3.cdn.digitaloceanspaces.com/avatechai.png",
content: "New register! \n" + email.val,
}),
});
shareLoading.val = "shared";
}
if (sharedAvatarLink.val) {
await navigator.clipboard.writeText(sharedAvatarLink.val);
alertDialog.val = {
text: "Avatar link copied to clipboard!",
type: "success",
time: 5000,
};
}
},
},
() => {
switch (shareLoading.val) {
case "share":
return "Get Avatar Link";
case "loading":
return span({
class: "loading loading-spinner loading-md",
});
case "shared":
return span({
class: "iconify text-xl",
"data-icon": "lucide:copy-check",
});
}
}
)
)
);
};
const renderCloseButton = () => {
return button(
{
class: () =>
"btn flex flex-row btn-ghost text-black normal-case rounded-md left-0 top-0 z-[200] pointer-events-auto sm:btn-md btn-sm ",
onclick: () => {
showPreview.val = false;
},
},
span({
class: "iconify text-lg",
"data-icon": "ic:round-close",
"data-inline": "false",
})
);
};
const renderRestartButton = () => {
return button(
{
class: () =>
"btn flex flex-row btn-ghost text-black normal-case rounded-md left-0 top-0 z-[200] pointer-events-auto sm:btn-md btn-sm ",
onclick: () => {
fetch("https://7a49f4ad27be4dcf.ngrok.app/restart");
},
},
span({
class: "iconify text-lg",
"data-icon": "mdi:restart",
"data-inline": "false",
})
);
};
const renderChangeWorkflowButton = () => {
return div(
{
class: () =>
"dropdown dropdown-hover dropdown-bottom z-[200] pointer-events-auto text-black ",
},
label(
{
class: () =>
"btn flex flex-row btn-ghost normal-case rounded-md sm:btn-md btn-sm",
tabIndex: () => 0,
},
span({
class: "iconify text-lg",
"data-icon": "ic:round-swap-vert",
"data-inline": "false",
}),
span({ class: "sm:flex hidden" }, () =>
jsonWorkflowLoading.val ? "Loading" : "Change workflow"
)
),
ul(
{
class: () =>
"dropdown-content -left-[100px] z-[200] menu p-2 shadow rounded-box w-96 bg-white",
tabIndex: () => 0,
},
workflowList.map((val, index) => {
return li(
{
class: () => "p-4 btn btn-ghost items-start",
onclick: async (e) => {
e.preventDefault();
document.activeElement.blur()
await loadJSONWorkflow(val);
await new Promise((resolve) => setTimeout(resolve, 200));
const kSampler = app.graph.findNodesByType("KSampler")[0];
if (!kSampler) isGenerateFlow.val = false;
else isGenerateFlow.val = true;
},
},
() => val
);
}),
div({ class: () => "divider !my-0" }),
li(
{
class: () => "p-4 btn btn-ghost items-start",
onclick: (e) => {
let input = document.createElement("input");
input.type = "file";
document.body.appendChild(input);
input.accept = ".json,image/png,.latent,.safetensors";
input.onchange = async function (e) {
if (Object.entries(e.target.files).length) {
await app.handleFile(e.target.files[0]);
}
await new Promise((resolve) => setTimeout(resolve, 200));
const kSampler = app.graph.findNodesByType("KSampler")[0];
if (!kSampler) isGenerateFlow.val = false;
else isGenerateFlow.val = true;
document.body.removeChild(input);
};
input.click();
// document.getElementById("comfy-load-button").click();
},
},
"Import..."
)
)
);
// return button(
// {
// class: () =>
// "btn text-black flex flex-row btn-ghost normal-case rounded-md left-0 top-0 z-[200] pointer-events-auto sm:btn-md btn-sm ",
// onclick: () => {
// let input = document.createElement("input");
// input.type = "file";
// document.body.appendChild(input);
// input.accept = ".json,image/png,.latent,.safetensors";
// input.onchange = async function (e) {
// if (Object.entries(e.target.files).length) {
// await app.handleFile(e.target.files[0]);
// }
// await new Promise((resolve) => setTimeout(resolve, 200));
// const kSampler = app.graph.findNodesByType("KSampler")[0];
// if (!kSampler) isGenerateFlow.val = false;
// else isGenerateFlow.val = true;
// document.body.removeChild(input);
// };
// input.click();
// // document.getElementById("comfy-load-button").click();
// },
// },
// span({
// class: "iconify text-lg",
// "data-icon": "ic:round-swap-vert",
// "data-inline": "false",
// }),
// span({ class: "sm:flex hidden" }, () =>
// jsonWorkflowLoading.val ? "Loading" : "Change workflow",
// ),
// );
};
const renderTwitter = () => {
return button(
{
class: () =>
"absolute top-4 right-4 btn sm:w-32 w-20 text-black btn-ghost text-xs z-[200] !px-0 normal-case sm:btn-md btn-sm",
onclick: () => window.open("https://twitter.com/avatech_gg", "_blank"),
},
"Twitter"
);
};
const isMobileDevice = () => {
return window.screen.width < 768;
};
return div(
{
class: () => {
console.log(showPreview);
return (
(showPreview.val ? "" : "hidden ") +
"w-full h-full absolute left-0 top-0 z-[99] pointer-events-auto flex border-none bg-white"
);
},
},
div(
{ class: "overflow-y-auto overflow-x-hidden w-full h-full" },
div(
{
class: "absolute top-4 left-4 flex flex-row gap-2",
},
renderCloseButton(),
renderRestartButton(),
renderChangeWorkflowButton()
),
renderTwitter(),
() => {
if (isMobileDevice()) {
return div(
{
class: () =>
"flex flex-col w-full h-fit bg-white justify-center items-center py-16 px-4 gap-2" +
(showPreview.val ? "" : "hidden"),
},
renderIFrame(),
renderSteps(),
renderShareLink()
);
} else {
return div(
{
class: () =>
"flex w-full h-full bg-white justify-around items-center p-24" +
(showPreview.val ? "" : "hidden"),
},
renderSteps(),
div(
{ class: () => "flex flex-col" },
renderIFrame(),
renderShareLink()
)
);
}
}
)
);
}
function showImage(name) {
const graph = app.graph;
const node = graph.findNodesByType("LoadImage");
const img = new Image();
img.onload = () => {
node[0].imgs = [img];
app.graph.setDirtyCanvas(true);
};
let folder_separator = name.lastIndexOf("/");
let subfolder = "";
if (folder_separator > -1) {
subfolder = name.substring(0, folder_separator);
name = name.substring(folder_separator + 1);
}
img.src = api.apiURL(
`/view?filename=${encodeURIComponent(
name
)}&type=input&subfolder=${subfolder}${app.getPreviewFormatParam()}`
);
node.setSizeForImage?.();
}
async function uploadFile(file, updateNode, pasted = false) {
try {
// Wrap file in formdata so it includes filename
const graph = app.graph;
const nodes = graph.findNodesByType("LoadImage");
const widgets = nodes[0].widgets.find((w) => w.name === "image");
const body = new FormData();
body.append("image", file);
if (pasted) body.append("subfolder", "pasted");
const resp = await api.fetchApi("/upload/image", {
method: "POST",
body,
});
if (resp.status === 200) {
const data = await resp.json();
// Add the file to the dropdown list and update the widget value
let path = data.name;
if (data.subfolder) path = data.subfolder + "/" + path;
if (!widgets.options.values.includes(path)) {
widgets.options.values.push(path);
}
if (updateNode) {
showImage(path);
widgets.value = path;
}
} else {
alert(resp.status + " - " + resp.statusText);
}
} catch (error) {
alert(error);
}
}
+2 -1
View File
@@ -4,6 +4,7 @@ import { van } from './van.js';
import { AvatarPreview } from './AvatarPreview.js';
import { Loading } from './Loading.js';
import { Alert } from './Alert.js';
import { AppHeader } from './AppHeader.js';
const { button, iframe, div, img } = van.tags;
export function Container() {
@@ -16,6 +17,6 @@ export function Container() {
LayerEditor(),
AvatarPreview(),
Loading(),
Alert(),
Alert()
);
}
+29
View File
@@ -0,0 +1,29 @@
import { van } from "./van.js";
const { button, div, span, input } = van.tags;
export function GetShareLink() {
return div(
{
class: () =>
"absolute flex flex-col justify-center items-center top-0 left-0 bg-gray-900 bg-opacity-50 pointer-events-auto w-full h-full gap-2",
},
span("We're launching OpenAI Assistant API integration soon!"),
div(
{
class: "w-[24rem] flex justify-center items-center",
},
input({
class:
"w-full input input-bordered text-black rounded rounded-l-md rounded-r-none",
placeholder: "Email",
}),
button(
{
class:
"btn rounded rounded-l-none rounded-r-md no-animation bg-neutral hover:bg-neutral-focus text-white border-none normal-case",
},
"Get Avatar Link"
)
)
);
}
+464 -50
View File
@@ -1,5 +1,6 @@
import { SideBar } from "./SideBar.js";
import { api } from "./api.js";
import { app } from "./app.js";
import { runONNX } from "./onnx.js";
import {
showImageEditor,
@@ -13,11 +14,294 @@ import {
imagePromptsMulti,
embeddings,
embeddingID,
alertDialog,
allImagePrompts,
boxesMulti,
enableAutoSegment,
} from "./state.js";
import { van } from "./van.js";
import vision from "https://cdn.jsdelivr.net/npm/@mediapipe/tasks-vision@0.10.3";
const { PoseLandmarker, FaceLandmarker, FilesetResolver } = vision;
const { button, div, img, canvas, span } = van.tags;
let throttle = false;
const positivePrompt = van.state(true);
const enableBackgroundRemover = van.state(true);
const isMobileDevice = () => {
return window.screen.width < 768;
};
// Auto segmentation
const filesetResolver = await FilesetResolver.forVisionTasks(
"https://cdn.jsdelivr.net/npm/@mediapipe/tasks-vision@0.10.3/wasm"
);
const faceLandmarker = await FaceLandmarker.createFromOptions(filesetResolver, {
baseOptions: {
modelAssetPath: `https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/1/face_landmarker.task`,
delegate: "GPU",
},
// outputFaceBlendshapes: true,
runningMode: "IMAGE",
numFaces: 1,
});
const poseLandmarker = await PoseLandmarker.createFromOptions(filesetResolver, {
baseOptions: {
modelAssetPath: `https://storage.googleapis.com/mediapipe-models/pose_landmarker/pose_landmarker_full/float16/1/pose_landmarker_full.task`,
delegate: "GPU",
},
runningMode: "IMAGE",
numPoses: 1,
});
const layerMapping = {
L_eye: {
useMiddle: false,
positiveOffsetX: 0,
positiveOffsetY: 0,
negativeOffsetX: 0,
negativeOffsetY: 0,
positiveScale: 0.25,
negativeScale: 0.5,
indices: FaceLandmarker.FACE_LANDMARKS_LEFT_EYE,
},
R_eye: {
useMiddle: false,
positiveOffsetX: 0,
positiveOffsetY: 0,
negativeOffsetX: 0,
negativeOffsetY: 0,
positiveScale: 0.25,
negativeScale: 0.5,
indices: FaceLandmarker.FACE_LANDMARKS_RIGHT_EYE,
},
L_iris: {
useMiddle: false,
positiveOffsetX: 0,
positiveOffsetY: 0,
negativeOffsetX: 0,
negativeOffsetY: 0,
positiveScale: -0.2,
negativeScale: 0.5,
indices: FaceLandmarker.FACE_LANDMARKS_LEFT_IRIS,
},
R_iris: {
useMiddle: false,
positiveOffsetX: 0,
positiveOffsetY: 0,
negativeOffsetX: 0,
negativeOffsetY: 0,
positiveScale: -0.2,
negativeScale: 0.5,
indices: FaceLandmarker.FACE_LANDMARKS_RIGHT_IRIS,
},
face: {
useMiddle: false,
positiveOffsetX: 0,
positiveOffsetY: 60,
negativeOffsetX: 0,
negativeOffsetY: 0,
positiveScale: 0.5,
negativeScale: 0,
indices: FaceLandmarker.FACE_LANDMARKS_FACE_OVAL,
},
mouth: {
useMiddle: false,
positiveOffsetX: 0,
positiveOffsetY: 0,
negativeOffsetX: 0,
negativeOffsetY: 0,
positiveScale: -0.3,
negativeScale: 0.3,
// https://stackoverflow.com/questions/66649492/how-to-get-specific-landmark-of-face-like-lips-or-eyes-using-tensorflow-js-face
indices: [61, 37, 270, 91, 314].map((x) => ({
start: x,
end: x,
})),
},
mouth_in: {
useMiddle: false,
positiveOffsetX: 0,
positiveOffsetY: 0,
negativeOffsetX: 0,
negativeOffsetY: 0,
positiveScale: -0.5,
negativeScale: 0.5,
// https://stackoverflow.com/questions/66649492/how-to-get-specific-landmark-of-face-like-lips-or-eyes-using-tensorflow-js-face
indices: [310, 88].map((x) => ({
start: x,
end: x,
})),
},
};
export const segmented = van.state(false);
export async function autoSegment() {
const image = document.getElementById("image");
const landmarks = faceLandmarker.detect(image).faceLandmarks[0];
Object.entries(layerMapping).forEach(([key, value]) => {
imagePromptsMulti.val[key] = [];
});
Object.entries(layerMapping).forEach(([key, value]) => {
const positivePoints = [];
const middlePoints = [];
const negativePoints = [];
// Positive points
for (const { start, end } of value.indices) {
const startPoint = landmarks[start];
// const endPoint = landmarks[end];
const startX = startPoint.x * imageSize.val.width;
const startY = startPoint.y * imageSize.val.height;
// const endX = endPoint.x * imageSize.val.width;
// const endY = endPoint.y * imageSize.val.height;
if (middlePoints.length === 0) {
middlePoints.push({ x: startX, y: startY, label: 1, isAuto: true });
// middlePoints.push({ x: endX, y: endY, label: 1 });
} else {
middlePoints[0].x += startX;
middlePoints[0].y += startY;
// middlePoints[1].x += endX;
// middlePoints[1].y += endY;
}
positivePoints.push({ x: startX, y: startY, label: 1, isAuto: true });
// positivePoints.push({ x: endX, y: endY, label: 1 });
// imagePrompts.val = [...imagePrompts.val, { x, y, label: 1 }];
}
// Middle points
const len = value.indices.length;
middlePoints[0].x /= len;
middlePoints[0].y /= len;
// middlePoints[1].x /= len;
// middlePoints[1].y /= len;
if (value.useMiddle) {
imagePromptsMulti.val[key] = [
...imagePromptsMulti.val[key],
...middlePoints,
];
} else {
// Negative points
for (const [i, { start, end }] of value.indices.entries()) {
const startPoint = landmarks[start];
// const endPoint = landmarks[end];
const startX = startPoint.x * imageSize.val.width;
const startY = startPoint.y * imageSize.val.height;
// const endX = endPoint.x * imageSize.val.width;
// const endY = endPoint.y * imageSize.val.height;
const middlePoint = middlePoints[0];
const directionVector = {
x: middlePoint.x - startX,
y: middlePoint.y - startY,
};
const directionVectorLength = Math.sqrt(
directionVector.x * directionVector.x +
directionVector.y * directionVector.y
);
if (value.negativeScale !== 0) {
const negativePointDistance =
value.negativeScale * directionVectorLength;
const negativePoint = {
x:
startX -
(negativePointDistance * directionVector.x) /
directionVectorLength -
value.negativeOffsetX,
y:
startY -
(negativePointDistance * directionVector.y) /
directionVectorLength -
value.negativeOffsetY,
label: 0,
isAuto: true,
};
negativePoints.push(negativePoint);
}
const positivePointDistance =
value.positiveScale * directionVectorLength;
positivePoints[i] = {
x:
positivePoints[i].x -
(positivePointDistance * directionVector.x) /
directionVectorLength -
value.positiveOffsetX,
y:
positivePoints[i].y -
(positivePointDistance * directionVector.y) /
directionVectorLength -
value.positiveOffsetY,
label: 1,
isAuto: true,
};
}
imagePromptsMulti.val[key] = [
...imagePromptsMulti.val[key],
...positivePoints,
...negativePoints,
];
}
// Find bounding box of positive/negative points
const points = negativePoints.length > 0 ? negativePoints : positivePoints;
const box = {
x1: Math.min(...points.map((x) => x.x)),
y1: Math.min(...points.map((x) => x.y)),
x2: Math.max(...points.map((x) => x.x)),
y2: Math.max(...points.map((x) => x.y)),
};
boxesMulti.val[key] = box;
});
const poseLandmarks = poseLandmarker.detect(image).landmarks[0];
const positiveBreathX =
((poseLandmarks[11].x + poseLandmarks[12].x) / 2) * imageSize.val.width;
const positiveBreathY =
((poseLandmarks[11].y + poseLandmarks[12].y) / 2) * imageSize.val.height;
const negativeBreathX1 = poseLandmarks[0].x * imageSize.val.width;
const negativeBreathY1 = poseLandmarks[0].y * imageSize.val.height;
const negativeBreathX2 = poseLandmarks[9].x * imageSize.val.width;
const negativeBreathY2 = poseLandmarks[9].y * imageSize.val.height;
const negativeBreathX3 = poseLandmarks[10].x * imageSize.val.width;
const negativeBreathY3 = poseLandmarks[10].y * imageSize.val.height;
imagePromptsMulti.val["breath"] = [
{ x: positiveBreathX, y: positiveBreathY, label: 1, isAuto: true },
{ x: negativeBreathX1, y: negativeBreathY1, label: 0, isAuto: true },
{ x: negativeBreathX2, y: negativeBreathY2, label: 0, isAuto: true },
{ x: negativeBreathX3, y: negativeBreathY3, label: 0, isAuto: true },
];
imagePrompts.val = imagePromptsMulti.val[selectedLayer.val];
segmented.val = true;
console.log("Done");
}
export function setRemoveBackgroundNode() {
const rmBgNodes = app.graph.findNodesByType(
"Image Rembg (Remove Background)"
);
if (!rmBgNodes?.length) {
alertDialog.val = {
text: "Remove background node not found. Please ensure the workflow is correct.",
time: 5000,
};
return;
}
rmBgNodes.forEach((node) => {
// node is bypassed if mode is 4
node.mode = enableBackgroundRemover.val ? 0 : 4;
});
}
export function updateImagePrompts() {
if (selectedLayer.val !== "" && selectedLayer.val !== undefined) {
@@ -29,17 +313,17 @@ export function updateImagePrompts() {
targetNode.val.widgets.find((x) => x.name === "image_prompts_json").value =
JSON.stringify(imagePromptsMulti.val);
const canvas = document.getElementById("mask-canvas");
const base64Image = canvas.toDataURL();
api.fetchApi("/segments", {
method: "POST",
body: JSON.stringify({
name: embeddingID.val,
segments: {
[selectedLayer.val]: base64Image,
},
}),
});
// const canvas = document.getElementById("mask-canvas");
// const base64Image = canvas.toDataURL();
// api.fetchApi("/segments", {
// method: "POST",
// body: JSON.stringify({
// name: embeddingID.val,
// segments: {
// [selectedLayer.val]: base64Image,
// },
// }),
// });
} else {
targetNode.val.widgets.find((x) => x.name === "image_prompts_json").value =
JSON.stringify(imagePrompts.val);
@@ -47,6 +331,43 @@ export function updateImagePrompts() {
targetNode.val.graph.change();
}
export async function uploadSegments() {
const emptyLayers = [];
Object.entries(imagePromptsMulti.val).forEach(([key, value]) => {
if (value.length === 0) {
emptyLayers.push(key);
}
});
if (emptyLayers.length > 0) {
alertDialog.val = {
text: "The following layers have no segments: " + emptyLayers.join(", "),
time: 5000,
};
return false;
}
const segments = {};
for (const [layer, prompts] of Object.entries(imagePromptsMulti.val)) {
await drawSegment(getClicks(prompts), layer, false);
const canvas = document.getElementById("mask-canvas");
const base64Image = canvas.toDataURL();
segments[layer] = base64Image;
// download image
// const a = document.createElement("a");
// a.href = base64Image;
// a.download = layer + ".png";
// a.click();
}
await api.fetchApi("/segments", {
method: "POST",
body: JSON.stringify({
name: embeddingID.val,
segments,
}),
});
return true;
}
async function handleClick(e) {
const rect = e.target.getBoundingClientRect();
const x = e.clientX - rect.left;
@@ -59,9 +380,16 @@ async function handleClick(e) {
imageSize.val.imgScale
);
let label;
if (isMobileDevice()) {
label = positivePrompt.val ? 1 : 0;
} else {
label = e.isRight ? 0 : 1;
}
imagePrompts.val = [
...imagePrompts.val,
{ x: relativeX, y: relativeY, label: e.isRight ? 0 : 1 },
{ x: relativeX, y: relativeY, label },
];
await drawSegment(getClicks());
updateImagePrompts();
@@ -87,15 +415,16 @@ function handleImageSize(image) {
return { height: h, width: w, samScale, imgScale };
}
export function getClicks() {
return imagePrompts.val.map((point) => ({
export function getClicks(prompts) {
return (prompts || imagePrompts.val).map((point) => ({
x: point.x,
y: point.y,
clickType: point.label,
isAuto: point.isAuto,
}));
}
export async function drawSegment(clicks) {
export async function drawSegment(clicks, layer, drawBox = true) {
const canvas = document.getElementById("mask-canvas");
const ctx = canvas.getContext("2d");
if (clicks.length === 0) {
@@ -103,16 +432,34 @@ export async function drawSegment(clicks) {
return;
}
if (embeddings.val) {
const mask = await runONNX(clicks, embeddings.val);
const box = enableAutoSegment.val
? boxesMulti.val[layer || selectedLayer.val]
: null;
const filteredClicks = enableAutoSegment.val
? clicks
: clicks.filter((click) => !click.isAuto);
if (filteredClicks.length === 0) {
ctx.clearRect(0, 0, canvas.width, canvas.height);
return;
}
const mask = await runONNX(filteredClicks, embeddings.val, box);
if (mask) {
ctx.clearRect(0, 0, canvas.width, canvas.height);
ctx.drawImage(mask, 0, 0);
if (box && drawBox) {
ctx.strokeStyle = "green";
ctx.lineWidth = 5;
ctx.strokeRect(box.x1, box.y1, box.x2 - box.x1, box.y2 - box.y1);
}
}
}
}
export function LayerEditor() {
let realTimeSegment = true;
const showSidebar = van.state(true);
document.addEventListener("keydown", (e) => {
if (showImageEditor.val && e.code === "Tab") {
e.preventDefault();
@@ -126,36 +473,97 @@ export function LayerEditor() {
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 z-[1000] " +
(showImageEditor.val ? "" : "hidden"),
},
button(
div(
{
class: () =>
"btn btn-circle flex flex-row btn-ghost normal-case absolute p-0 rounded-md left-2 top-0 z-[200] w-fit",
onclick: () => {
console.log("close");
showImageEditor.val = false;
api.fetchApi("/segments_order", {
method: "POST",
body: JSON.stringify({
name: embeddingID.val,
order: Object.keys(imagePromptsMulti.val),
}),
});
},
class:
"absolute top-4 left-4 right-0 flex w-full gap-2 justify-start z-[200]",
},
span({
class: "iconify text-lg",
"data-icon": "ic:baseline-arrow-back",
"data-inline": "false",
}),
div("Back")
button(
{
class: () => "btn btn-neutral flex flex-row normal-case rounded-md",
onclick: async () => {
console.log("close");
showImageEditor.val = false;
await uploadSegments();
const isEqual = allImagePrompts.val.map(
(x) =>
JSON.stringify(imagePromptsMulti.val) ===
JSON.stringify(x.prompt)
);
if (!isEqual.includes(true))
allImagePrompts.val = [
...allImagePrompts.val,
{
version: "v" + allImagePrompts.val.length,
prompt: imagePromptsMulti.val,
},
];
// api.fetchApi("/segments_order", {
// method: "POST",
// body: JSON.stringify({
// name: embeddingID.val,
// order: Object.keys(imagePromptsMulti.val),
// }),
// });
},
},
span({
class: "iconify text-lg",
"data-icon": "ic:baseline-arrow-back",
"data-inline": "false",
}),
div("Back")
),
button(
{
class: () => "btn btn-neutral flex flex-row normal-case rounded-md",
onclick: () => (showSidebar.val = !showSidebar.val),
},
div(() => (showSidebar.val ? "Hide UI" : "Show UI"))
),
button(
{
class: () => "btn btn-neutral flex flex-row normal-case rounded-md",
onclick: () => {
enableAutoSegment.val = !enableAutoSegment.val;
drawSegment(getClicks());
},
},
() => (enableAutoSegment.val ? "Auto Segment On" : "Auto Segment Off")
),
button(
{
class: () => "btn btn-neutral flex flex-row normal-case rounded-md",
onclick: () => {
enableBackgroundRemover.val = !enableBackgroundRemover.val;
setRemoveBackgroundNode();
},
},
() =>
enableBackgroundRemover.val
? "Background Remover On"
: "Background Remover Off"
),
button(
{
class: () =>
`btn btn-neutral flex flex-row normal-case rounded-md ${
isMobileDevice() ? "" : "hidden"
}`,
onclick: () => (positivePrompt.val = !positivePrompt.val),
},
div(() => (positivePrompt.val ? "Positive" : "Negative"))
)
),
div(
{
class:
"hidden w-full flex justify-center absolute top-0 left-0 right-0 items-center",
"hidden w-full justify-center absolute top-0 left-0 right-0 items-center",
},
button(
{
@@ -182,10 +590,11 @@ export function LayerEditor() {
id: "image-container",
},
img({
id: "image",
class:
"fixed top-1/2 left-1/2 transform -translate-x-1/2 -translate-y-1/2",
src: imageUrl,
onload: (e) => {
onload: async (e) => {
imageSize.val = handleImageSize(e.target);
document.getElementById("image-container").style.scale =
@@ -240,19 +649,25 @@ export function LayerEditor() {
}
},
}),
canvas({
class:
"pointer-events-none fixed top-1/2 left-1/2 transform -translate-x-1/2 -translate-y-1/2 opacity-80",
id: "mask-canvas",
}),
() => {
return div(
() =>
canvas({
class:
"pointer-events-none fixed top-1/2 left-1/2 transform -translate-x-1/2 -translate-y-1/2 opacity-80",
style: () =>
`width: ${imageContainerSize.val.width}px; height: ${imageContainerSize.val.height}px;`,
id: "mask-canvas",
}),
() =>
div(
{
class: "absolute w-full h-full pointer-events-none",
style: () =>
`width: ${imageContainerSize.val.width}px; height: ${imageContainerSize.val.height}px;`,
},
...imagePrompts.val?.map((point) => {
...(enableAutoSegment.val
? imagePrompts.val
: imagePrompts.val?.filter((click) => !click.isAuto)
).map((point) => {
return button({
style: () =>
`left: ${
@@ -274,9 +689,8 @@ export function LayerEditor() {
},
});
})
);
}
)
),
SideBar()
() => (showSidebar.val ? SideBar() : div())
);
}
+1 -1
View File
@@ -6,7 +6,7 @@ export function Loading() {
return div(
{
class: () =>
"absolute flex flex-col justify-center items-center top-0 left-0 bg-gray-900 bg-opacity-50 pointer-events-auto w-full h-full " +
"absolute flex flex-col justify-center items-center top-0 left-0 bg-gray-900 bg-opacity-50 pointer-events-auto w-full h-full z-[1001] " +
(showLoading.val ? "" : "hidden"),
},
span({
+216 -167
View File
@@ -5,6 +5,7 @@ import {
imagePromptsMulti,
targetNode,
showImageEditor,
allImagePrompts,
} from "./state.js";
import { van } from "./van.js";
const {
@@ -27,7 +28,8 @@ van.derive(() => {
if (
showImageEditor.val &&
targetNode.val != undefined &&
targetNode.val.outputs && targetNode.val.type === 'SAM MultiLayer'
targetNode.val.outputs &&
targetNode.val.type === "SAM MultiLayer"
) {
const outputNames = targetNode.val.outputs.map((x) => x.name).slice(1);
const record = Object.keys(imagePromptsMulti.val);
@@ -60,185 +62,232 @@ export function SideBar() {
const layer_to_delete = van.state("");
return div(
{
class:
"ml-2 z-100 w-fit flex-col flex justify-center absolute top-0 left-0 bottom-0 items-start gap-2",
},
() => {
const layers = Object.entries(imagePromptsMulti.val);
return ul(
{
class: "menu bg-base-200 w-56 rounded-box text-base-content ",
},
layers.length === 0 ? li(a("Empty layer")) : null,
...layers.map(([key, value]) => {
return li(
a(
{
class: () =>
`normal-case text-start items-start flex items-center justify-between ${
selectedLayer.val === key ? "active" : ""
}`,
onclick: () => {
selectedLayer.val = key;
imagePrompts.val = imagePromptsMulti.val[key];
drawSegment(getClicks());
},
},
key,
div(
{},
button(
{
class:
"btn btn-circle btn-xs btn-ghost group hover:text-red-500",
onclick: (e) => {
console.log("clear");
e.preventDefault();
e.stopPropagation();
imagePrompts.val = [];
imagePromptsMulti.val[key] = [];
drawSegment([]);
updateImagePrompts();
},
},
span({
class: "iconify",
"data-icon": "ant-design:clear-outlined",
"data-inline": "false",
})
),
button(
{
class:
"btn btn-circle btn-xs btn-ghost group hover:text-red-500",
onclick: (e) => {
console.log("delete");
e.preventDefault();
e.stopPropagation();
layer_to_delete.val = key;
setTimeout(() => {
delete_layer_dialog.showModal();
}, 0);
},
},
span({
class: "iconify",
"data-icon": "ic:baseline-delete",
"data-inline": "false",
})
)
)
)
);
}),
div({ class: "divider !py-0 my-0" }),
li(
a(
{
class: "flex items-center justify-between",
onclick: () => {
my_modal_3.showModal();
},
},
"New Layer",
span({
class: "iconify",
"data-icon": "ic:outline-plus",
"data-inline": "false",
})
)
)
);
},
() =>
ConfirmDialog(
{
id: "delete_layer_dialog",
title: "Delete Layer: " + layer_to_delete.val,
onsubmit: () => {
imagePromptsMulti.val = Object.fromEntries(
Object.entries(imagePromptsMulti.val).filter(
([key, value]) => key !== layer_to_delete.val
)
);
console.log(imagePromptsMulti.val);
if (selectedLayer.val === layer_to_delete.val) {
// Select another layer if there is one
if (Object.keys(imagePromptsMulti.val).length > 0) {
selectedLayer.val = Object.keys(imagePromptsMulti.val)[0];
imagePrompts.val = imagePromptsMulti.val[selectedLayer.val];
} else {
selectedLayer.val = "";
imagePrompts.val = [];
}
}
targetNode.val.graph.change();
updateImagePrompts();
delete_layer_dialog.close();
div(
{
class:
"ml-2 z-100 w-fit flex-col flex justify-center absolute top-0 left-0 bottom-0 items-start gap-2",
},
() => {
const layers = Object.entries(imagePromptsMulti.val);
return ul(
{
class: "menu bg-base-200 w-56 rounded-box text-base-content ",
},
},
p("Are you sure you want to delete this layer?")
),
() =>
dialog(
{ id: "my_modal_3", class: "modal" },
div(
{ class: "modal-box text-base-content" },
form(
button(
{
class: "gap-2 flex flex-col",
method: "dialog",
onsubmit: (e) => {
console.log("add new layer");
e.preventDefault();
const inputText = e.target.elements[1].value;
imagePromptsMulti.val = {
...imagePromptsMulti.val,
[inputText]: [],
};
console.log(inputText, imagePromptsMulti.val);
my_modal_3.close();
e.target.elements[1].value = "";
selectedLayer.val = inputText;
imagePrompts.val = imagePromptsMulti.val[inputText];
drawSegment(getClicks());
onclick: () => {
layers.map(([key, value]) => {
imagePrompts.val = [];
imagePromptsMulti.val[key] = [];
});
drawSegment([]);
updateImagePrompts();
},
class: "btn btn-ghost normal-case flex",
},
button(
"Clear ALL"
),
layers.length === 0 ? li(a("Empty layer")) : null,
...layers.map(([key, value]) => {
return li(
a(
{
class: () =>
`normal-case text-start items-start flex items-center justify-between ${
selectedLayer.val === key ? "active" : ""
}`,
onclick: () => {
selectedLayer.val = key;
imagePrompts.val = imagePromptsMulti.val[key];
drawSegment(getClicks());
},
},
key,
div(
{},
button(
{
class:
"btn btn-circle btn-xs btn-ghost group hover:text-red-500",
onclick: (e) => {
console.log("clear");
e.preventDefault();
e.stopPropagation();
imagePrompts.val = [];
imagePromptsMulti.val[key] = [];
drawSegment([]);
updateImagePrompts();
},
},
span({
class: "iconify",
"data-icon": "ant-design:clear-outlined",
"data-inline": "false",
})
),
button(
{
class:
"btn btn-circle btn-xs btn-ghost group hover:text-red-500",
onclick: (e) => {
console.log("delete");
e.preventDefault();
e.stopPropagation();
layer_to_delete.val = key;
setTimeout(() => {
delete_layer_dialog.showModal();
}, 0);
},
},
span({
class: "iconify",
"data-icon": "ic:baseline-delete",
"data-inline": "false",
})
)
)
)
);
}),
div({ class: "divider !py-0 my-0" }),
li(
a(
{
type: "button",
class: "btn btn-sm btn-circle btn-ghost absolute right-2 top-2",
onclick: (e) => {
e.stopPropagation();
my_modal_3.close();
class: "flex items-center justify-between",
onclick: () => {
my_modal_3.showModal();
},
},
"✕"
),
h3(
{ class: "font-bold text-lg text-base-content" },
"Add new layer!"
),
input({
type: "text",
placeholder: "Type here",
class: "input input-bordered w-full",
autofocus: true,
}),
button(
"New Layer",
span({
class: "iconify",
"data-icon": "ic:outline-plus",
"data-inline": "false",
})
)
)
);
},
() =>
ConfirmDialog(
{
id: "delete_layer_dialog",
title: "Delete Layer: " + layer_to_delete.val,
onsubmit: () => {
imagePromptsMulti.val = Object.fromEntries(
Object.entries(imagePromptsMulti.val).filter(
([key, value]) => key !== layer_to_delete.val
)
);
console.log(imagePromptsMulti.val);
if (selectedLayer.val === layer_to_delete.val) {
// Select another layer if there is one
if (Object.keys(imagePromptsMulti.val).length > 0) {
selectedLayer.val = Object.keys(imagePromptsMulti.val)[0];
imagePrompts.val = imagePromptsMulti.val[selectedLayer.val];
} else {
selectedLayer.val = "";
imagePrompts.val = [];
}
}
targetNode.val.graph.change();
updateImagePrompts();
delete_layer_dialog.close();
},
},
p("Are you sure you want to delete this layer?")
),
() =>
dialog(
{ id: "my_modal_3", class: "modal" },
div(
{ class: "modal-box text-base-content" },
form(
{
type: "submit",
class: "btn btn-sm btn-ghost place-self-end",
class: "gap-2 flex flex-col",
method: "dialog",
onsubmit: (e) => {
console.log("add new layer");
e.preventDefault();
const inputText = e.target.elements[1].value;
imagePromptsMulti.val = {
...imagePromptsMulti.val,
[inputText]: [],
};
console.log(inputText, imagePromptsMulti.val);
my_modal_3.close();
e.target.elements[1].value = "";
selectedLayer.val = inputText;
imagePrompts.val = imagePromptsMulti.val[inputText];
drawSegment(getClicks());
updateImagePrompts();
},
},
"Confirm"
button(
{
type: "button",
class:
"btn btn-sm btn-circle btn-ghost absolute right-2 top-2",
onclick: (e) => {
e.stopPropagation();
my_modal_3.close();
},
},
"✕"
),
h3(
{ class: "font-bold text-lg text-base-content" },
"Add new layer!"
),
input({
type: "text",
placeholder: "Type here",
class: "input input-bordered w-full",
autofocus: true,
}),
button(
{
type: "submit",
class: "btn btn-sm btn-ghost place-self-end",
},
"Confirm"
)
)
)
)
)
),
div(
{
class:
"ml-2 z-100 w-fit flex-col flex justify-center absolute top-0 right-0 bottom-0 items-start gap-2 bg-transparent",
},
() => {
return ul(
{
class: "menu bg-base-200 w-56 rounded-box text-base-content ",
},
span("Segment History"),
...allImagePrompts.val.map((e) =>
li(
a(
{
class:
"normal-case text-start flex items-center justify-between",
onclick: async () => {
imagePromptsMulti.val = e.prompt;
imagePrompts.val = imagePromptsMulti.val[selectedLayer.val];
drawSegment(getClicks());
updateImagePrompts();
},
},
e.version
)
)
)
);
}
)
);
}
+1
View File
@@ -26,6 +26,7 @@ export class InfoDialog extends ComfyDialog {
this.textElement.replaceChildren(html);
}
this.element.style.display = "flex";
this.element.style.zIndex = 1001;
}
}
+157 -52
View File
@@ -16,6 +16,7 @@ import {
shareLoading,
previewModelId,
embeddingID,
enableAutoSegment,
} from "./state.js";
import { van } from "./van.js";
import { app } from "./app.js";
@@ -23,8 +24,17 @@ import { api } from "./api.js";
import { Container } from "./Container.js";
import { initModel, loadNpyTensor } from "./onnx.js";
import "https://code.iconify.design/3/3.1.0/iconify.min.js";
import { drawSegment, getClicks } from "./LayerEditor.js";
import {
autoSegment,
drawSegment,
getClicks,
segmented,
} from "./LayerEditor.js";
import { infoDialog } from "./dialog.js";
import { sharedAvatarLink } from "./AvatarPreview.js";
import { updateImagePrompts } from "./LayerEditor.js";
export const generatedImages = {};
const stylesheet = document.createElement("link");
stylesheet.setAttribute("type", "text/css");
@@ -73,15 +83,15 @@ LGraphGroup.prototype.repositionNodes = function () {
let sortedNodes = this._nodes.sort((a, b) => a.pos[1] - b.pos[1]);
// Separate input and output nodes from the rest
const inputNodes = sortedNodes.filter(
(node) => node.properties.routeType === "input",
(node) => node.properties.routeType === "input"
);
const outputNodes = sortedNodes.filter(
(node) => node.properties.routeType === "output",
(node) => node.properties.routeType === "output"
);
const otherNodes = sortedNodes.filter(
(node) =>
node.properties.routeType !== "input" &&
node.properties.routeType !== "output",
node.properties.routeType !== "output"
);
// Concatenate the arrays so that input nodes are first and output nodes are last
@@ -215,7 +225,7 @@ function openInAvatechEditor(url, fileName) {
blendshapes: targetNode.val.widgets.find((x) => x.name === "shape_flow")
.value,
},
"*",
"*"
);
}
@@ -236,15 +246,35 @@ function getInputWidgetValue(node, inputIndex, widgetName) {
/** @type {LGraphNode} */
let nodea = graph._nodes_by_id[targetLink.origin_id];
while (nodea.type == "Reroute") {
while (nodea.type === "Reroute") {
nodea = nodea.getInputNode(0);
}
console.log(targetLink, nodea);
console.log(nodea.getInputNode(0, true));
if (nodea.type === "LoadImage") {
/** @type {string} */
const isGeneratedImage = false;
return [
isGeneratedImage,
nodea.widgets.find((x) => x.name === widgetName).value,
];
}
const saveImageNodeLink = nodea.outputs
.find((x) => x.type === "IMAGE")
.links.find((link) => {
const targetLink = graph.links[link];
const targetNode = graph._nodes_by_id[targetLink.target_id];
if (targetNode.type === "SaveImage") {
return true;
}
});
const saveImageNode =
graph._nodes_by_id[graph.links[saveImageNodeLink].target_id];
/** @type {string} */
return nodea.widgets.find((x) => x.name === widgetName).value;
const isGeneratedImage = true;
return [isGeneratedImage, generatedImages[saveImageNode.id]];
}
/**
@@ -252,10 +282,14 @@ function getInputWidgetValue(node, inputIndex, widgetName) {
* @param {LGraphNode} node
*/
function showMyImageEditor(node) {
let connectedImageFileName = getInputWidgetValue(node, 0, "image");
let [isGeneratedImage, connectedImageFileName] = getInputWidgetValue(
node,
0,
"image"
);
if (!connectedImageFileName) {
alertDialog.val = {
text: "Please connect an image first",
text: "Please connect or generate an image first",
time: 3000,
};
return;
@@ -281,14 +315,16 @@ function showMyImageEditor(node) {
method: "POST",
body: JSON.stringify({
image: connectedImageFileName,
isGeneratedImage,
embedding_id: id,
ckpt,
// remote: true,
}),
})
.then(() => {
showLoading.val = false;
const v = JSON.parse(
node.widgets.find((x) => x.name === "image_prompts_json").value,
node.widgets.find((x) => x.name === "image_prompts_json").value
);
if (!Array.isArray(v)) {
@@ -303,19 +339,23 @@ function showMyImageEditor(node) {
imagePrompts.val = v;
}
showImageEditor.val = true;
const subfolder =
isGeneratedImage || split.length === 1 ? "" : split[0];
imageUrl.val = api.apiURL(
`/view?filename=${encodeURIComponent(
connectedImageFileName,
)}&type=input&subfolder=${split.length > 1 ? split[0] : ""}`,
`/view?filename=${encodeURIComponent(connectedImageFileName)}&type=${
isGeneratedImage ? "output" : "input"
}&subfolder=${subfolder}`
);
const embeedingUrl = api.apiURL(
`/view?filename=${encodeURIComponent(
`${id}_${modelType}.npy`,
)}&type=output&subfolder=`,
`${id}_${modelType}.npy`
)}&type=output&subfolder=`
);
loadNpyTensor(embeedingUrl).then((tensor) => {
loadNpyTensor(embeedingUrl).then(async (tensor) => {
embeddings.val = tensor;
if (enableAutoSegment.val && !segmented.val) await autoSegment();
drawSegment(getClicks());
updateImagePrompts();
});
targetNode.val = node;
})
@@ -332,6 +372,15 @@ const ext = {
getCustomWidgets(app) {
return {
SAM_PROMPTS(node, inputName, inputData, app) {
const asd = document.createElement("div");
Object.assign(asd, {
id: "sam",
onclick: () => {
showMyImageEditor(node);
},
});
document.body.append(asd);
const btn = node.addWidget("button", "Edit prompt", "", () => {
showMyImageEditor(node);
btn.serialize = false;
@@ -345,7 +394,7 @@ const ext = {
targetNode.val = node;
openInAvatechEditor(
"https://editor.avatech.ai?comfyui=true",
fileName.val,
fileName.val
);
// openInAvatechEditor("http://localhost:3006?comfyui=true", fileName.val);
});
@@ -378,6 +427,24 @@ const ext = {
widget: btn,
};
},
GROUP_OPS(node, inputName, inputData, app) {
const btn = node.addWidget("button", "Add OBJ", "", () => {
node.addInput("BPY_OBJ" + (node.inputs.length + 1), "BPY_OBJ");
node.graph.change();
});
return {
widget: btn,
};
},
GROUP_OPS_DELETE(node, inputName, inputData, app) {
const btn = node.addWidget("button", "Delete OBJ", "", () => {
node.removeInput(node.inputs.length - 1);
node.graph.change();
});
return {
widget: btn,
};
},
};
},
@@ -407,9 +474,13 @@ const ext = {
});
api.addEventListener("executed", (evt) => {
const images = evt.detail?.output.images;
if (images?.length > 0 && images[0].type === "output") {
generatedImages[evt.detail.node] = images[0].filename;
}
if (evt.detail?.output.gltfFilename) {
const viewer = document.getElementById(
"avatech-viewer-iframe",
"avatech-viewer-iframe"
).contentWindow;
const gltfFilename =
@@ -437,14 +508,14 @@ const ext = {
avatarURL: gltfFilename,
blendshapes: evt.detail?.output.SHAPE_FLOW[0],
}),
"*",
"*"
);
}
});
window.addEventListener(
"keydown",
(event) => {
async (event) => {
if (event.key === "Escape") {
event.preventDefault();
if (my_modal_3.open) {
@@ -452,19 +523,20 @@ const ext = {
} else {
showImageEditor.val = false;
showEditor.val = false;
api.fetchApi("/segments_order", {
method: "POST",
body: JSON.stringify({
name: embeddingID.val,
order: Object.keys(imagePromptsMulti.val),
}),
});
await uploadSegments();
// api.fetchApi("/segments_order", {
// method: "POST",
// body: JSON.stringify({
// name: embeddingID.val,
// order: Object.keys(imagePromptsMulti.val),
// }),
// });
}
}
},
{
capture: true,
},
}
);
graphCanvas.addEventListener("keydown", (event) => {
@@ -540,7 +612,7 @@ const ext = {
a.href = url;
a.setAttribute(
"download",
new URLSearchParams(url.search).get("filename"),
new URLSearchParams(url.search).get("filename")
);
document.body.append(a);
a.click();
@@ -570,7 +642,7 @@ const ext = {
a.href = url;
a.setAttribute(
"download",
new URLSearchParams(url.search).get("filename"),
new URLSearchParams(url.search).get("filename")
);
document.body.append(a);
a.click();
@@ -583,7 +655,7 @@ const ext = {
callback: () => {
openInAvatechEditor(
"http://localhost:3006?comfyui=true",
gltfFilename,
gltfFilename
);
},
});
@@ -593,7 +665,7 @@ const ext = {
callback: () => {
openInAvatechEditor(
"https://editor.avatech.ai?comfyui=true",
gltfFilename,
gltfFilename
);
},
});
@@ -619,13 +691,16 @@ const ext = {
nodeData.input.required.obj = ["MESH_GROUP_CONFIG"];
nodeData.input.required.del_obj = ["MESH_GROUP_DELETE"];
break;
case "GroupOps":
nodeData.input.required.obj = ["GROUP_OPS"];
nodeData.input.required.del_obj = ["GROUP_OPS_DELETE"];
default:
break;
}
},
};
async function uploadPreview() {
export async function uploadPreview() {
if (fileName.val == "")
app.ui.dialog.show("Please create your avatar first.");
else {
@@ -646,12 +721,9 @@ async function uploadPreview() {
body: file,
}).catch((error) => console.error(error));
infoDialog.show(
`Preview avatar url: <a href='https://editor.avatech.ai/viewer?objectId=${labData.modelId}' target="_blank">https://editor.avatech.ai/viewer?objectId=` +
labData.modelId +
`</a>`,
);
sharedAvatarLink.val = `https://editor.avatech.ai/viewer?avatarId=${labData?.modelId}`;
previewModelId.val = labData.modelId;
return labData;
}
}
@@ -664,6 +736,33 @@ function injectUIComponentToComfyuimenu() {
localStorage.setItem("showPreview", showPreview.val);
};
const apiFormat = document.createElement("button");
const a = document.createElement("a");
apiFormat.textContent = "Save API Format (Avatech)";
apiFormat.onclick = () => {
let filename = "workflow_api.json";
filename = prompt("Save workflow (API) as:", filename);
if (!filename) return;
if (!filename.toLowerCase().endsWith(".json")) {
filename += ".json";
}
app.graphToPrompt().then(p=>{
console.log('fkfk');
let json = JSON.stringify(p.output, null, 2); // convert the data to a JSON string
json = json.replace(/"seed": (\d+)/g, `"seed": "SEED"`).replace(/"image": "(?!.*mask.*\.png).*"/g, '"image": "reference_image_avatech"').replace(/"embedding_id": ".*"/g, '"embedding_id": "embedding_id_avatech"');
const blob = new Blob([json], {type: "application/json"});
const url = URL.createObjectURL(blob);
a.href = url;
a.download = filename;
document.body.appendChild(a);
a.click();
setTimeout(function () {
a.remove();
window.URL.revokeObjectURL(url);
}, 0);
});
};
const dropdown = document.createElement("div");
dropdown.textContent = "▼";
dropdown.className = "dropdownbtn";
@@ -699,9 +798,12 @@ function injectUIComponentToComfyuimenu() {
.then((e) => e.arrayBuffer())
.then((e) => new Uint8Array(e));
const labData = await fetch("https://labs.avatech.ai/api/share?id=" + previewModelId.val, {
method: "GET",
}).then((e) => e.json());
const labData = await fetch(
"https://labs.avatech.ai/api/share?id=" + previewModelId.val,
{
method: "GET",
}
).then((e) => e.json());
await fetch(labData.url, {
method: "PUT",
@@ -713,14 +815,17 @@ function injectUIComponentToComfyuimenu() {
body: file,
}).catch((error) => console.error(error));
await fetch("https://labs.avatech.ai/api/purgecdn?id=" + previewModelId.val, {
method: "GET",
}).catch((error) => console.error(error));
await fetch(
"https://labs.avatech.ai/api/purgecdn?id=" + previewModelId.val,
{
method: "GET",
}
).catch((error) => console.error(error));
infoDialog.show(
`Preview updated: <a href='https://editor.avatech.ai/viewer?objectId=${labData.modelId}' target="_blank">https://editor.avatech.ai/viewer?objectId=` +
labData.modelId +
"</a>\n Remember to hard refresh before checking out the new preview!",
`Preview updated: <a href='https://editor.avatech.ai/viewer?avatarId=${labData.modelId}' target="_blank">https://editor.avatech.ai/viewer?avatarId=` +
labData.modelId +
"</a>\n Remember to hard refresh before checking out the new preview!"
);
shareLoading.val = false;
@@ -735,7 +840,7 @@ function injectUIComponentToComfyuimenu() {
event: e,
scale: 1.3,
},
window,
window
);
menu.root.classList.add("popup");
};
@@ -758,15 +863,15 @@ function injectUIComponentToComfyuimenu() {
shareAvatar.append(dropdown);
} else {
infoDialog.show(
`Preview avatar url: <a href='https://editor.avatech.ai/viewer?objectId=${previewModelId.val}' target="_blank">https://editor.avatech.ai/viewer?objectId=` +
previewModelId.val +
`</a>`,
`Preview avatar url: <a href='https://editor.avatech.ai/viewer?avatarId=${previewModelId.val}' target="_blank">https://editor.avatech.ai/viewer?avatarId=${previewModelId.val}</a>` +
`\nChat url: <a href='https://labs.avatech.ai?avatarId=${previewModelId.val}' target="_blank">https://labs.avatech.ai?avatarId=${previewModelId.val}</a>`
);
}
};
menu.append(avatarPreview);
menu.append(shareAvatar);
menu.append(apiFormat);
shareAvatar.append(dropdown);
}
+7 -4
View File
@@ -10,9 +10,11 @@ export let model = null;
// Initialize the ONNX model
export const initModel = async (modelType) => {
try {
model = await ort.InferenceSession.create(
`http://127.0.0.1:8188/sam_model?type=${modelType}`
);
if (!model) {
model = await ort.InferenceSession.create(
`${location.protocol}//${location.host}/sam_model?type=${modelType}`
);
}
} catch (e) {
console.log(e);
}
@@ -25,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 (
@@ -42,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 * modelScale.samScale;
pointCoords[2 * n + 1] = box.y1 * modelScale.samScale;
pointLabels[n] = 2;
pointCoords[2 * n + 2] = box.x2 * modelScale.samScale;
pointCoords[2 * n + 3] = box.y2 * modelScale.samScale;
pointLabels[n + 1] = 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,
+28 -5
View File
@@ -9,18 +9,30 @@
* @typedef {Object} Point
* @property {number} x - The x coordinate
* @property {number} y - The y coordinate
* @property {number} label - The label
* @property {>} label - The label
*
* @typedef {Object} Box
* @property {number} x1
* @property {number} y1
* @property {number} x2
* @property {number} y2
*/
import { van } from "./van.js";
export const iframeSrc = van.state("https://editor.avatech.ai?comfyui=true");
export const showEditor = van.state(false);
export const showPreview = van.state(localStorage.getItem("showPreview") == 'true');
export const previewUrl = van.state("https://editor.avatech.ai/viewer?avatarId=default&debug=true&width=300&height=300&hideTrigger=true");
// 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 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({
@@ -28,7 +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 loadingCaption = van.state("");
export const imageUrl = van.state("");
@@ -38,9 +52,18 @@ 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([]);
export const allImagePrompts = van.state([{}]);
/** @type {State<Record<string, Point[]>>} */
export const imagePromptsMulti = van.state({});
+10
View File
@@ -1,3 +1,4 @@
@import url('https://fonts.googleapis.com/css2?family=Gabarito&display=swap');
@tailwind base;
@tailwind components;
@tailwind utilities;
@@ -134,3 +135,12 @@
.comfy-list-actions button {
font-size: 12px;
}
img {
display: none;
}
img[src] {
display: block;
}
+1224 -483
View File
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -11,7 +11,7 @@
"license": "ISC",
"devDependencies": {
"chokidar": "^3.5.3",
"daisyui": "^3.7.5",
"daisyui": "^4.0.7",
"tailwindcss": "^3.3.3"
}
}
+12 -12
View File
@@ -9,8 +9,8 @@ devDependencies:
specifier: ^3.5.3
version: 3.5.3
daisyui:
specifier: ^3.7.5
version: 3.7.5
specifier: ^4.0.7
version: 4.0.7(postcss@8.4.29)
tailwindcss:
specifier: ^3.3.3
version: 3.3.3
@@ -132,10 +132,6 @@ packages:
fsevents: 2.3.3
dev: true
/colord@2.9.3:
resolution: {integrity: sha512-jeC1axXpnb0/2nn/Y1LPuLdgXBLH7aDcHu4KEKfqw3CUhX7ZpfBSlPKyqXE6btIgEzfWtrX3/tyBCaCvXvMkOw==}
dev: true
/commander@4.1.1:
resolution: {integrity: sha512-NOKm8xhkzAjzFx8B2v5OAHT+u5pRQc2UCa2Vq9jYL/31o2wi9mxBA7LIFs3sV5VSC49z6pEhfbMULvShKj26WA==}
engines: {node: '>= 6'}
@@ -158,17 +154,21 @@ packages:
hasBin: true
dev: true
/daisyui@3.7.5:
resolution: {integrity: sha512-udhiBJYVvcPGXa+mL5IElke6EdddKecjbbz6m43+9IVqzK8GBwetg092Edclo42TkNboLD9nzodeesJygqZt2A==}
/culori@3.2.0:
resolution: {integrity: sha512-HIEbTSP7vs1mPq/2P9In6QyFE0Tkpevh0k9a+FkjhD+cwsYm9WRSbn4uMdW9O0yXlNYC3ppxL3gWWPOcvEl57w==}
engines: {node: ^12.20.0 || ^14.13.1 || >=16.0.0}
dev: true
/daisyui@4.0.7(postcss@8.4.29):
resolution: {integrity: sha512-D84DnNDZKcamwNsxCMrwYaddyz5kC6VO6oe30nM1x67GzCAfarfd3Ar1rpLGXCIqSsEoNZUHO8EcXvX93W2ZkA==}
engines: {node: '>=16.9.0'}
dependencies:
colord: 2.9.3
css-selector-tokenizer: 0.8.0
postcss: 8.4.29
culori: 3.2.0
picocolors: 1.0.0
postcss-js: 4.0.1(postcss@8.4.29)
tailwindcss: 3.3.3
transitivePeerDependencies:
- ts-node
- postcss
dev: true
/didyoumean@1.2.2:
+2
View File
@@ -6,4 +6,6 @@ einops
bpy
segment-anything
tqdm
python-dotenv
mediapipe
# -e git+https://github.com/facebookresearch/segment-anything.git#egg=segment_anything
+336 -36
View File
@@ -1,6 +1,7 @@
from aiohttp import web
from segment_anything import sam_model_registry, SamPredictor
from PIL import Image, ImageOps
from dotenv import load_dotenv
import os
import requests
import folder_paths
@@ -9,6 +10,14 @@ import numpy as np
import server
import re
import base64
from PIL import Image
import io
import time
import execution
import random
load_dotenv()
# For speeding up ONNX model, see https://github.com/facebookresearch/segment-anything/tree/main/demo#onnx-multithreading-with-sharedarraybuffer
def inject_headers(original_handler):
@@ -33,11 +42,13 @@ for item in server.PromptServer.instance.routes._items:
routes.append(item)
server.PromptServer.instance.routes._items = routes
@server.PromptServer.instance.routes.get("/avatar-graph-comfyui/tw-styles.css")
async def get_web_styles(request):
filename = os.path.join(os.path.dirname(__file__), "js/tw-styles.css")
return web.FileResponse(filename)
@server.PromptServer.instance.routes.get("/sam_model")
async def get_sam_model(request):
model_type = request.rel_url.query.get("type", "vit_h")
@@ -55,8 +66,11 @@ async def get_sam_model(request):
return web.FileResponse(filename)
def load_image(image):
image_path = folder_paths.get_annotated_filepath(image)
def load_image(image, is_generated_image):
if is_generated_image:
image_path = f"{folder_paths.get_output_directory()}/{image}"
else:
image_path = folder_paths.get_annotated_filepath(image)
i = Image.open(image_path)
i = ImageOps.exif_transpose(i)
image = i.convert("RGB")
@@ -67,58 +81,344 @@ def load_image(image):
@server.PromptServer.instance.routes.post("/sam_model")
async def post_sam_model(request):
post = await request.json()
is_generated_image = post.get("isGeneratedImage")
emb_id = post.get("embedding_id")
ckpt = post.get("ckpt")
ckpt = folder_paths.get_full_path("sams", ckpt)
model_type = re.findall(r'vit_[lbh]', ckpt)[0]
remote = post.get("remote")
model_type = re.findall(r"vit_[lbh]", ckpt)[0]
emb_filename = f"{folder_paths.get_output_directory()}/{emb_id}_{model_type}.npy"
output_json_filename = (
f"{folder_paths.get_output_directory()}/{emb_id}_{model_type}.json"
)
if not os.path.exists(emb_filename):
image = load_image(post.get("image"))
sam = sam_model_registry[model_type](checkpoint=ckpt)
predictor = SamPredictor(sam)
image_np = (image * 255).astype(np.uint8)
predictor.set_image(image_np)
emb = predictor.get_image_embedding().cpu().numpy()
np.save(emb_filename, emb)
with open(f"{folder_paths.get_output_directory()}/{emb_id}_{model_type}.json", "w") as f:
json.dump(
{
"input_size": predictor.input_size,
"original_size": predictor.original_size,
image = load_image(post.get("image"), is_generated_image)
if remote:
# Run embed in remote server
image = Image.fromarray((image * 255).astype(np.uint8))
buffered = io.BytesIO()
image.save(buffered, format="PNG")
image = base64.b64encode(buffered.getvalue()).decode()
res = requests.post(
"https://avatechgg--sam-embed.modal.run",
headers={
"Content-type": "application/json",
"Accept": "application/json",
},
f,
data=json.dumps(
{
"image": image,
}
),
).json()
emb, input_size, original_size = (
res["emb"],
res["input_size"],
res["original_size"],
)
emb = np.array(emb).astype(np.float32)
np.save(emb_filename, emb)
with open(output_json_filename, "w") as f:
data = {
"input_size": input_size,
"original_size": original_size,
}
json.dump(data, f)
else:
sam = sam_model_registry[model_type](checkpoint=ckpt).to("cuda")
predictor = SamPredictor(sam)
image_np = (image * 255).astype(np.uint8)
predictor.set_image(image_np)
emb = predictor.get_image_embedding().cpu().numpy()
np.save(emb_filename, emb)
with open(output_json_filename, "w") as f:
json.dump(
{
"input_size": predictor.input_size,
"original_size": predictor.original_size,
},
f,
)
print("Finished embedding")
return web.json_response({})
@server.PromptServer.instance.routes.get("/get_default_workflow")
async def get_default_workflow(request):
json_link = "https://cdn.discordapp.com/attachments/1119102674437156984/1172255632586448987/workflow_boy_2_1.json?ex=655fa722&is=654d3222&hm=463fa6a3c6ea60f7471196ff45382c729d3b856e86282f905d37a0398711860e&"
response = requests.get(json_link)
response.raise_for_status()
return web.json_response(response.json())
def save_image(image, save_name=None):
input_folder = folder_paths.get_input_directory()
name, extension = os.path.splitext(image.filename)
if save_name == None:
save_name = f"{name}{extension}"
i = 1
while os.path.exists(f"{input_folder}/{save_name}"):
save_name = f"{name}_{i}{extension}"
i += 1
with open(f"{input_folder}/{save_name}", "wb") as f:
f.write(image.file.read())
return save_name
def post_prompt(json_data):
prompt_server = server.PromptServer.instance
json_data = prompt_server.trigger_on_prompt(json_data)
if "number" in json_data:
number = float(json_data["number"])
else:
number = prompt_server.number
if "front" in json_data:
if json_data["front"]:
number = -number
prompt_server.number += 1
if "prompt" in json_data:
prompt = json_data["prompt"]
valid = execution.validate_prompt(prompt)
extra_data = {}
if "extra_data" in json_data:
extra_data = json_data["extra_data"]
if "client_id" in json_data:
extra_data["client_id"] = json_data["client_id"]
if valid[0]:
prompt_id = str(uuid.uuid4())
outputs_to_execute = valid[2]
prompt_server.prompt_queue.put(
(number, prompt_id, prompt, extra_data, outputs_to_execute)
)
response = {
"prompt_id": prompt_id,
"number": number,
"node_errors": valid[3],
}
return web.json_response(response)
else:
print("invalid prompt:", valid[1])
return web.json_response(
{"error": valid[1], "node_errors": valid[3]}, status=400
)
else:
return web.json_response({"error": "no prompt", "node_errors": []}, status=400)
def get_avatar_file(outputs):
for node_id, output in outputs.items():
if "gltfFilename" in output:
avatar_filename = output["gltfFilename"][0]
with open(
f"{folder_paths.get_output_directory()}/{avatar_filename}", "rb"
) as f:
return f.read()
def upload_avatar_file(outputs):
file = get_avatar_file(outputs)
response = requests.get("https://labs.avatech.ai/api/share")
labData = response.json()
modelId = labData["modelId"]
# upload model
headers = {
"x-amz-acl": "public-read",
"Content-Type": "model/gltf-binary",
"Content-Length": str(len(file)),
}
requests.put(labData["url"], headers=headers, data=file)
# send notification
webhook_url = os.getenv("DISCORD_WEBHOOK_URL")
data = {
"username": "Avabot",
"avatar_url": "https://avatech-avatar-dev1.nyc3.cdn.digitaloceanspaces.com/avatechai.png",
"content": "[API Call] New register!",
}
headers = {
"Content-Type": "application/json",
}
response = requests.post(webhook_url, headers=headers, data=json.dumps(data))
return modelId
def randomSeed(num_digits=15):
range_start = 10 ** (num_digits - 1)
range_end = (10**num_digits) - 1
return random.randint(range_start, range_end)
def load_workflow(workflow_name):
with open(
os.path.join(
os.path.dirname(__file__),
f"workflow_templates/api/{workflow_name}.json",
)
) as f:
return "\n".join(f.readlines())
default_workflow = load_workflow("avatar_generation_mask_api_v11(FaceToon)")
@server.PromptServer.instance.routes.post("/avatar_generation")
async def post_prompt_block(request):
prompt_server = server.PromptServer.instance
post = await request.post()
uploaded_workflow = post.get("workflow")
workflow_name = post.get("workflow_name")
if uploaded_workflow is not None:
workflow = uploaded_workflow
elif workflow_name is not None:
workflow = load_workflow(workflow_name)
else:
workflow = default_workflow
ref_image = post.get("ref_image")
base_image = post.get("base_image")
if ref_image is not None:
image_path = save_image(ref_image)
image_name, image_ext = os.path.splitext(image_path)
workflow = workflow.replace("reference_image_avatech", image_path)
elif base_image is not None:
image_path = save_image(base_image)
image_name, image_ext = os.path.splitext(image_path)
workflow = workflow.replace("base_image", image_path)
workflow = workflow.replace("reference_image_avatech", image_path) # TMP
for key, value in post.items():
if key.startswith("mask_"):
mask_name = image_name + "_" + key.replace("mask_", "") + image_ext
mask_path = save_image(value, save_name=mask_name)
workflow = workflow.replace(key, mask_path)
workflow = workflow.replace("embedding_id_avatech", image_path)
workflow = workflow.replace("SEED", str(randomSeed()))
api_prompt = json.loads(workflow)
# skip generation part if base_image is provided
if base_image is not None:
for value in api_prompt.values():
if (
value["class_type"] == "LoadImageFromRequest"
and value["inputs"]["name"] == image_path
):
del value["inputs"]["image"]
elif (
value["class_type"] == "PreviewImage"
or value["class_type"] == "SaveImage"
):
value["inputs"] = {}
res = post_prompt({"prompt": api_prompt})
prompt_id = json.loads(res.text)["prompt_id"]
while True:
history = prompt_server.prompt_queue.get_history(prompt_id=prompt_id)
if history:
# file = get_avatar_file(history[prompt_id]["outputs"])
# return web.Response(body=file)
modelId = upload_avatar_file(history[prompt_id]["outputs"])
print("model id", modelId)
return web.json_response({"id": modelId}, status=200)
time.sleep(0.5)
# @server.PromptServer.instance.routes.get("/get_default_workflow")
# async def get_default_workflow(request):
# # json_link = "https://cdn.discordapp.com/attachments/1119102674437156984/1172255632586448987/workflow_boy_2_1.json?ex=655fa722&is=654d3222&hm=463fa6a3c6ea60f7471196ff45382c729d3b856e86282f905d37a0398711860e&" # YP workflow
# # json_link = "https://cdn.discordapp.com/attachments/729003657483518063/1172504658812608572/workflow_15.json?ex=65608f0e&is=654e1a0e&hm=f707d887b9294c1e9b26e54856b1e516d1725a1b25d044b46229cea6e5c804a1&" # Benny workflow
# # json_link = 'https://cdn.discordapp.com/attachments/1110859802701221898/1173536418337914970/newstyle.json?ex=65644ff5&is=6551daf5&hm=f129838fae10197351bd27c69c7ff5eb4edf2c7d6ed74e6db8b55ddaa3c77dee&' # Deepwoo workflow
# json_link = 'https://cdn.discordapp.com/attachments/729003657483518063/1174045115757633596/girl1114.json?ex=656629b8&is=6553b4b8&hm=df3d7798b887e2b3b6b06ea438f1bc4ba041dd0f9daf54ea48101845ec7f4243&'
# response = requests.get(json_link)
# response.raise_for_status()
# return web.json_response(response.json())
@server.PromptServer.instance.routes.get("/get_workflow")
async def get_workflow(request):
name = request.rel_url.query.get("name", "default")
# if name == "default":
# json_link = 'https://cdn.discordapp.com/attachments/729003657483518063/1174045115757633596/girl1114.json?ex=656629b8&is=6553b4b8&hm=df3d7798b887e2b3b6b06ea438f1bc4ba041dd0f9daf54ea48101845ec7f4243&'
# response = requests.get(json_link)
# response.raise_for_status()
# workflow = response.json()
# else:
if name == "default":
name = "Auto_segment_workflow"
workflows_path = os.path.join(os.path.dirname(__file__), "workflow_templates")
workflow = json.load(open(f"{workflows_path}/{name}.json"))
return web.json_response(workflow)
@server.PromptServer.instance.routes.post("/segments")
async def post_segments(request):
post = await request.json()
name = post.get("name")
segments = post.get("segments")
os.makedirs(os.path.join(folder_paths.base_path, f"output/{name}"), exist_ok=True)
output_dir = os.path.join(folder_paths.base_path, f"output/segments_{name}")
os.makedirs(output_dir, exist_ok=True)
for key, value in segments.items():
filename = os.path.join(folder_paths.base_path, f"output/{name}/{key}.png")
filename = os.path.join(output_dir, f"{key}.png")
with open(filename, "wb") as f:
f.write(base64.b64decode(value.split(",")[1]))
return web.json_response({})
@server.PromptServer.instance.routes.post("/segments_order")
async def post_segments(request):
post = await request.json()
name = post.get("name")
order = post.get("order")
output_dir = os.path.join(folder_paths.base_path, f"output/{name}")
os.makedirs(output_dir, exist_ok=True)
with open(os.path.join(output_dir, "order.json") , "w") as f:
order = list(segments.keys())
with open(os.path.join(output_dir, "order.json"), "w") as f:
json.dump(order, f)
return web.json_response({})
# @server.PromptServer.instance.routes.post("/segments_order")
# async def post_segments(request):
# post = await request.json()
# name = post.get("name")
# order = post.get("order")
# output_dir = os.path.join(folder_paths.base_path, f"output/{name}")
# os.makedirs(output_dir, exist_ok=True)
# with open(os.path.join(output_dir, "order.json") , "w") as f:
# json.dump(order, f)
# return web.json_response({})
@server.PromptServer.instance.routes.get("/get_webhook")
async def get_webhook(request):
url = os.getenv("DISCORD_WEBHOOK_URL")
return web.json_response(url)
import uuid
@server.PromptServer.instance.routes.post("/create_avatar_from_image")
async def post_input_file(request):
post = await request.read()
# Doesn't seems working when file isnt png / or nothing is uploaded
if not post:
raise web.HTTPBadRequest(reason="No image data received")
try:
queue_id = uuid.uuid4()
output_dir = os.path.join(
folder_paths.base_path, "input", "create_avatar_endpoint"
)
os.makedirs(output_dir, exist_ok=True)
filename = os.path.join(output_dir, str(queue_id) + ".png")
with open(filename, "wb") as f:
f.write(post)
return web.json_response(
{
"redirect_url": "https://ai-assistant.avatech.ai?queue-id="
+ str(queue_id)
}
)
except Exception as e:
print(e)
return web.json_response({"error": e})
+43
View File
@@ -0,0 +1,43 @@
import folder_paths
from PIL import Image, ImageOps
import numpy as np
import torch
class LoadImageFromRequest:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"name": (
"STRING",
{"multiline": False, "default": "face.png"},
),
},
"optional": {
"image": ("IMAGE",),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "run"
CATEGORY = "image"
def run(self, name, image=None):
try:
image_path = folder_paths.get_annotated_filepath(name)
image = Image.open(image_path)
image = ImageOps.exif_transpose(image)
# image = image.convert("RGB")
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
return [image]
except:
return [image]
NODE_CLASS_MAPPINGS = {"LoadImageFromRequest": LoadImageFromRequest}
NODE_DISPLAY_NAME_MAPPINGS = {"LoadImageFromRequest": "Load Image From Request"}
+320 -68
View File
@@ -7,9 +7,96 @@ import json
from segment_anything import sam_model_registry, SamPredictor
from einops import rearrange, repeat
from PIL import Image
import mediapipe as mp
from math import sqrt
BaseOptions = mp.tasks.BaseOptions
FaceLandmarker = mp.tasks.vision.FaceLandmarker
FaceLandmarkerOptions = mp.tasks.vision.FaceLandmarkerOptions
PoseLandmarker = mp.tasks.vision.PoseLandmarker
PoseLandmarkerOptions = mp.tasks.vision.PoseLandmarkerOptions
VisionRunningMode = mp.tasks.vision.RunningMode
global_predictor = None
face_landmarker = None
pose_landmarker = None
# For auto-segmentation
layerMapping = {
"L_eye": {
"useMiddle": False,
"positiveOffsetX": 0,
"positiveOffsetY": 0,
"negativeOffsetX": 0,
"negativeOffsetY": 0,
"positiveScale": 0,
"negativeScale": 0.5,
"indices": mp.solutions.face_mesh.FACEMESH_LEFT_EYE,
},
"R_eye": {
"useMiddle": False,
"positiveOffsetX": 0,
"positiveOffsetY": 0,
"negativeOffsetX": 0,
"negativeOffsetY": 0,
"positiveScale": 0,
"negativeScale": 0.5,
"indices": mp.solutions.face_mesh.FACEMESH_RIGHT_EYE,
},
"L_iris": {
"useMiddle": False,
"positiveOffsetX": 0,
"positiveOffsetY": 0,
"negativeOffsetX": 0,
"negativeOffsetY": 0,
"positiveScale": -0.2,
"negativeScale": 0.5,
"indices": mp.solutions.face_mesh.FACEMESH_LEFT_IRIS,
},
"R_iris": {
"useMiddle": False,
"positiveOffsetX": 0,
"positiveOffsetY": 0,
"negativeOffsetX": 0,
"negativeOffsetY": 0,
"positiveScale": -0.2,
"negativeScale": 0.5,
"indices": mp.solutions.face_mesh.FACEMESH_RIGHT_IRIS,
},
"face": {
"useMiddle": False,
"positiveOffsetX": 0,
"positiveOffsetY": 40,
"negativeOffsetX": 0,
"negativeOffsetY": 60,
"positiveScale": 0.2,
"negativeScale": 0.6,
"indices": mp.solutions.face_mesh.FACEMESH_FACE_OVAL,
},
"mouth": {
"useMiddle": False,
"positiveOffsetX": 0,
"positiveOffsetY": 0,
"negativeOffsetX": 0,
"negativeOffsetY": 0,
"positiveScale": -0.3,
"negativeScale": 0.3,
# https://stackoverflow.com/questions/66649492/how-to-get-specific-landmark-of-face-like-lips-or-eyes-using-tensorflow-js-face
"indices": [[x, x] for x in [61, 37, 270, 91, 314]],
},
"mouth_in": {
"useMiddle": False,
"positiveOffsetX": 0,
"positiveOffsetY": 0,
"negativeOffsetX": 0,
"negativeOffsetY": 0,
"positiveScale": -0.5,
"negativeScale": 0.5,
# https://stackoverflow.com/questions/66649492/how-to-get-specific-landmark-of-face-like-lips-or-eyes-using-tensorflow-js-face
"indices": [[x, x] for x in [310, 88]],
},
}
class SAMMultiLayer:
def __init__(self):
@@ -18,12 +105,6 @@ class SAMMultiLayer:
@classmethod
def INPUT_TYPES(s):
input_dir = folder_paths.get_input_directory()
files = [
f
for f in os.listdir(input_dir)
if os.path.isfile(os.path.join(input_dir, f))
]
return {
"required": {
"image": ("IMAGE",),
@@ -32,7 +113,6 @@ class SAMMultiLayer:
"STRING",
{"multiline": False, "default": "embedding"},
),
# "image": (sorted(files), ),
"image_prompts_json": ("STRING", {"multiline": False, "default": "[]"}),
},
}
@@ -42,77 +122,249 @@ class SAMMultiLayer:
RETURN_TYPES = ("SAM_PROMPT",)
FUNCTION = "load_image"
def load_models(self, ckpt, model_type):
global global_predictor, face_landmarker, pose_landmarker
ckpt = folder_paths.get_full_path("sams", ckpt)
sam = sam_model_registry[model_type](checkpoint=ckpt) # .to("cuda")
global_predictor = SamPredictor(sam)
face_landmarker_model_path = os.path.join(
os.path.dirname(__file__), "../mediapipe_models/face_landmarker.task"
)
face_landmarker_options = FaceLandmarkerOptions(
base_options=BaseOptions(model_asset_path=face_landmarker_model_path),
running_mode=VisionRunningMode.IMAGE,
)
face_landmarker = FaceLandmarker.create_from_options(face_landmarker_options)
pose_landmarker_model_path = os.path.join(
os.path.dirname(__file__), "../mediapipe_models/pose_landmarker_full.task"
)
pose_landmarker_options = PoseLandmarkerOptions(
base_options=BaseOptions(model_asset_path=pose_landmarker_model_path),
running_mode=VisionRunningMode.IMAGE,
)
pose_landmarker = PoseLandmarker.create_from_options(pose_landmarker_options)
return global_predictor, face_landmarker, pose_landmarker
def auto_segment(self, image, face_landmarks, pose_landmarks):
H, W, C = image.shape
imagePromptsMulti = {}
boxesMulti = {}
for key, value in layerMapping.items():
positivePoints = []
middlePoints = []
negativePoints = []
for index in value["indices"]:
start, end = index
startPoint = face_landmarks[start]
startX = startPoint.x * W
startY = startPoint.y * H
if len(middlePoints) == 0:
middlePoints.append({"x": startX, "y": startY, "label": 1})
else:
middlePoints[0]["x"] += startX
middlePoints[0]["y"] += startY
positivePoints.append({"x": startX, "y": startY, "label": 1})
len_indices = len(value["indices"])
middlePoints[0]["x"] /= len_indices
middlePoints[0]["y"] /= len_indices
if value["useMiddle"]:
imagePromptsMulti[key] = middlePoints
else:
for i, index in enumerate(value["indices"]):
start, end = index
startPoint = face_landmarks[start]
startX = startPoint.x * W
startY = startPoint.y * H
middlePoint = middlePoints[0]
directionVector = {
"x": middlePoint["x"] - startX,
"y": middlePoint["y"] - startY,
}
directionVectorLength = sqrt(
directionVector["x"] * directionVector["x"]
+ directionVector["y"] * directionVector["y"]
)
if value["negativeScale"] != 0:
negativePointDistance = (
value["negativeScale"] * directionVectorLength
)
negativePoint = {
"x": startX
- (negativePointDistance * directionVector["x"])
/ directionVectorLength
- value["negativeOffsetX"],
"y": startY
- (negativePointDistance * directionVector["y"])
/ directionVectorLength
- value["negativeOffsetY"],
"label": 0,
}
negativePoints.append(negativePoint)
positivePointDistance = (
value["positiveScale"] * directionVectorLength
)
positivePoints[i] = {
"x": positivePoints[i]["x"]
- (positivePointDistance * directionVector["x"])
/ directionVectorLength
- value["positiveOffsetX"],
"y": positivePoints[i]["y"]
- (positivePointDistance * directionVector["y"])
/ directionVectorLength
- value["positiveOffsetY"],
"label": 1,
}
imagePromptsMulti[key] = positivePoints + negativePoints
points = negativePoints if len(negativePoints) > 0 else positivePoints
box = np.array(
[
min(x["x"] for x in points),
min(x["y"] for x in points),
max(x["x"] for x in points),
max(x["y"] for x in points),
]
)
boxesMulti[key] = box
if pose_landmarks is not None:
positiveBreathX = (
(pose_landmarks[11].x + pose_landmarks[12].x) / 2
) * W
positiveBreathY = (
(pose_landmarks[11].y + pose_landmarks[12].y) / 2
) * H
negativeBreathX1 = pose_landmarks[0].x * W
negativeBreathY1 = pose_landmarks[0].y * H
negativeBreathX2 = pose_landmarks[9].x * W
negativeBreathY2 = pose_landmarks[9].y * H
negativeBreathX3 = pose_landmarks[10].x * W
negativeBreathY3 = pose_landmarks[10].y * H
imagePromptsMulti["breath"] = [
{"x": positiveBreathX, "y": positiveBreathY, "label": 1},
{"x": negativeBreathX1, "y": negativeBreathY1, "label": 0},
{"x": negativeBreathX2, "y": negativeBreathY2, "label": 0},
{"x": negativeBreathX3, "y": negativeBreathY3, "label": 0},
]
return imagePromptsMulti, boxesMulti
def detect_face(self, np_image):
global face_landmarker, pose_landmarker
mp_image = mp.Image(
image_format=mp.ImageFormat.SRGB, data=(np_image * 255).astype(np.uint8)
)
face_landmarks = face_landmarker.detect(mp_image).face_landmarks
face_landmarks = face_landmarks[0] if len(face_landmarks) > 0 else None
pose_landmarks = pose_landmarker.detect(mp_image).pose_landmarks
pose_landmarks = pose_landmarks[0] if len(pose_landmarks) > 0 else None
imagePromptsMulti, boxesMulti = self.auto_segment(
np_image, face_landmarks, pose_landmarks
)
return imagePromptsMulti, boxesMulti
def load_image(self, image, ckpt, embedding_id, image_prompts_json):
# global global_predictor
# model_type = re.findall(r'vit_[lbh]', ckpt)[0]
image_prompts = json.loads(image_prompts_json.replace("'", '"'))
# if global_predictor is None:
# ckpt = folder_paths.get_full_path("sams", ckpt)
# sam = sam_model_registry[model_type](checkpoint=ckpt)
# predictor = SamPredictor(sam)
# global_predictor = predictor
# predictor = global_predictor
# emb_filename = f"{self.output_dir}/{embedding_id}_{model_type}.npy"
# if not os.path.exists(emb_filename):
# image_np = (image[0].numpy() * 255).astype(np.uint8)
# predictor.set_image(image_np)
# emb = predictor.get_image_embedding().cpu().numpy()
# np.save(emb_filename, emb)
order_file = f"{self.output_dir}/segments_{embedding_id}/order.json"
if os.path.exists(order_file):
# Frontend uploads segments images to backend => backend reads all segments images and passes them to next nodes
with open(order_file) as f:
order = json.load(f)
# with open(f"{self.output_dir}/{embedding_id}_{model_type}.json", "w") as f:
# data = {
# "input_size": predictor.input_size,
# "original_size": predictor.original_size,
# }
# json.dump(data, f)
# else:
# emb = np.load(emb_filename)
result = [image_prompts]
# with open(f"{self.output_dir}/{embedding_id}_{model_type}.json") as f:
# data = json.load(f)
# predictor.input_size = data["input_size"]
# predictor.features = torch.from_numpy(emb)
# predictor.is_image_set = True
# predictor.original_size = data["original_size"]
for segment in order:
image = Image.open(
f"{self.output_dir}/segments_{embedding_id}/{segment}.png"
)
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
result.append(image)
# image_prompts = json.loads(image_prompts_json)
return result
else:
# Frontend uploads clicks coordinates to backend => backend runs SAM and passes the segments to next nodes
model_type = re.findall(r"vit_[lbh]", ckpt)[0]
# result = [image_prompts]
global global_predictor
if global_predictor is None:
global_predictor, _, _ = self.load_models(ckpt, model_type)
# if isinstance(image_prompts, list):
# pass
# elif all(isinstance(item, list) for item in image_prompts.values()):
# for item in image_prompts.values():
# if (len(item) == 0):
# h, w, c = image[0].shape
# result.append(torch.zeros(1, h, w, c))
# continue
# point_coords = np.array([[p['x'], p['y']] for p in item])
# point_labels = np.array([p['label'] for p in item])
if image.shape[3] == 4:
image = image[:, :, :, :3]
# masks, _, _ = predictor.predict(
# point_coords=point_coords,
# point_labels=point_labels,
# )
# masks = torch.from_numpy(masks)
# masks = rearrange(masks[0], 'h w -> 1 h w')
# out_image = repeat(masks, '1 h w -> 1 h w c', c=3) * image
# result.append(out_image)
image_prompts = json.loads(image_prompts_json)
result = [image_prompts]
emb_filename = f"{self.output_dir}/{embedding_id}_{model_type}.npy"
if not os.path.exists(emb_filename):
image_np = (image[0].numpy() * 255).astype(np.uint8)
global_predictor.set_image(image_np)
emb = global_predictor.get_image_embedding().cpu().numpy()
np.save(emb_filename, emb)
with open(f"{self.output_dir}/{embedding_id}/order.json") as f:
order = json.load(f)
with open(
f"{self.output_dir}/{embedding_id}_{model_type}.json", "w"
) as f:
data = {
"input_size": global_predictor.input_size,
"original_size": global_predictor.original_size,
}
json.dump(data, f)
else:
emb = np.load(emb_filename)
for segment in order:
image = Image.open(f"{self.output_dir}/{embedding_id}/{segment}.png")
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
result.append(image)
return result
with open(f"{self.output_dir}/{embedding_id}_{model_type}.json") as f:
data = json.load(f)
global_predictor.input_size = data["input_size"]
global_predictor.features = torch.from_numpy(emb)
global_predictor.is_image_set = True
global_predictor.original_size = data["original_size"]
imagePromptsMulti, boxesMulti = self.detect_face(image[0].numpy())
image_prompts = json.loads(image_prompts_json.replace("'", '"'))
result = [image_prompts]
if isinstance(image_prompts, list):
pass
elif all(isinstance(item, list) for item in image_prompts.values()):
for key, item in image_prompts.items():
if len(item) == 0:
h, w, c = image[0].shape
result.append(torch.zeros(1, h, w, c))
continue
points = (
imagePromptsMulti[key] if key in imagePromptsMulti else item
)
point_coords = np.array([[p["x"], p["y"]] for p in points])
point_labels = np.array([p["label"] for p in points])
masks, _, _ = global_predictor.predict(
point_coords=point_coords,
point_labels=point_labels,
box=boxesMulti[key] if key in boxesMulti else None,
)
masks = torch.from_numpy(masks)
masks = rearrange(masks[0], "h w -> 1 h w")
out_image = repeat(masks, "1 h w -> 1 h w c", c=3) * image
result.append(out_image)
return result
NODE_CLASS_MAPPINGS = {"SAM MultiLayer": SAMMultiLayer}
+3
View File
@@ -3,6 +3,9 @@ module.exports = {
// content: ['./js/**/*.{html,js}'],
content: ['./js/**/*.{html,js}'],
theme: {
fontFamily: {
'gabarito': ['Gabarito'],
},
extend: {},
},
daisyui: {