17 Commits
Author SHA1 Message Date
EdwinWong 42848ac555 fix: onnx runtime version 2024-02-06 16:54:30 +08:00
Radionic 2615a9ec51 feat: return black image when no contours found 2024-01-24 17:10:24 +08:00
Radionic 5406bb2780 fix: sam input image 2024-01-24 16:47:59 +08:00
Radionic 542030e349 feat: separate face and pose detection 2024-01-24 16:47:46 +08:00
Radionic 3d37cd2b6d refactor: sam and mediapipe 2024-01-24 16:35:32 +08:00
EdwinWong 82e3f4dbec fix: preview display 2024-01-24 11:53:03 +08:00
EdwinWong 10d93e0823 fix: notification 2024-01-19 19:50:27 +08:00
EdwinWong 2ded553837 fix: add output type in main output 2024-01-19 18:06:14 +08:00
EdwinWong 00a7d5a26c fix: nanoid 2024-01-19 18:05:17 +08:00
Radionic 980408b363 fix: route 2024-01-16 11:35:38 +08:00
Radionic eaaab6e346 feat: data generation 2024-01-11 15:26:16 +08:00
Radionic eccc4bca31 feat: render blender image 2024-01-09 12:56:34 +08:00
EdwinWong 8ed66b3260 fix: emb id 2024-01-04 17:59:34 +08:00
Radionic 7d0cfb6e5b fix: onnx model loading 2024-01-03 12:14:48 +08:00
Radionic a9c3075b70 feat: update sam node and combine points node 2024-01-03 12:14:35 +08:00
Radionic 0f29e2c27d fix: duplicated input 2024-01-02 11:15:41 +08:00
Radionic b3ee46fdf2 feat: extract boundary points and combine points node 2023-12-29 18:13:32 +08:00
23 changed files with 1363 additions and 585 deletions
+1 -1
View File
@@ -9,7 +9,6 @@ import sys
sys.path.append(os.path.join(os.path.dirname(__file__))) sys.path.append(os.path.join(os.path.dirname(__file__)))
import routes
import inspect import inspect
import sys import sys
import importlib import importlib
@@ -115,6 +114,7 @@ for path in paths:
NODE_CLASS_MAPPINGS = {} NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {} NODE_DISPLAY_NAME_MAPPINGS = {}
import routes
import blender_node import blender_node
base_class = blender_node.ObjectOps base_class = blender_node.ObjectOps
+1
View File
@@ -79,6 +79,7 @@ class AvatarMainOutput(blender_node.ObjectOps):
"files": [{ "files": [{
"filename": filepath.replace(f"{self.output_dir}/", ""), "filename": filepath.replace(f"{self.output_dir}/", ""),
"content_type": "model/gltf+json", "content_type": "model/gltf+json",
"type": "output"
},], },],
"SHAPE_FLOW": {SHAPE_FLOW}, "SHAPE_FLOW": {SHAPE_FLOW},
"auto_save": {'true' if auto_save else 'false'}, "auto_save": {'true' if auto_save else 'false'},
+78 -50
View File
@@ -3,7 +3,7 @@ import subprocess
import os import os
import folder_paths import folder_paths
import requests import requests
import json
def genreate_mesh_from_texture(bpy, image): def genreate_mesh_from_texture(bpy, image):
import torch import torch
@@ -14,11 +14,12 @@ def genreate_mesh_from_texture(bpy, image):
gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY) gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)
gray = (gray * 255).astype(np.uint8) gray = (gray * 255).astype(np.uint8)
# Find contours # Find contours
contours, _ = cv2.findContours( contours, _ = cv2.findContours(gray, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
gray, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
if len(contours) == 0: 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).") print("Warning: No contours found. The image may have 0 segment.")
black_image = torch.zeros(1, *image.shape)
return (black_image, None)
# Get the largest contour # Get the largest contour
areas = [cv2.contourArea(contour) for contour in contours] areas = [cv2.contourArea(contour) for contour in contours]
@@ -38,11 +39,12 @@ def genreate_mesh_from_texture(bpy, image):
for contour in contours: for contour in contours:
normalized_contour = [] normalized_contour = []
for vertex in contour: for vertex in contour:
normalized_vertex = [normalize_vertices( normalized_vertex = [
vertex[0][0], width), normalize_vertices(vertex[0][1], height) * -1] normalize_vertices(vertex[0][0], width),
normalize_vertices(vertex[0][1], height) * -1,
]
normalized_contour.append(normalized_vertex) normalized_contour.append(normalized_vertex)
normalized_contours.append( normalized_contours.append(np.array(normalized_contour, dtype=np.float32))
np.array(normalized_contour, dtype=np.float32))
meshes = [] meshes = []
# print(len(normalized_contours)) # print(len(normalized_contours))
@@ -63,12 +65,12 @@ def genreate_mesh_from_texture(bpy, image):
mesh.from_pydata(ordered_vertices, [], [face]) mesh.from_pydata(ordered_vertices, [], [face])
# Create a default shape key for the mesh # Create a default shape key for the mesh
sk_basis = obj.shape_key_add(name='Basis') sk_basis = obj.shape_key_add(name="Basis")
meshes.append(obj) # Add the object to the list of meshes meshes.append(obj) # Add the object to the list of meshes
# Draw contours on the original image # Draw contours on the original image
if not image.flags['C_CONTIGUOUS']: if not image.flags["C_CONTIGUOUS"]:
image = np.ascontiguousarray(image) image = np.ascontiguousarray(image)
cv2.drawContours(image, contours, -1, (0, 255, 0), 3) cv2.drawContours(image, contours, -1, (0, 255, 0), 3)
@@ -80,7 +82,7 @@ def genreate_mesh_from_texture(bpy, image):
def assign_texture(bpy, BPY_OBJ, texture, texture_name): def assign_texture(bpy, BPY_OBJ, texture, texture_name):
import numpy as np import numpy as np
import time import time
# Start the timer # Start the timer
start_time = time.time() start_time = time.time()
@@ -92,7 +94,8 @@ def assign_texture(bpy, BPY_OBJ, texture, texture_name):
# Create an image with the required dimensions # Create an image with the required dimensions
img = bpy.data.images.new( img = bpy.data.images.new(
texture_name, width=texture.shape[1], height=texture.shape[0], alpha = True) texture_name, width=texture.shape[1], height=texture.shape[0], alpha=True
)
# If there is no alpha channel, append one full of 1's # If there is no alpha channel, append one full of 1's
if texture.shape[2] == 3: if texture.shape[2] == 3:
@@ -109,7 +112,7 @@ def assign_texture(bpy, BPY_OBJ, texture, texture_name):
print(f"Time taken (texture.ravel) : {end_time - start_time} seconds") print(f"Time taken (texture.ravel) : {end_time - start_time} seconds")
# End the timer and print the time taken # End the timer and print the time taken
# Pack image to store it within .blend file # Pack image to store it within .blend file
img.pack() img.pack()
@@ -124,28 +127,27 @@ def assign_texture(bpy, BPY_OBJ, texture, texture_name):
# Create a material # Create a material
mat = bpy.data.materials.new("MaterialName") mat = bpy.data.materials.new("MaterialName")
mat.use_nodes = True mat.use_nodes = True
mat.blend_method = 'BLEND' mat.blend_method = "BLEND"
nodes = mat.node_tree.nodes nodes = mat.node_tree.nodes
for node in nodes: for node in nodes:
nodes.remove(node) nodes.remove(node)
# Add a new texture node # Add a new texture node
texture_node = nodes.new(type='ShaderNodeTexImage') texture_node = nodes.new(type="ShaderNodeTexImage")
texture_node.image = img texture_node.image = img
# Add a new BSDF node # Add a new BSDF node
bsdf_node = nodes.new(type='ShaderNodeBsdfPrincipled') bsdf_node = nodes.new(type="ShaderNodeBsdfPrincipled")
# Add a new output node # Add a new output node
output_node = nodes.new(type='ShaderNodeOutputMaterial') output_node = nodes.new(type="ShaderNodeOutputMaterial")
# Link nodes together # Link nodes together
links = mat.node_tree.links links = mat.node_tree.links
links.new(bsdf_node.inputs['Base Color'], links.new(bsdf_node.inputs["Base Color"], texture_node.outputs["Color"])
texture_node.outputs['Color']) links.new(output_node.inputs["Surface"], bsdf_node.outputs["BSDF"])
links.new(output_node.inputs['Surface'], bsdf_node.outputs['BSDF'])
links.new(bsdf_node.inputs['Alpha'], texture_node.outputs['Alpha']) links.new(bsdf_node.inputs["Alpha"], texture_node.outputs["Alpha"])
# Assign the material to the active object # Assign the material to the active object
if obj.data.materials: if obj.data.materials:
@@ -160,21 +162,30 @@ def assign_texture(bpy, BPY_OBJ, texture, texture_name):
blender_process_global = [] blender_process_global = []
def open_in_blender(blender_process, blender_path, output_file, camera_location=(0, 0, 0), camera_rotation=(0, 0, 0), shading="Material"): def open_in_blender(
blender_process,
blender_path,
output_file,
camera_location=(0, 0, 0),
camera_rotation=(0, 0, 0),
shading="Material",
):
import global_bpy import global_bpy
import mathutils import mathutils
bpy = global_bpy.get_bpy() bpy = global_bpy.get_bpy()
# Change shading mode and viewport # Change shading mode and viewport
for area in bpy.context.screen.areas: for area in bpy.context.screen.areas:
if area.type == 'VIEW_3D': if area.type == "VIEW_3D":
for space in area.spaces: for space in area.spaces:
if space.type == 'VIEW_3D': if space.type == "VIEW_3D":
space.shading.type = shading.upper() space.shading.type = shading.upper()
rv3d = space.region_3d rv3d = space.region_3d
rv3d.view_location = camera_location rv3d.view_location = camera_location
rv3d.view_rotation = mathutils.Euler( rv3d.view_rotation = mathutils.Euler(
camera_rotation).to_quaternion() camera_rotation
).to_quaternion()
# Open blender # Open blender
if blender_process != None: if blender_process != None:
@@ -187,8 +198,8 @@ def open_in_blender(blender_process, blender_path, output_file, camera_location=
os.remove(output_file) os.remove(output_file)
bpy.ops.wm.save_as_mainfile(filepath=output_file) bpy.ops.wm.save_as_mainfile(filepath=output_file)
print('blender_path', blender_path) print("blender_path", blender_path)
print('output_file', output_file) print("output_file", output_file)
blender_process = subprocess.Popen([blender_path, output_file]) blender_process = subprocess.Popen([blender_path, output_file])
# append to global list so it doesn't get garbage collected # append to global list so it doesn't get garbage collected
blender_process_global.append(blender_process) blender_process_global.append(blender_process)
@@ -198,18 +209,20 @@ def open_in_blender(blender_process, blender_path, output_file, camera_location=
# detects when the python process is killed, and kills the blender process # detects when the python process is killed, and kills the blender process
@atexit.register @atexit.register
def kill_blender_process(): def kill_blender_process():
print('blender_process_global', blender_process_global) print("blender_process_global", blender_process_global)
for process in blender_process_global: for process in blender_process_global:
process.kill() process.kill()
def export_gltf(output_dir, bpy_objects, filename, model_type, write_mode, metadata): def export_gltf(output_dir, bpy_objects, filename, model_type, write_mode, metadata):
import global_bpy import global_bpy
bpy = global_bpy.get_bpy() bpy = global_bpy.get_bpy()
# print(bpy, bpy_objects) # print(bpy, bpy_objects)
# deselect all objects # deselect all objects
override = bpy.context.copy() override = bpy.context.copy()
override["selected_objects"] = list(bpy_objects) override["selected_objects"] = list(bpy_objects)
@@ -230,36 +243,51 @@ def export_gltf(output_dir, bpy_objects, filename, model_type, write_mode, metad
return ".ava" return ".ava"
ext = get_file_extension(model_type) ext = get_file_extension(model_type)
filepath = output_dir + "/" + filename + ext + (".glb" if model_type == "AVA" else "") filepath = (
output_dir + "/" + filename + ext + (".glb" if model_type == "AVA" else "")
)
if write_mode == "Increment": if write_mode == "Increment":
count = 0 count = 0
# while file exists, increment count # while file exists, increment count
while os.path.exists(output_dir + "/" + filename + '_' + str(count) + ext): while os.path.exists(output_dir + "/" + filename + "_" + str(count) + ext):
count += 1 count += 1
filepath = output_dir + "/" + filename + '_' + str(count) + ext + (".glb" if model_type == "AVA" else "") filepath = (
output_dir
+ "/"
+ filename
+ "_"
+ str(count)
+ ext
+ (".glb" if model_type == "AVA" else "")
)
with bpy.context.temp_override(**override): with bpy.context.temp_override(**override):
bpy.ops.export_scene.gltf(filepath=filepath, export_format="GLB" if model_type == "AVA" else model_type, use_selection=True, export_extras=True) bpy.ops.export_scene.gltf(
filepath=filepath,
export_format="GLB" if model_type == "AVA" else model_type,
use_selection=True,
export_extras=True,
)
# print(filepath) # print(filepath)
if filepath.endswith('.ava.glb'): if filepath.endswith(".ava.glb"):
new_filepath = filepath.replace('.ava.glb', '.ava') new_filepath = filepath.replace(".ava.glb", ".ava")
os.replace(filepath, new_filepath) os.replace(filepath, new_filepath)
filepath = new_filepath filepath = new_filepath
return filepath return filepath
def get_avatar_file(output): def get_avatar_file(output):
avatar_filename = output["gltfFilename"][0] avatar_filename = output["gltfFilename"][0]
with open( with open(f"{folder_paths.get_output_directory()}/{avatar_filename}", "rb") as f:
f"{folder_paths.get_output_directory()}/{avatar_filename}", "rb"
) as f:
return f.read() return f.read()
def upload_avatar_file(output): def upload_avatar_file(output):
file = get_avatar_file(output) file = get_avatar_file(output)
response = requests.get("https://labs.avatech.ai/api/share") response = requests.get("https://labs.avatech.ai/api/share?version=v2")
labData = response.json() labData = response.json()
modelId = labData["modelId"] modelId = labData["modelId"]
@@ -272,15 +300,15 @@ def upload_avatar_file(output):
requests.put(labData["url"], headers=headers, data=file) requests.put(labData["url"], headers=headers, data=file)
# send notification # send notification
webhook_url = os.getenv("DISCORD_WEBHOOK_URL") # webhook_url = os.getenv("DISCORD_WEBHOOK_URL")
data = { # data = {
"username": "Avabot", # "username": "Avabot",
"avatar_url": "https://avatech-avatar-dev1.nyc3.cdn.digitaloceanspaces.com/avatechai.png", # "avatar_url": "https://avatech-avatar-dev1.nyc3.cdn.digitaloceanspaces.com/avatechai.png",
"content": "[API Call] New register!", # "content": "[API Call] New register!",
} # }
headers = { # headers = {
"Content-Type": "application/json", # "Content-Type": "application/json",
} # }
response = requests.post(webhook_url, headers=headers, data=json.dumps(data)) # response = requests.post(webhook_url, headers=headers, data=json.dumps(data))
return modelId return modelId
+3 -1
View File
@@ -24,6 +24,8 @@ class Object_CreateMeshLayer(blender_node.ObjectOps):
def blender_process(self, bpy, image, convex_hull, shape_threshold, mesh_layer_name, scale_x,scale_y , extrude_x, extrude_y, seed): def blender_process(self, bpy, image, convex_hull, shape_threshold, mesh_layer_name, scale_x,scale_y , extrude_x, extrude_y, seed):
image, BPY_OBJ = genreate_mesh_from_texture(bpy, image) image, BPY_OBJ = genreate_mesh_from_texture(bpy, image)
if BPY_OBJ is None:
return (None, image)
bpy.context.view_layer.objects.active = BPY_OBJ bpy.context.view_layer.objects.active = BPY_OBJ
@@ -41,7 +43,7 @@ class Object_CreateMeshLayer(blender_node.ObjectOps):
bpy.ops.mesh.select_all(action='SELECT') bpy.ops.mesh.select_all(action='SELECT')
bpy.ops.mesh.edge_face_add() bpy.ops.mesh.edge_face_add()
bpy.ops.transform.resize(value=(scale_x, scale_y, 1)) bpy.ops.transform.resize(value=(float(scale_x), float(scale_y), 1))
bpy.context.object.vertex_groups.new(name=mesh_layer_name) bpy.context.object.vertex_groups.new(name=mesh_layer_name)
bpy.ops.object.vertex_group_assign() bpy.ops.object.vertex_group_assign()
@@ -28,6 +28,8 @@ class Object_CreateMeshLayer_Advanced(blender_node.ObjectOps):
def blender_process(self, bpy, image, convex_hull, shape_threshold, mesh_layer_name, scale_x,scale_y , extrude_x, extrude_y, inner_translate_x, inner_translate_y, outer_translate_x, outer_translate_y, seed): def blender_process(self, bpy, image, convex_hull, shape_threshold, mesh_layer_name, scale_x,scale_y , extrude_x, extrude_y, inner_translate_x, inner_translate_y, outer_translate_x, outer_translate_y, seed):
image, BPY_OBJ = genreate_mesh_from_texture(bpy, image) image, BPY_OBJ = genreate_mesh_from_texture(bpy, image)
if BPY_OBJ is None:
return (None, image)
bpy.context.view_layer.objects.active = BPY_OBJ bpy.context.view_layer.objects.active = BPY_OBJ
+3 -6
View File
@@ -2,14 +2,12 @@ import blender_node
class Mesh_JoinMesh(blender_node.ObjectOps): class Mesh_JoinMesh(blender_node.ObjectOps):
EXTRA_INPUT_TYPES = { EXTRA_INPUT_TYPES = {"BPY_OBJ2": (blender_node.BPY_OBJ,)}
"BPY_OBJ2": (blender_node.BPY_OBJ,)
}
CUSTOM_NAME = "Join Meshes" CUSTOM_NAME = "Join Meshes"
def blender_process(self, bpy, BPY_OBJ, **props): def blender_process(self, bpy, BPY_OBJ, **props):
prop_values = props.values() prop_values = props.values()
for obj in list(prop_values) + [BPY_OBJ]: for obj in list(prop_values) + [BPY_OBJ]:
if obj is not None: if obj is not None:
obj.select_set(True) obj.select_set(True)
@@ -18,4 +16,3 @@ class Mesh_JoinMesh(blender_node.ObjectOps):
bpy.ops.object.join() bpy.ops.object.join()
return (BPY_OBJ,) return (BPY_OBJ,)
+26
View File
@@ -0,0 +1,26 @@
import blender_node
class Mesh_SetShapeKeyValue(blender_node.ObjectOps):
CUSTOM_NAME = "Set Shape Key Value"
EXTRA_INPUT_TYPES = {
"shape_key_name": ("STRING", {
"multiline": False,
"default": "my_shape_key",
}),
"value": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "display": "number"}),
}
def blender_process(self, bpy, BPY_OBJ, shape_key_name, value):
# Check if the object has shape keys
if BPY_OBJ.data.shape_keys:
# Check if the specified shape key exists
if shape_key_name in BPY_OBJ.data.shape_keys.key_blocks:
BPY_OBJ.data.shape_keys.key_blocks[shape_key_name].value = float(value)
else:
print(f"The shape key {shape_key_name} does not exist on the object.")
else:
print("The object does not have any shape keys.")
return (BPY_OBJ,)
+132
View File
@@ -0,0 +1,132 @@
import blender_node
import math
import folder_paths
import torch
import numpy as np
import os
from PIL import Image, ImageOps
def get_incremented_filename(folder_path, base_filename):
# Initialize the counter and create the full initial path
counter = 0
output_path = f"{folder_path}/{base_filename}.png"
# Check if the file exists and increment the counter until the file does not exist
while os.path.exists(output_path):
counter += 1
output_path = f"{folder_path}/{base_filename}_{counter}.png"
return output_path
class BlenderRenderImage(blender_node.ObjectOps):
def __init__(self):
pass
EXTRA_INPUT_TYPES = {
}
# OUTPUT_NODE = True
RETURN_TYPES = ("BPY_OBJ", "IMAGE")
def add_light(self, bpy):
# Check if there is at least one light source in the scene
light_exists = any(ob for ob in bpy.data.objects if ob.type == 'LIGHT')
if not light_exists:
# Create a new Area light datablock for ambient light
light_data = bpy.data.lights.new(name='AmbientLight', type='AREA')
light_object = bpy.data.objects.new(name='AmbientLight', object_data=light_data)
bpy.context.collection.objects.link(light_object)
# Position the light in the scene
light_object.location = (0, 0, 10)
# Set light size for soft shadows and ambient effect
light_data.size = 10
light_data.energy = 1000
print("Added an ambient light source to the scene.")
def add_camera(self, bpy):
# Check if there is a camera in the scene
if bpy.context.scene.camera:
return bpy.context.scene.camera
# If not, create a new camera
cam_data = bpy.data.cameras.new(name='Camera')
cam = bpy.data.objects.new(name='Camera', object_data=cam_data)
bpy.context.collection.objects.link(cam)
# Set the new camera to the active camera
bpy.context.scene.camera = cam
# Position the camera to a default view
cam.location = (0, 0, 10)
return cam
def get_texture_size(self, obj):
# Get the first material slot
mat = obj.data.materials[0]
# Check if the material has a node tree
if mat.node_tree:
nodes = mat.node_tree.nodes
# Find an image texture node in the node tree
for node in nodes:
if node.type == 'TEX_IMAGE':
texture = node.image
if texture:
return texture.size
print("No image texture node found in the material's node tree.")
else:
print("Material has no node tree.")
def blender_process(self, bpy, BPY_OBJ=None):
cam = self.add_camera(bpy)
plane = BPY_OBJ
if plane:
tex_width, tex_height = self.get_texture_size(plane)
# Calculate the aspect ratio of the plane
aspect_ratio_plane = tex_width / tex_height
# Set the render resolution to match the plane's aspect ratio
# Choose an arbitrary resolution for the longer side of the plane
base_resolution = 512
if aspect_ratio_plane > 1:
# Plane is wider than it is tall
bpy.context.scene.render.resolution_x = base_resolution
bpy.context.scene.render.resolution_y = int(base_resolution / aspect_ratio_plane)
else:
# Plane is taller than it is wide
bpy.context.scene.render.resolution_x = int(base_resolution * aspect_ratio_plane)
bpy.context.scene.render.resolution_y = base_resolution
bpy.context.scene.render.resolution_percentage = 100
ortho_scale = max(plane.dimensions.x, plane.dimensions.y)
cam.data.type = 'ORTHO'
cam.data.ortho_scale = ortho_scale
self.add_light(bpy)
# Update the scene to reflect changes
bpy.context.view_layer.update()
# Set render engine (e.g., 'BLENDER_EEVEE', 'CYCLES', 'BLENDER_WORKBENCH')
bpy.context.scene.render.engine = "BLENDER_EEVEE"
# Specify the render output path
output_path = get_incremented_filename(folder_paths.get_output_directory(), "render")
bpy.context.scene.render.filepath = output_path
# Render the image
bpy.ops.render.render(write_still=True)
# Load the image
i = Image.open(output_path)
i = ImageOps.exif_transpose(i)
image = i.convert("RGB")
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
# print(image.shape)
return (BPY_OBJ, image)
+1 -1
View File
@@ -706,7 +706,7 @@ export function AvatarPreview() {
console.log(showPreview); console.log(showPreview);
return ( return (
(showPreview.val ? "" : "hidden ") + (showPreview.val && !showEditor.val ? "" : "hidden ") +
"absolute w-[360px] h-[360px] rounded-xl overflow-hidden right-0 top-0 z-[99] pointer-events-auto flex border-none bg-transparent" "absolute w-[360px] h-[360px] rounded-xl overflow-hidden right-0 top-0 z-[99] pointer-events-auto flex border-none bg-transparent"
); );
}, },
+146
View File
@@ -0,0 +1,146 @@
import { combinePointsNode, samPrompts } from "./state.js";
import { van } from "./van.js";
const { div, dialog, form, button, h3, input, span } = van.tags;
van.derive(() => {
if (
combinePointsNode.val != undefined &&
combinePointsNode.val.type === "Combine Points"
) {
const inputNames = combinePointsNode.val.inputs?.map((x) => x.name) || [];
const record = Object.keys(samPrompts.val);
const diff = inputNames.filter((x) => !record.includes(x));
const missingDiff = record.filter((x) => !inputNames.includes(x));
if (diff.length > 0) {
diff.forEach((x) => {
combinePointsNode.val.removeInput(
combinePointsNode.val.findInputSlot(x)
);
});
combinePointsNode.val.graph.change();
}
if (missingDiff.length > 0) {
missingDiff.forEach((x) => {
combinePointsNode.val.addInput(x, "POINTS");
});
combinePointsNode.val.graph.change();
}
}
});
export function CombinePointsDialog() {
const showAddLayer = van.state(false);
return div(
{
class: () =>
"absolute z-[100] top-0 left-0 flex justify-center w-full h-full ",
},
() =>
dialog(
{ id: "combine_points_dialog", class: "modal" },
div(
{ class: "modal-box text-base-content" },
form(
{
class: "gap-2 flex flex-col",
method: "dialog",
onsubmit: (e) => {
e.preventDefault();
combine_points_dialog.close();
},
},
button(
{
type: "button",
class: "btn btn-sm btn-circle btn-ghost absolute right-2 top-2",
onclick: (e) => {
e.stopPropagation();
combine_points_dialog.close();
},
},
"✕"
),
h3({ class: "font-bold text-lg text-base-content" }, "Edit points"),
() =>
div(
{ class: "flex flex-col gap-2 mb-2" },
...Object.keys(samPrompts.val).map((key) => {
return span(key);
})
),
() =>
showAddLayer.val
? div(
{
class:
"flex flex-row justify-center items-center border rounded-md pr-2",
},
input({
type: "text",
placeholder: "Type here",
id: "layerName",
class:
"input input-ghost w-full focus:ring-0 focus:border-none focus:outline-none",
autofocus: true,
}),
button(
{
onclick: (e) => {
e.stopPropagation();
e.preventDefault();
showAddLayer.val = false;
},
},
span({
class: "iconify text-2xl",
"data-icon": "iconoir:cancel",
})
),
button(
{
onclick: (e) => {
e.stopPropagation();
e.preventDefault();
showAddLayer.val = false;
const inputText =
document.getElementById("layerName").value;
samPrompts.val = {
...samPrompts.val,
[inputText]: [],
};
},
},
span({
class: "iconify text-2xl",
"data-icon": "mdi:tick",
})
)
)
: button(
{
class: "btn btn-outline btn",
onclick: (e) => {
e.stopPropagation();
e.preventDefault();
showAddLayer.val = true;
},
},
"Add new layer"
),
button(
{
type: "submit",
class: "btn btn-sm btn-ghost place-self-end",
},
"Confirm"
)
)
)
)
);
}
+11 -9
View File
@@ -1,18 +1,20 @@
import { LayerEditor } from './LayerEditor.js'; import { LayerEditor } from "./LayerEditor.js";
import { ShapeFlowEditor } from './ShapeFlowEditor.js'; import { ShapeFlowEditor } from "./ShapeFlowEditor.js";
import { van } from './van.js'; import { van } from "./van.js";
import { AvatarPreview } from './AvatarPreview.js'; import { AvatarPreview } from "./AvatarPreview.js";
import { Loading } from './Loading.js'; import { Loading } from "./Loading.js";
import { Alert } from './Alert.js'; import { Alert } from "./Alert.js";
import { AppHeader } from './AppHeader.js'; import { AppHeader } from "./AppHeader.js";
import { CombinePointsDialog } from "./CombinePointsDialog.js";
const { button, iframe, div, img } = van.tags; const { button, iframe, div, img } = van.tags;
export function Container() { export function Container() {
return div( return div(
{ {
class: 'fixed left-0 top-0 w-full h-full z-[1000] pointer-events-none', class: "fixed left-0 top-0 w-full h-full z-[1000] pointer-events-none",
id: 'avatech-editor', id: "avatech-editor",
}, },
CombinePointsDialog(),
ShapeFlowEditor(), ShapeFlowEditor(),
LayerEditor(), LayerEditor(),
AvatarPreview(), AvatarPreview(),
+29 -24
View File
@@ -6,6 +6,7 @@ import {
targetNode, targetNode,
showImageEditor, showImageEditor,
allImagePrompts, allImagePrompts,
samPrompts,
} from "./state.js"; } from "./state.js";
import { van } from "./van.js"; import { van } from "./van.js";
const { const {
@@ -24,6 +25,33 @@ const {
span, span,
} = van.tags; } = van.tags;
export const updateOutputs = () => {
const outputNames = targetNode.val.outputs.map((x) => x.name).slice(1);
const record = Object.keys(imagePromptsMulti.val);
const diff = outputNames.filter((x) => !record.includes(x));
const missingDiff = record.filter((x) => !outputNames.includes(x));
if (diff.length > 0) {
console.log("Cleaning up missing output slots", diff);
diff.forEach((x) => {
targetNode.val.removeOutput(targetNode.val.findOutputSlot(x));
});
targetNode.val.graph.change();
}
if (missingDiff.length > 0) {
console.log("Adding missing output slots", diff);
missingDiff.forEach((x) => {
targetNode.val.addOutput(
x,
targetNode.val.type === "SAM MultiLayer" ? "IMAGE" : "SAM_PROMPT"
);
});
targetNode.val.graph.change();
}
};
van.derive(() => { van.derive(() => {
if ( if (
showImageEditor.val && showImageEditor.val &&
@@ -31,30 +59,7 @@ van.derive(() => {
targetNode.val.outputs && targetNode.val.outputs &&
targetNode.val.type === "SAM MultiLayer" targetNode.val.type === "SAM MultiLayer"
) { ) {
const outputNames = targetNode.val.outputs.map((x) => x.name).slice(1); updateOutputs();
const record = Object.keys(imagePromptsMulti.val);
const diff = outputNames.filter((x) => !record.includes(x));
const missingDiff = record.filter((x) => !outputNames.includes(x));
if (diff.length > 0) {
console.log("Cleaning up missing output slots", diff);
diff.forEach((x) => {
targetNode.val.removeOutput(targetNode.val.findOutputSlot(x));
});
targetNode.val.graph.change();
}
if (missingDiff.length > 0) {
console.log("Adding missing output slots", diff);
missingDiff.forEach((x) => {
targetNode.val.addOutput(
x,
targetNode.val.type === "SAM MultiLayer" ? "IMAGE" : "SAM_PROMPT"
);
});
targetNode.val.graph.change();
}
} }
}); });
+46 -6
View File
@@ -5,6 +5,7 @@ import {
imageUrl, imageUrl,
imagePrompts, imagePrompts,
targetNode, targetNode,
combinePointsNode,
fileName, fileName,
embeddings, embeddings,
imagePromptsMulti, imagePromptsMulti,
@@ -17,6 +18,7 @@ import {
previewModelId, previewModelId,
embeddingID, embeddingID,
enableAutoSegment, enableAutoSegment,
samPrompts,
} from "./state.js"; } from "./state.js";
import { van } from "./van.js"; import { van } from "./van.js";
import { app } from "./app.js"; import { app } from "./app.js";
@@ -33,6 +35,7 @@ import {
import { infoDialog } from "./dialog.js"; import { infoDialog } from "./dialog.js";
import { sharedAvatarLink } from "./AvatarPreview.js"; import { sharedAvatarLink } from "./AvatarPreview.js";
import { updateImagePrompts } from "./LayerEditor.js"; import { updateImagePrompts } from "./LayerEditor.js";
import { updateOutputs } from "./SideBar.js";
export const generatedImages = {}; export const generatedImages = {};
@@ -318,7 +321,6 @@ function showMyImageEditor(node) {
isGeneratedImage, isGeneratedImage,
embedding_id: id, embedding_id: id,
ckpt, ckpt,
// remote: true,
}), }),
}) })
.then(() => { .then(() => {
@@ -357,7 +359,6 @@ function showMyImageEditor(node) {
drawSegment(getClicks()); drawSegment(getClicks());
updateImagePrompts(); updateImagePrompts();
}); });
targetNode.val = node;
}) })
.catch((err) => { .catch((err) => {
console.log(err); console.log(err);
@@ -385,6 +386,25 @@ const ext = {
showMyImageEditor(node); showMyImageEditor(node);
btn.serialize = false; btn.serialize = false;
}); });
targetNode.val = node;
node.onConnectInput = (node, slot, targetSlot) => {
if (targetSlot.name === "SAM_PROMPTS") {
imagePromptsMulti.val = samPrompts.val;
updateOutputs();
}
};
return {
widget: btn,
};
},
COMBINE_POINTS(node, inputName, inputData, app) {
const btn = node.addWidget("button", "Edit points", "", () => {
console.log("Edit points");
combine_points_dialog.showModal();
combinePointsNode.val = node;
});
return { return {
widget: btn, widget: btn,
}; };
@@ -675,6 +695,18 @@ const ext = {
nodeData.input.required.sam = ["SAM_PROMPTS"]; nodeData.input.required.sam = ["SAM_PROMPTS"];
// nodeData.input.required.upload = ['IMAGEUPLOAD']; // nodeData.input.required.upload = ['IMAGEUPLOAD'];
// nodeData.input.required.prompts_points = ["IMAGEUPLOAD"]; // nodeData.input.required.prompts_points = ["IMAGEUPLOAD"];
addMenuHandler(nodeType, function (_, options) {
options.unshift({
content: "Open In Points Editor (Local)",
callback: () => {
showMyImageEditor(this);
},
});
});
break;
case "Combine Points":
nodeData.input.required.sam = ["COMBINE_POINTS"];
addMenuHandler(nodeType, function (_, options) { addMenuHandler(nodeType, function (_, options) {
options.unshift({ options.unshift({
content: "Open In Points Editor (Local)", content: "Open In Points Editor (Local)",
@@ -746,11 +778,19 @@ function injectUIComponentToComfyuimenu() {
if (!filename.toLowerCase().endsWith(".json")) { if (!filename.toLowerCase().endsWith(".json")) {
filename += ".json"; filename += ".json";
} }
app.graphToPrompt().then(p=>{ app.graphToPrompt().then((p) => {
console.log('fkfk');
let json = JSON.stringify(p.output, null, 2); // convert the data to a JSON string 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"'); json = json
const blob = new Blob([json], {type: "application/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); const url = URL.createObjectURL(blob);
a.href = url; a.href = url;
a.download = filename; a.download = filename;
+6 -4
View File
@@ -1,16 +1,18 @@
import "https://cdn.jsdelivr.net/npm/onnxruntime-web/dist/ort.min.js"; import "https://cdn.jsdelivr.net/npm/onnxruntime-web@1.16.3/dist/ort.min.js";
import npyjs from "https://esm.sh/npyjs"; import npyjs from "https://esm.sh/npyjs";
import { imageSize } from "./state.js"; import { imageSize } from "./state.js";
import { modelData, onnxMaskToImage } from "./onnx_helper.js"; import { modelData, onnxMaskToImage } from "./onnx_helper.js";
ort.env.wasm.wasmPaths = "https://cdn.jsdelivr.net/npm/onnxruntime-web/dist/"; ort.env.wasm.wasmPaths = "https://cdn.jsdelivr.net/npm/onnxruntime-web@1.16.3/dist/";
export let model = null; export let model = null;
let modelType = null;
// Initialize the ONNX model // Initialize the ONNX model
export const initModel = async (modelType) => { export const initModel = async (type) => {
try { try {
if (!model) { if (!model || modelType !== type) {
modelType = type;
model = await ort.InferenceSession.create( model = await ort.InferenceSession.create(
`${location.protocol}//${location.host}/sam_model?type=${modelType}` `${location.protocol}//${location.host}/sam_model?type=${modelType}`
); );
+8 -3
View File
@@ -19,11 +19,13 @@
*/ */
import { van } from "./van.js"; import { van } from "./van.js";
export const iframeSrc = van.state("https://editor.avatech.ai?comfyui=true"); export const iframeSrc = van.state("https://editor.avatech.ai?comfyui=true");
export const showEditor = van.state(false); export const showEditor = van.state(false);
// localStorage.getItem("showPreview") == 'true' // localStorage.getItem("showPreview") == 'true'
export const showPreview = van.state(true); console.log(localStorage.getItem("showPreview"));
if (localStorage.getItem("showPreview") == null)
localStorage.setItem("showPreview", 'true')
export const showPreview = van.state(localStorage.getItem("showPreview") == 'true');
export const previewUrl = van.state( export const previewUrl = van.state(
"https://editor.avatech.ai/viewer?avatarId=default&debug=true&width=350&height=350&hideTrigger=true&voiceSelection=true&hideUI=true" "https://editor.avatech.ai/viewer?avatarId=default&debug=true&width=350&height=350&hideTrigger=true&voiceSelection=true&hideUI=true"
); );
@@ -63,7 +65,6 @@ export const imagePrompts = van.state([]);
export const allImagePrompts = van.state([{}]); export const allImagePrompts = van.state([{}]);
/** @type {State<Record<string, Point[]>>} */ /** @type {State<Record<string, Point[]>>} */
export const imagePromptsMulti = van.state({}); export const imagePromptsMulti = van.state({});
@@ -76,3 +77,7 @@ export const targetNode = van.state();
export const imageSize = van.state({ width: 0, height: 0, samScale: 0 }); export const imageSize = van.state({ width: 0, height: 0, samScale: 0 });
export const embeddings = van.state(); export const embeddings = van.state();
export const embeddingID = van.state("Test"); export const embeddingID = van.state("Test");
/** @type {State<LGraphNode>} */
export const combinePointsNode = van.state();
export const samPrompts = van.state({});
+261 -88
View File
@@ -1,7 +1,7 @@
@import url('https://fonts.googleapis.com/css2?family=Gabarito&display=swap'); @import url('https://fonts.googleapis.com/css2?family=Gabarito&display=swap');
/* /*
! tailwindcss v3.4.0 | MIT License | https://tailwindcss.com ! tailwindcss v3.3.5 | MIT License | https://tailwindcss.com
*/ */
/* /*
@@ -34,11 +34,9 @@
4. Use the user's configured `sans` font-family by default. 4. Use the user's configured `sans` font-family by default.
5. Use the user's configured `sans` font-feature-settings by default. 5. Use the user's configured `sans` font-feature-settings by default.
6. Use the user's configured `sans` font-variation-settings by default. 6. Use the user's configured `sans` font-variation-settings by default.
7. Disable tap highlights on iOS
*/ */
html, html {
:host {
line-height: 1.5; line-height: 1.5;
/* 1 */ /* 1 */
-webkit-text-size-adjust: 100%; -webkit-text-size-adjust: 100%;
@@ -48,14 +46,12 @@ html,
-o-tab-size: 4; -o-tab-size: 4;
tab-size: 4; tab-size: 4;
/* 3 */ /* 3 */
font-family: ui-sans-serif, system-ui, sans-serif, "Apple Color Emoji", "Segoe UI Emoji", "Segoe UI Symbol", "Noto Color Emoji"; font-family: ui-sans-serif, system-ui, -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, "Helvetica Neue", Arial, "Noto Sans", sans-serif, "Apple Color Emoji", "Segoe UI Emoji", "Segoe UI Symbol", "Noto Color Emoji";
/* 4 */ /* 4 */
font-feature-settings: normal; font-feature-settings: normal;
/* 5 */ /* 5 */
font-variation-settings: normal; font-variation-settings: normal;
/* 6 */ /* 6 */
-webkit-tap-highlight-color: transparent;
/* 7 */
} }
/* /*
@@ -127,10 +123,8 @@ strong {
} }
/* /*
1. Use the user's configured `mono` font-family by default. 1. Use the user's configured `mono` font family by default.
2. Use the user's configured `mono` font-feature-settings by default. 2. Correct the odd `em` font sizing in all browsers.
3. Use the user's configured `mono` font-variation-settings by default.
4. Correct the odd `em` font sizing in all browsers.
*/ */
code, code,
@@ -139,12 +133,8 @@ samp,
pre { pre {
font-family: ui-monospace, SFMono-Regular, Menlo, Monaco, Consolas, "Liberation Mono", "Courier New", monospace; font-family: ui-monospace, SFMono-Regular, Menlo, Monaco, Consolas, "Liberation Mono", "Courier New", monospace;
/* 1 */ /* 1 */
font-feature-settings: normal;
/* 2 */
font-variation-settings: normal;
/* 3 */
font-size: 1em; font-size: 1em;
/* 4 */ /* 2 */
} }
/* /*
@@ -1011,6 +1001,99 @@ html{
} }
} }
.btn-outline:hover{
--tw-border-opacity: 1;
border-color: var(--fallback-bc,oklch(var(--bc)/var(--tw-border-opacity)));
--tw-bg-opacity: 1;
background-color: var(--fallback-bc,oklch(var(--bc)/var(--tw-bg-opacity)));
--tw-text-opacity: 1;
color: var(--fallback-b1,oklch(var(--b1)/var(--tw-text-opacity)));
}
.btn-outline.btn-primary:hover{
--tw-text-opacity: 1;
color: var(--fallback-pc,oklch(var(--pc)/var(--tw-text-opacity)));
}
@supports (color: color-mix(in oklab, black, black)){
.btn-outline.btn-primary:hover{
background-color: color-mix(in oklab, var(--fallback-p,oklch(var(--p)/1)) 90%, black);
border-color: color-mix(in oklab, var(--fallback-p,oklch(var(--p)/1)) 90%, black);
}
}
.btn-outline.btn-secondary:hover{
--tw-text-opacity: 1;
color: var(--fallback-sc,oklch(var(--sc)/var(--tw-text-opacity)));
}
@supports (color: color-mix(in oklab, black, black)){
.btn-outline.btn-secondary:hover{
background-color: color-mix(in oklab, var(--fallback-s,oklch(var(--s)/1)) 90%, black);
border-color: color-mix(in oklab, var(--fallback-s,oklch(var(--s)/1)) 90%, black);
}
}
.btn-outline.btn-accent:hover{
--tw-text-opacity: 1;
color: var(--fallback-ac,oklch(var(--ac)/var(--tw-text-opacity)));
}
@supports (color: color-mix(in oklab, black, black)){
.btn-outline.btn-accent:hover{
background-color: color-mix(in oklab, var(--fallback-a,oklch(var(--a)/1)) 90%, black);
border-color: color-mix(in oklab, var(--fallback-a,oklch(var(--a)/1)) 90%, black);
}
}
.btn-outline.btn-success:hover{
--tw-text-opacity: 1;
color: var(--fallback-suc,oklch(var(--suc)/var(--tw-text-opacity)));
}
@supports (color: color-mix(in oklab, black, black)){
.btn-outline.btn-success:hover{
background-color: color-mix(in oklab, var(--fallback-su,oklch(var(--su)/1)) 90%, black);
border-color: color-mix(in oklab, var(--fallback-su,oklch(var(--su)/1)) 90%, black);
}
}
.btn-outline.btn-info:hover{
--tw-text-opacity: 1;
color: var(--fallback-inc,oklch(var(--inc)/var(--tw-text-opacity)));
}
@supports (color: color-mix(in oklab, black, black)){
.btn-outline.btn-info:hover{
background-color: color-mix(in oklab, var(--fallback-in,oklch(var(--in)/1)) 90%, black);
border-color: color-mix(in oklab, var(--fallback-in,oklch(var(--in)/1)) 90%, black);
}
}
.btn-outline.btn-warning:hover{
--tw-text-opacity: 1;
color: var(--fallback-wac,oklch(var(--wac)/var(--tw-text-opacity)));
}
@supports (color: color-mix(in oklab, black, black)){
.btn-outline.btn-warning:hover{
background-color: color-mix(in oklab, var(--fallback-wa,oklch(var(--wa)/1)) 90%, black);
border-color: color-mix(in oklab, var(--fallback-wa,oklch(var(--wa)/1)) 90%, black);
}
}
.btn-outline.btn-error:hover{
--tw-text-opacity: 1;
color: var(--fallback-erc,oklch(var(--erc)/var(--tw-text-opacity)));
}
@supports (color: color-mix(in oklab, black, black)){
.btn-outline.btn-error:hover{
background-color: color-mix(in oklab, var(--fallback-er,oklch(var(--er)/1)) 90%, black);
border-color: color-mix(in oklab, var(--fallback-er,oklch(var(--er)/1)) 90%, black);
}
}
.btn-disabled:hover, .btn-disabled:hover,
.btn[disabled]:hover, .btn[disabled]:hover,
.btn:disabled:hover{ .btn:disabled:hover{
@@ -1353,6 +1436,43 @@ html{
} }
} }
@supports (color: color-mix(in oklab, black, black)){
.btn-outline.btn-primary.btn-active{
background-color: color-mix(in oklab, var(--fallback-p,oklch(var(--p)/1)) 90%, black);
border-color: color-mix(in oklab, var(--fallback-p,oklch(var(--p)/1)) 90%, black);
}
.btn-outline.btn-secondary.btn-active{
background-color: color-mix(in oklab, var(--fallback-s,oklch(var(--s)/1)) 90%, black);
border-color: color-mix(in oklab, var(--fallback-s,oklch(var(--s)/1)) 90%, black);
}
.btn-outline.btn-accent.btn-active{
background-color: color-mix(in oklab, var(--fallback-a,oklch(var(--a)/1)) 90%, black);
border-color: color-mix(in oklab, var(--fallback-a,oklch(var(--a)/1)) 90%, black);
}
.btn-outline.btn-success.btn-active{
background-color: color-mix(in oklab, var(--fallback-su,oklch(var(--su)/1)) 90%, black);
border-color: color-mix(in oklab, var(--fallback-su,oklch(var(--su)/1)) 90%, black);
}
.btn-outline.btn-info.btn-active{
background-color: color-mix(in oklab, var(--fallback-in,oklch(var(--in)/1)) 90%, black);
border-color: color-mix(in oklab, var(--fallback-in,oklch(var(--in)/1)) 90%, black);
}
.btn-outline.btn-warning.btn-active{
background-color: color-mix(in oklab, var(--fallback-wa,oklch(var(--wa)/1)) 90%, black);
border-color: color-mix(in oklab, var(--fallback-wa,oklch(var(--wa)/1)) 90%, black);
}
.btn-outline.btn-error.btn-active{
background-color: color-mix(in oklab, var(--fallback-er,oklch(var(--er)/1)) 90%, black);
border-color: color-mix(in oklab, var(--fallback-er,oklch(var(--er)/1)) 90%, black);
}
}
.btn:focus-visible{ .btn:focus-visible{
outline-style: solid; outline-style: solid;
outline-width: 2px; outline-width: 2px;
@@ -1399,6 +1519,95 @@ html{
background-color: var(--fallback-bc,oklch(var(--bc)/0.2)); background-color: var(--fallback-bc,oklch(var(--bc)/0.2));
} }
.btn-outline{
border-color: currentColor;
background-color: transparent;
--tw-text-opacity: 1;
color: var(--fallback-bc,oklch(var(--bc)/var(--tw-text-opacity)));
--tw-shadow: 0 0 #0000;
--tw-shadow-colored: 0 0 #0000;
box-shadow: var(--tw-ring-offset-shadow, 0 0 #0000), var(--tw-ring-shadow, 0 0 #0000), var(--tw-shadow);
}
.btn-outline.btn-active{
--tw-border-opacity: 1;
border-color: var(--fallback-bc,oklch(var(--bc)/var(--tw-border-opacity)));
--tw-bg-opacity: 1;
background-color: var(--fallback-bc,oklch(var(--bc)/var(--tw-bg-opacity)));
--tw-text-opacity: 1;
color: var(--fallback-b1,oklch(var(--b1)/var(--tw-text-opacity)));
}
.btn-outline.btn-primary{
--tw-text-opacity: 1;
color: var(--fallback-p,oklch(var(--p)/var(--tw-text-opacity)));
}
.btn-outline.btn-primary.btn-active{
--tw-text-opacity: 1;
color: var(--fallback-pc,oklch(var(--pc)/var(--tw-text-opacity)));
}
.btn-outline.btn-secondary{
--tw-text-opacity: 1;
color: var(--fallback-s,oklch(var(--s)/var(--tw-text-opacity)));
}
.btn-outline.btn-secondary.btn-active{
--tw-text-opacity: 1;
color: var(--fallback-sc,oklch(var(--sc)/var(--tw-text-opacity)));
}
.btn-outline.btn-accent{
--tw-text-opacity: 1;
color: var(--fallback-a,oklch(var(--a)/var(--tw-text-opacity)));
}
.btn-outline.btn-accent.btn-active{
--tw-text-opacity: 1;
color: var(--fallback-ac,oklch(var(--ac)/var(--tw-text-opacity)));
}
.btn-outline.btn-success{
--tw-text-opacity: 1;
color: var(--fallback-su,oklch(var(--su)/var(--tw-text-opacity)));
}
.btn-outline.btn-success.btn-active{
--tw-text-opacity: 1;
color: var(--fallback-suc,oklch(var(--suc)/var(--tw-text-opacity)));
}
.btn-outline.btn-info{
--tw-text-opacity: 1;
color: var(--fallback-in,oklch(var(--in)/var(--tw-text-opacity)));
}
.btn-outline.btn-info.btn-active{
--tw-text-opacity: 1;
color: var(--fallback-inc,oklch(var(--inc)/var(--tw-text-opacity)));
}
.btn-outline.btn-warning{
--tw-text-opacity: 1;
color: var(--fallback-wa,oklch(var(--wa)/var(--tw-text-opacity)));
}
.btn-outline.btn-warning.btn-active{
--tw-text-opacity: 1;
color: var(--fallback-wac,oklch(var(--wac)/var(--tw-text-opacity)));
}
.btn-outline.btn-error{
--tw-text-opacity: 1;
color: var(--fallback-er,oklch(var(--er)/var(--tw-text-opacity)));
}
.btn-outline.btn-error.btn-active{
--tw-text-opacity: 1;
color: var(--fallback-erc,oklch(var(--erc)/var(--tw-text-opacity)));
}
.btn.btn-disabled, .btn.btn-disabled,
.btn[disabled], .btn[disabled],
.btn:disabled{ .btn:disabled{
@@ -1570,6 +1779,18 @@ details.collapse summary::-webkit-details-marker{
outline-color: var(--fallback-bc,oklch(var(--bc)/0.2)); outline-color: var(--fallback-bc,oklch(var(--bc)/0.2));
} }
.input-ghost{
--tw-bg-opacity: 0.05;
}
.input-ghost:focus,
.input-ghost:focus-within{
--tw-bg-opacity: 1;
--tw-text-opacity: 1;
color: var(--fallback-bc,oklch(var(--bc)/var(--tw-text-opacity)));
box-shadow: none;
}
.input-disabled, .input-disabled,
.input:disabled, .input:disabled,
.input[disabled]{ .input[disabled]{
@@ -2248,6 +2469,10 @@ details.collapse summary::-webkit-details-marker{
z-index: 200; z-index: 200;
} }
.z-\[999\]{
z-index: 999;
}
.z-\[99\]{ .z-\[99\]{
z-index: 99; z-index: 99;
} }
@@ -2314,6 +2539,10 @@ details.collapse summary::-webkit-details-marker{
height: 24rem; height: 24rem;
} }
.h-\[360px\]{
height: 360px;
}
.h-\[394px\]{ .h-\[394px\]{
height: 394px; height: 394px;
} }
@@ -2322,35 +2551,6 @@ details.collapse summary::-webkit-details-marker{
height: 100%; height: 100%;
} }
.h-fit{
height: -moz-fit-content;
height: fit-content;
}
.h-\[400px\]{
height: 400px;
}
.h-\[420px\]{
height: 420px;
}
.h-\[410px\]{
height: 410px;
}
.h-\[360px\]{
height: 360px;
}
.min-h-\[400px\]{
min-height: 400px;
}
.min-h-\[380px\]{
min-height: 380px;
}
.min-h-\[350px\]{ .min-h-\[350px\]{
min-height: 350px; min-height: 350px;
} }
@@ -2379,6 +2579,10 @@ details.collapse summary::-webkit-details-marker{
width: 32rem; width: 32rem;
} }
.w-\[360px\]{
width: 360px;
}
.w-fit{ .w-fit{
width: -moz-fit-content; width: -moz-fit-content;
width: fit-content; width: fit-content;
@@ -2388,38 +2592,6 @@ details.collapse summary::-webkit-details-marker{
width: 100%; width: 100%;
} }
.w-\[\]{
width: ;
}
.w-\[400px\]{
width: 400px;
}
.w-\[420px\]{
width: 420px;
}
.w-\[410px\]{
width: 410px;
}
.w-\[360px\]{
width: 360px;
}
.min-w-\[400px\]{
min-width: 400px;
}
.min-w-\[380\]{
min-width: 380;
}
.min-w-\[380px\]{
min-width: 380px;
}
.min-w-\[350px\]{ .min-w-\[350px\]{
min-width: 350px; min-width: 350px;
} }
@@ -2514,10 +2686,6 @@ details.collapse summary::-webkit-details-marker{
border-radius: 0.125rem; border-radius: 0.125rem;
} }
.rounded-2xl{
border-radius: 1rem;
}
.rounded-xl{ .rounded-xl{
border-radius: 0.75rem; border-radius: 0.75rem;
} }
@@ -2739,10 +2907,6 @@ details.collapse summary::-webkit-details-marker{
padding: 1rem; padding: 1rem;
} }
.p-24{
padding: 6rem;
}
.\!px-0{ .\!px-0{
padding-left: 0px !important; padding-left: 0px !important;
padding-right: 0px !important; padding-right: 0px !important;
@@ -2763,9 +2927,8 @@ details.collapse summary::-webkit-details-marker{
padding-bottom: 0.5rem; padding-bottom: 0.5rem;
} }
.py-16{ .pr-2{
padding-top: 4rem; padding-right: 0.5rem;
padding-bottom: 4rem;
} }
.text-start{ .text-start{
@@ -3092,11 +3255,21 @@ img[src] {
color: rgb(239 68 68 / var(--tw-text-opacity)); color: rgb(239 68 68 / var(--tw-text-opacity));
} }
.focus\:border-none:focus{
border-style: none;
}
.focus\:outline-none:focus{ .focus\:outline-none:focus{
outline: 2px solid transparent; outline: 2px solid transparent;
outline-offset: 2px; outline-offset: 2px;
} }
.focus\:ring-0:focus{
--tw-ring-offset-shadow: var(--tw-ring-inset) 0 0 0 var(--tw-ring-offset-width) var(--tw-ring-offset-color);
--tw-ring-shadow: var(--tw-ring-inset) 0 0 0 calc(0px + var(--tw-ring-offset-width)) var(--tw-ring-color);
box-shadow: var(--tw-ring-offset-shadow), var(--tw-ring-shadow), var(--tw-shadow, 0 0 #0000);
}
@media (min-width: 640px){ @media (min-width: 640px){
.sm\:flex{ .sm\:flex{
display: flex; display: flex;
+61 -74
View File
@@ -1,18 +1,19 @@
from aiohttp import web from aiohttp import web
from segment_anything import sam_model_registry, SamPredictor
from PIL import Image, ImageOps
from dotenv import load_dotenv from dotenv import load_dotenv
from blender.mesh_utils import upload_avatar_file from blender.mesh_utils import upload_avatar_file
from sam_utils import (
sam_ckpt_to_type,
compute_image_embedding,
check_embedding_exists,
save_embedding,
load_image,
)
import os import os
import requests import requests
import folder_paths import folder_paths
import json import json
import numpy as np
import server import server
import re
import base64 import base64
from PIL import Image
import io
import time import time
import execution import execution
import random import random
@@ -67,80 +68,19 @@ async def get_sam_model(request):
return web.FileResponse(filename) return web.FileResponse(filename)
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")
image = np.array(image).astype(np.float32) / 255.0
return image
@server.PromptServer.instance.routes.post("/sam_model") @server.PromptServer.instance.routes.post("/sam_model")
async def post_sam_model(request): async def post_sam_model(request):
post = await request.json() post = await request.json()
is_generated_image = post.get("isGeneratedImage") is_generated_image = post.get("isGeneratedImage")
emb_id = post.get("embedding_id") emb_id = post.get("embedding_id")
ckpt = post.get("ckpt") ckpt = post.get("ckpt")
ckpt = folder_paths.get_full_path("sams", ckpt) model_type = sam_ckpt_to_type[ckpt]
remote = post.get("remote") if not check_embedding_exists(emb_id, model_type):
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"), is_generated_image) image = load_image(post.get("image"), is_generated_image)
if remote: emb, img_model_input_size, img_original_size = compute_image_embedding(
# Run embed in remote server image, model_type
image = Image.fromarray((image * 255).astype(np.uint8)) )
buffered = io.BytesIO() save_embedding(emb_id, model_type, emb, img_model_input_size, img_original_size)
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",
},
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)
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") print("Finished embedding")
return web.json_response({}) return web.json_response({})
@@ -222,7 +162,6 @@ def load_workflow(workflow_name):
return "\n".join(f.readlines()) return "\n".join(f.readlines())
@server.PromptServer.instance.routes.post("/avatar_generation") @server.PromptServer.instance.routes.post("/avatar_generation")
async def post_prompt_block(request): async def post_prompt_block(request):
prompt_server = server.PromptServer.instance prompt_server = server.PromptServer.instance
@@ -286,6 +225,54 @@ async def post_prompt_block(request):
time.sleep(0.5) time.sleep(0.5)
# TODO: refactor the code
@server.PromptServer.instance.routes.post("/rendering_generation")
async def post_data_generation(request):
prompt_server = server.PromptServer.instance
post = await request.json()
workflow_name = post.get("workflow_name")
workflow = load_workflow(workflow_name)
workflow = workflow.replace("SEED", str(randomSeed()))
inputs = post.get("inputs")
for key, value in inputs.items():
workflow = workflow.replace(f'"{key}"', f'"{str(value)}"')
res = post_prompt({"prompt": json.loads(workflow)})
prompt_id = json.loads(res.text)["prompt_id"]
while True:
history = prompt_server.prompt_queue.get_history(prompt_id=prompt_id)
if history:
outputs = history[prompt_id]["outputs"]
for node_id, output in outputs.items():
if "images" in output:
filename = output["images"][0]["filename"]
if filename.startswith("rendered"):
return web.json_response({"image": filename}, status=200)
time.sleep(0.5)
@server.PromptServer.instance.routes.post("/image_generation")
async def post_image_generation(request):
prompt_server = server.PromptServer.instance
workflow = load_workflow("generation")
workflow = workflow.replace("SEED", str(randomSeed()))
res = post_prompt({"prompt": json.loads(workflow)})
prompt_id = json.loads(res.text)["prompt_id"]
while True:
history = prompt_server.prompt_queue.get_history(prompt_id=prompt_id)
if history:
outputs = history[prompt_id]["outputs"]
for node_id, output in outputs.items():
if "images" in output:
filename = output["images"][0]["filename"]
if filename.startswith("avatar"):
return web.json_response({"image": filename}, status=200)
time.sleep(0.5)
# @server.PromptServer.instance.routes.get("/get_default_workflow") # @server.PromptServer.instance.routes.get("/get_default_workflow")
# async def get_default_workflow(request): # 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/1119102674437156984/1172255632586448987/workflow_boy_2_1.json?ex=655fa722&is=654d3222&hm=463fa6a3c6ea60f7471196ff45382c729d3b856e86282f905d37a0398711860e&" # YP workflow
+27
View File
@@ -0,0 +1,27 @@
import json
class CombinePoints:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
},
}
RETURN_NAMES = ("SAM_PROMPTS",)
RETURN_TYPES = ("STRING",)
FUNCTION = "run"
CATEGORY = "image"
# OUTPUT_NODE = True
def run(self, *args, **kwargs):
sam_prompts = json.dumps(kwargs, default=str)
return (sam_prompts,)
NODE_CLASS_MAPPINGS = {"Combine Points": CombinePoints}
NODE_DISPLAY_NAME_MAPPINGS = {"Combine Points": "Combine Points"}
+66
View File
@@ -0,0 +1,66 @@
import cv2
import numpy as np
import torch
class ExtractBoundaryPoints:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE",),
"n_points": ("INT", {"default": -1, "min": -1, "max": 100}),
},
}
RETURN_TYPES = ("POINTS", "IMAGE")
FUNCTION = "run"
CATEGORY = "image"
def find_main_contour(self, image, n_points):
image = np.copy(image[0].numpy())
gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)
gray = (gray * 255).astype(np.uint8)
# Find contours
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]
max_area_index = areas.index(max(areas))
largest_contour = contours[max_area_index]
if n_points > 0:
divided_by = int(largest_contour.shape[0] / n_points)
divided_by = min(divided_by, largest_contour.shape[0])
largest_contour = largest_contour[::divided_by].astype(int)
contours = [largest_contour]
if not image.flags["C_CONTIGUOUS"]:
image = np.ascontiguousarray(image)
cv2.drawContours(image, contours, -1, (0, 255, 0), 3)
points = []
for point in largest_contour:
points.append(
{"x": point[0][0], "y": point[0][1], "label": 1, "isAuto": True}
)
return image, points
def run(self, image, n_points):
contour_image, points = self.find_main_contour(image, n_points)
contour_image = torch.from_numpy(np.expand_dims(contour_image, axis=0))
print(points)
return (points, contour_image)
NODE_CLASS_MAPPINGS = {"Extract Boundary Points": ExtractBoundaryPoints}
NODE_DISPLAY_NAME_MAPPINGS = {"Extract Boundary Points": "Extract Boundary Points"}
+31
View File
@@ -0,0 +1,31 @@
class LoadValueFromRequest:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"name": (
"STRING",
{"multiline": False, "default": "key_name"},
),
},
"optional": {
"value": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01, "display": "number"}),
}
}
RETURN_TYPES = ("FLOAT",)
RETURN_NAMES = ("value",)
FUNCTION = "run"
CATEGORY = "image"
def run(self, name, value=None):
if name:
value = name
return (value,)
NODE_CLASS_MAPPINGS = {"LoadValueFromRequest": LoadValueFromRequest}
NODE_DISPLAY_NAME_MAPPINGS = {"LoadValueFromRequest": "Load Value From Request"}
+254
View File
@@ -0,0 +1,254 @@
# For auto-segmentation
import mediapipe as mp
import numpy as np
import os
from math import sqrt
face_landmarker = None
pose_landmarker = None
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
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]],
},
}
def load_mediapipe_models():
global face_landmarker, pose_landmarker
if face_landmarker is None and pose_landmarker is None:
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 face_landmarker, pose_landmarker
def auto_segment_face(image, face_landmarks):
H, W, C = image.shape
layer_points = {}
layer_bboxes = {}
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"]:
layer_points[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,
}
layer_points[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),
]
)
layer_bboxes[key] = box
return layer_points, layer_bboxes
def auto_segment_pose(image, pose_landmarks):
H, W, C = image.shape
layer_points = {}
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
layer_points["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 layer_points
def detect_face(np_image):
face_landmarker, pose_landmarker = load_mediapipe_models()
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
if len(face_landmarks) > 0:
face_points, face_bboxes = auto_segment_face(np_image, face_landmarks[0])
else:
face_points, face_bboxes = {}, {}
print("Warning: no face detected")
pose_landmarks = pose_landmarker.detect(mp_image).pose_landmarks
if len(pose_landmarks) > 0:
pose_points = auto_segment_pose(np_image, pose_landmarks[0])
else:
pose_points = {}
print("Warning: no pose detected")
layer_points = {**face_points, **pose_points}
layer_bboxes = {**face_bboxes}
return layer_points, layer_bboxes
+73 -318
View File
@@ -4,105 +4,20 @@ import numpy as np
import torch import torch
import re import re
import json import json
from segment_anything import sam_model_registry, SamPredictor import uuid
from sam_utils import (
load_model,
check_embedding_exists,
compute_image_embedding,
save_embedding,
load_embdding,
load_image,
)
from mediapipe_utils import detect_face
from einops import rearrange, repeat 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: class SAMMultiLayer:
def __init__(self):
self.predictor = None
self.output_dir = folder_paths.get_output_directory()
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
return { return {
@@ -119,183 +34,31 @@ class SAMMultiLayer:
CATEGORY = "image" CATEGORY = "image"
RETURN_TYPES = ("SAM_PROMPT",) RETURN_TYPES = ["SAM_PROMPT"] # + ["IMAGE"] * 100
FUNCTION = "load_image" FUNCTION = "run"
def load_models(self, ckpt, model_type): def run(self, image, ckpt, embedding_id, image_prompts_json):
global global_predictor, face_landmarker, pose_landmarker if (
"COMFY_DEPLOY" in os.environ
and os.getenv("COMFY_DEPLOY", "FALSE") == "TRUE"
):
embedding_id = str(uuid.uuid4())
layer_points = json.loads(image_prompts_json.replace("'", '"'))
ckpt = folder_paths.get_full_path("sams", ckpt) order_file = (
sam = sam_model_registry[model_type](checkpoint=ckpt) # .to("cuda") f"{folder_paths.get_output_directory()}/segments_{embedding_id}/order.json"
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):
image_prompts = json.loads(image_prompts_json.replace("'", '"'))
order_file = f"{self.output_dir}/segments_{embedding_id}/order.json"
if os.path.exists(order_file): if os.path.exists(order_file):
# Frontend uploads segments images to backend => backend reads all segments images and passes them to next nodes # Frontend uploads segments images to backend => backend reads all segments images and passes them to next nodes
with open(order_file) as f: with open(order_file) as f:
order = json.load(f) order = json.load(f)
result = [image_prompts] result = [layer_points]
for segment in order: for segment in order:
image = Image.open( image = load_image(
f"{self.output_dir}/segments_{embedding_id}/{segment}.png" f"{folder_paths.get_output_directory()}/segments_{embedding_id}/{segment}.png",
comfyui_format=True,
) )
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
result.append(image) result.append(image)
return result return result
@@ -303,68 +66,60 @@ class SAMMultiLayer:
# Frontend uploads clicks coordinates to backend => backend runs SAM and passes the segments to next nodes # 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] model_type = re.findall(r"vit_[lbh]", ckpt)[0]
global global_predictor # get first image from batch
if global_predictor is None: image = image[0]
global_predictor, _, _ = self.load_models(ckpt, model_type) if image.shape[2] == 4:
# to RGB
image = image[:, :, :3]
if image.shape[3] == 4: if not check_embedding_exists(embedding_id, model_type):
image = image[:, :, :, :3] emb, img_model_input_size, img_original_size = compute_image_embedding(
image, model_type
emb_filename = f"{self.output_dir}/{embedding_id}_{model_type}.npy" )
if not os.path.exists(emb_filename): save_embedding(
image_np = (image[0].numpy() * 255).astype(np.uint8) embedding_id,
global_predictor.set_image(image_np) model_type,
emb = global_predictor.get_image_embedding().cpu().numpy() emb,
np.save(emb_filename, emb) img_model_input_size,
img_original_size,
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: else:
emb = np.load(emb_filename) load_embdding(embedding_id, model_type)
with open(f"{self.output_dir}/{embedding_id}_{model_type}.json") as f: detected_points, detected_bboxes = detect_face(image.numpy())
data = json.load(f) result = [layer_points]
global_predictor.input_size = data["input_size"] for layer, points in layer_points.items():
global_predictor.features = torch.from_numpy(emb) if detected_points is not None and layer in detected_points:
global_predictor.is_image_set = True # use detected points by mediapipe
global_predictor.original_size = data["original_size"] points = detected_points[layer]
imagePromptsMulti, boxesMulti = self.detect_face(image[0].numpy()) if len(points) == 0:
# no points, append a black image
h, w, c = image.shape
result.append(torch.zeros(1, h, w, c))
print("No points for layer", layer)
continue
image_prompts = json.loads(image_prompts_json.replace("'", '"')) # prepare for SAM inferencing
result = [image_prompts] point_coords = np.array([[p["x"], p["y"]] for p in points])
point_labels = np.array([p["label"] for p in points])
bbox = (
detected_bboxes[layer]
if detected_bboxes is not None and layer in detected_bboxes
else None
)
if isinstance(image_prompts, list): sam_predictor = load_model(model_type)["predictor"]
pass masks, _, _ = sam_predictor.predict(
elif all(isinstance(item, list) for item in image_prompts.values()): point_coords=point_coords,
for key, item in image_prompts.items(): point_labels=point_labels,
if len(item) == 0: box=bbox,
h, w, c = image[0].shape )
result.append(torch.zeros(1, h, w, c)) masks = torch.from_numpy(masks)
continue masks = rearrange(masks[0], "h w -> 1 h w")
out_image = repeat(masks, "1 h w -> 1 h w c", c=3) * image.unsqueeze(0)
points = ( result.append(out_image)
imagePromptsMulti[key] if key in imagePromptsMulti else item return result
)
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} NODE_CLASS_MAPPINGS = {"SAM MultiLayer": SAMMultiLayer}
+97
View File
@@ -0,0 +1,97 @@
from segment_anything import sam_model_registry, SamPredictor
from PIL import Image, ImageOps
import folder_paths
import numpy as np
import torch
import os
import json
sam_type_to_ckpt = {
"vit_h": "sam_vit_h_4b8939.pth",
"vit_l": "sam_vit_l_0b3195.pth",
"vit_b": "sam_vit_b_01ec64.pth",
}
sam_ckpt_to_type = {v: k for k, v in sam_type_to_ckpt.items()}
sam_instance = {"model_type": None, "model": None, "predictor": None}
def load_model(model_type):
global sam_instance
if sam_instance["model"] is None or sam_instance["model_type"] != model_type:
ckpt = sam_type_to_ckpt[model_type]
ckpt = folder_paths.get_full_path("sams", ckpt)
sam_instance["model_type"] = model_type
sam_instance["model"] = sam_model_registry[model_type](checkpoint=ckpt)
if torch.cuda.is_available():
sam_instance["model"].cuda()
sam_instance["predictor"] = SamPredictor(sam_instance["model"])
return sam_instance
def check_embedding_exists(emb_id, model_type):
emb_filename = f"{folder_paths.get_output_directory()}/{emb_id}_{model_type}.npy"
return os.path.exists(emb_filename)
def save_embedding(emb_id, model_type, emb, img_input_size, img_original_size):
emb_filename = f"{folder_paths.get_output_directory()}/{emb_id}_{model_type}.npy"
np.save(emb_filename, emb)
json_filename = f"{folder_paths.get_output_directory()}/{emb_id}_{model_type}.json"
with open(json_filename, "w") as f:
json.dump(
{
"input_size": img_input_size,
"original_size": img_original_size,
},
f,
)
def load_embdding(emb_id, model_type):
emb_filename = f"{folder_paths.get_output_directory()}/{emb_id}_{model_type}.npy"
emb = np.load(emb_filename)
json_filename = f"{folder_paths.get_output_directory()}/{emb_id}_{model_type}.json"
with open(json_filename, "r") as f:
sizes = json.load(f)
predictor = load_model(model_type)["predictor"]
predictor.input_size = sizes["input_size"]
predictor.features = torch.from_numpy(emb)
predictor.is_image_set = True
predictor.original_size = sizes["original_size"]
@torch.no_grad()
def compute_image_embedding(image, model_type="vit_h"):
sam = load_model(model_type)
predictor = sam["predictor"]
# if image.shape[3] == 4:
# image = image[:, :, :, :3]
if torch.is_tensor(image):
image = image.numpy()
image_np = (image * 255).astype(np.uint8)
predictor.set_image(image_np)
emb = predictor.get_image_embedding().cpu().numpy()
return emb, predictor.input_size, predictor.original_size
def load_image(image, is_generated_image=False, comfyui_format=False):
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")
image = np.array(image).astype(np.float32) / 255.0
if comfyui_format:
# to torch and create batch dimension
image = torch.from_numpy(image)[None,]
return image