Update HumanSegmentation (#826)

* Change to new mask components on humanSegmentation

* Change segformer index

* Add segformer_b3_clothes and fashion

* Add face_parsing

* Fix face_parsing output error images

---------

Co-authored-by: yolain <me@yolain.com>
This commit is contained in:
yolain
2025-07-10 18:32:06 +08:00
committed by GitHub
co-authored by yolain
parent 54614079ca
commit 560be6aee7
16 changed files with 163 additions and 89 deletions
+2 -2
View File
@@ -6620,7 +6620,7 @@
"name": "温度"
},
"max_tokens": {
"name": "最大词令牌数"
"name": "最大词元数"
},
"caption_type": {
"name": "提示词类型"
@@ -6651,7 +6651,7 @@
"name": "温度"
},
"max_tokens": {
"name": "最大词令牌数"
"name": "最大词元数"
},
"caption_type": {
"name": "提示词类型"
+9
View File
@@ -370,6 +370,15 @@ HUMANPARSING_MODELS = {
},
"human-parts":{
"model_url":"https://huggingface.co/Metal3d/deeplabv3p-resnet50-human/resolve/main/deeplabv3p-resnet50-human.onnx",
},
"segformer_b3_clothes":{
"model_name": "sayeed99/segformer-b3-clothes",
},
"segformer_b3_fashion":{
"model_name": "sayeed99/segformer-b3-fashion",
},
"face_parsing":{
"model_name": "jonathandinu/face-parsing"
}
}
+115 -50
View File
@@ -4,12 +4,13 @@ import torch
import numpy as np
import comfy.utils
import comfy.model_management
import shutil
from comfy_extras.nodes_compositing import JoinImageWithAlpha
from server import PromptServer
from nodes import MAX_RESOLUTION, NODE_CLASS_MAPPINGS as ALL_NODE_CLASS_MAPPINGS
from PIL import Image, ImageDraw, ImageFilter, ImageOps
import torch.nn.functional as F
from torchvision.transforms import Resize, CenterCrop, GaussianBlur
from torchvision.transforms import Resize, CenterCrop, GaussianBlur, ToPILImage
from torchvision.transforms.functional import to_pil_image
from ..libs.log import log_node_info
from ..libs.utils import AlwaysEqualProxy, ByPassTypeTuple
@@ -1277,13 +1278,22 @@ class humanSegmentation:
@classmethod
def INPUT_TYPES(cls):
return {
"required":{
"image": ("IMAGE",),
"method": (["selfie_multiclass_256x256", "human_parsing_lip", "human_parts (deeplabv3p)"],),
"method": (["selfie_multiclass_256x256", "human_parsing_lip", "human_parts (deeplabv3p)", "segformer_b3_clothes", "segformer_b3_fashion", "face_parsing"],),
"confidence": ("FLOAT", {"default": 0.4, "min": 0.05, "max": 0.95, "step": 0.01},),
"crop_multi": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001},),
"mask_components":(
"EASY_COMBO",{
"options": [{'label':'Background','value':0}],
"multi_select": {
"placeholder": "select mask components",
"chip": True,
"max_selected_labels": 4,
}
}
)
},
"hidden": {
"prompt": "PROMPT",
@@ -1309,12 +1319,7 @@ class humanSegmentation:
numpy_image = cv2.cvtColor(numpy_image, cv2.COLOR_BGR2RGB)
return mp.Image(image_format=image_format, data=numpy_image)
def parsing(self, image, confidence, method, crop_multi, prompt=None, my_unique_id=None):
mask_components = []
if my_unique_id in prompt:
if prompt[my_unique_id]["inputs"]['mask_components']:
mask_components = prompt[my_unique_id]["inputs"]['mask_components'].split(',')
mask_components = list(map(int, mask_components))
def parsing(self, image, confidence, method, crop_multi, mask_components, prompt=None, my_unique_id=None):
if method == 'selfie_multiclass_256x256':
try:
import mediapipe as mp
@@ -1442,6 +1447,107 @@ class humanSegmentation:
output_image = torch.cat(ret_images, dim=0)
mask = torch.cat(ret_masks, dim=0)
elif method in ["segformer_b3_clothes", "segformer_b3_fashion", "face_parsing"]:
from transformers import SegformerImageProcessor, AutoModelForSemanticSegmentation
# 分割
def get_segmentation_from_model(tensor_image, model, processor):
cloth = tensor2pil(tensor_image)
inputs = processor(images=cloth, return_tensors="pt")
outputs = model(**inputs)
logits = outputs.logits.cpu()
upsampled_logits = F.interpolate(logits, size=cloth.size[::-1], mode="bilinear",
align_corners=False)
pred_seg = upsampled_logits.argmax(dim=1)[0].numpy()
return pred_seg, cloth
if method in cache:
_, (processor, model) = cache[method][1]
else:
model_folder_path = os.path.join(folder_paths.models_dir, method)
if os.path.exists(model_folder_path):
print(f"Start to load existing model...")
else:
from huggingface_hub import snapshot_download
PromptServer.instance.send_sync("easyuse-toast", {"content": f"Model not found locally. Downloading {method}...", "type": 'loading', "duration": 10000})
print(f"Model not found locally. Downloading {method}...")
model_path_cache = os.path.join(folder_paths.models_dir, "cache-"+method)
snapshot_download(
repo_id=HUMANPARSING_MODELS[method]['model_name'],
local_dir=model_path_cache,
local_dir_use_symlinks=False,
resume_download=True
)
shutil.move(model_path_cache, model_folder_path)
print(f"Model downloaded to {model_folder_path}...")
try:
model_folder_path = os.path.normpath(folder_paths.folder_names_and_paths[method][0][0])
except:
pass
processor = SegformerImageProcessor.from_pretrained(model_folder_path)
model = AutoModelForSemanticSegmentation.from_pretrained(model_folder_path)
update_cache(method, 'human_segmentation', (False, (processor, model)))
ret_images = []
ret_masks = []
if method == "face_parsing":
import matplotlib
import torchvision.transforms as T
transform = ToPILImage()
colormap = matplotlib.colormaps['viridis']
device = model.device
results = []
images = []
for img in image:
size = img.shape[:2]
inputs = processor(images=transform(img.permute(2, 0, 1)), return_tensors="pt")
inputs = {k: v.to(device) for k, v in inputs.items()}
outputs = model(**inputs)
logits = outputs.logits
upsampled_logits = F.interpolate(
logits,
size=size,
mode="bilinear",
align_corners=False)
pred_seg = upsampled_logits.argmax(dim=1)[0]
pred_seg_np = pred_seg.cpu().detach().numpy().astype(np.uint8)
results.append(torch.tensor(pred_seg_np))
results_out = torch.stack(results, dim=0)
for img, result_item in zip(image, results_out):
mask = torch.zeros(result_item.shape, dtype=torch.uint8)
for i in mask_components:
mask = mask | torch.where(result_item == i, 1, 0)
# 将mask转换为numpy数组,并确保数据类型正确
mask_np = (mask * 255).numpy().astype(np.uint8)
_mask = Image.fromarray(mask_np)
# 处理图像输出
ret_image = RGB2RGBA(tensor2pil(img).convert('RGB'), _mask.convert('L'))
ret_images.append(pil2tensor(ret_image))
ret_masks.append(image2mask(_mask))
else:
for img in image:
pred_seg, cloth = get_segmentation_from_model(img, model, processor)
i = torch.unsqueeze(img, 0)
i = pil2tensor(tensor2pil(i).convert('RGB'))
mask = np.isin(pred_seg, mask_components).astype(np.uint8)
_mask = Image.fromarray(mask * 255)
ret_image = RGB2RGBA(tensor2pil(img).convert('RGB'), _mask.convert('L'))
ret_images.append(pil2tensor(ret_image))
ret_masks.append(image2mask(_mask))
output_image = torch.cat(ret_images, dim=0)
mask = torch.cat(ret_masks, dim=0)
# use crop
bbox = [[0, 0, 0, 0]]
if crop_multi > 0.0:
@@ -1892,47 +1998,6 @@ class loadImagesForLoop:
"result": tuple(["stub", index, image, mask, name] + outputs),
"expand": graph.finalize(),
}
# 姿势编辑器
# class poseEditor:
# @classmethod
# def INPUT_TYPES(self):
# temp_dir = folder_paths.get_temp_directory()
#
# if not os.path.isdir(temp_dir):
# os.makedirs(temp_dir)
#
# temp_dir = folder_paths.get_temp_directory()
#
# return {"required":
# {"image": (sorted(os.listdir(temp_dir)),)},
# }
#
# RETURN_TYPES = ("IMAGE",)
# FUNCTION = "output_pose"
#
# CATEGORY = "EasyUse/🚫 Deprecated"
#
# def output_pose(self, image):
# image_path = os.path.join(folder_paths.get_temp_directory(), image)
# # print(f"Create: {image_path}")
#
# i = Image.open(image_path)
# image = i.convert("RGB")
# image = np.array(image).astype(np.float32) / 255.0
# image = torch.from_numpy(image)[None,]
#
# return (image,)
#
# @classmethod
# def IS_CHANGED(self, image):
# image_path = os.path.join(
# folder_paths.get_temp_directory(), image)
# # print(f'Change: {image_path}')
#
# m = hashlib.sha256()
# with open(image_path, 'rb') as f:
# m.update(f.read())
# return m.digest().hex()
class makeImageForICRepaint:
@classmethod
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
import{x as e,y as t,j as n,o,n as i,i as l,C as r,R as u,e as a,S as s}from"./vue-CYk3Qk0T.js";function c(n){return!!e()&&(t(n),!0)}const v="undefined"!=typeof window&&"undefined"!=typeof document;"undefined"!=typeof WorkerGlobalScope&&(globalThis,WorkerGlobalScope);const f=e=>null!=e,d=Object.prototype.toString,p=e=>"[object Object]"===d.call(e);function m(e){return Array.isArray(e)?e:[e]}function h(e,t=!0,n){l()?o(e,n):t?e():i(e)}const w=v?window:void 0;function y(e){var t;const n=u(e);return null!=(t=null==n?void 0:n.$el)?t:n}function b(...e){const t=[],o=()=>{t.forEach((e=>e())),t.length=0},i=r((()=>{const t=m(u(e[0])).filter((e=>null!=e));return t.every((e=>"string"!=typeof e))?t:void 0})),l=(s=([e,n,i,l])=>{if(o(),!(null==e?void 0:e.length)||!(null==n?void 0:n.length)||!(null==i?void 0:i.length))return;const r=p(l)?{...l}:l;t.push(...e.flatMap((e=>n.flatMap((t=>i.map((n=>((e,t,n,o)=>(e.addEventListener(t,n,o),()=>e.removeEventListener(t,n,o)))(e,t,n,r))))))))},v={flush:"post"},n((()=>{var t,n;return[null!=(n=null==(t=i.value)?void 0:t.map((e=>y(e))))?n:[w].filter((e=>null!=e)),m(u(i.value?e[1]:e[0])),m(a(i.value?e[2]:e[1])),u(i.value?e[3]:e[2])]}),s,{...v,immediate:!0}));var s,v;return c(o),()=>{l(),o()}}function g(e){const t=function(){const e=s(!1),t=l();return t&&o((()=>{e.value=!0}),t),e}();return r((()=>(t.value,Boolean(e()))))}function O(e,t={}){const{reset:o=!0,windowResize:i=!0,windowScroll:l=!0,immediate:a=!0,updateTiming:v="sync"}=t,d=s(0),p=s(0),O=s(0),S=s(0),j=s(0),x=s(0),z=s(0),A=s(0);function R(){const t=y(e);if(!t)return void(o&&(d.value=0,p.value=0,O.value=0,S.value=0,j.value=0,x.value=0,z.value=0,A.value=0));const n=t.getBoundingClientRect();d.value=n.height,p.value=n.bottom,O.value=n.left,S.value=n.right,j.value=n.top,x.value=n.width,z.value=n.x,A.value=n.y}function E(){"sync"===v?R():"next-frame"===v&&requestAnimationFrame((()=>R()))}return function(e,t,o={}){const{window:i=w,...l}=o;let a;const s=g((()=>i&&"ResizeObserver"in i)),v=()=>{a&&(a.disconnect(),a=void 0)},f=r((()=>{const t=u(e);return Array.isArray(t)?t.map((e=>y(e))):[y(t)]})),d=n(f,(e=>{if(v(),s.value&&i){a=new ResizeObserver(t);for(const t of e)t&&a.observe(t,l)}}),{immediate:!0,flush:"post"}),p=()=>{v(),d()};c(p)}(e,E),n((()=>y(e)),(e=>!e&&E())),function(e,t,o={}){const{window:i=w,...l}=o;let a;const s=g((()=>i&&"MutationObserver"in i)),v=()=>{a&&(a.disconnect(),a=void 0)},d=r((()=>{const t=m(u(e)).map(y).filter(f);return new Set(t)})),p=n((()=>d.value),(e=>{v(),s.value&&e.size&&(a=new MutationObserver(t),e.forEach((e=>a.observe(e,l))))}),{immediate:!0,flush:"post"}),h=()=>{p(),v()};c(h)}(e,E,{attributeFilter:["style","class"]}),l&&b("scroll",E,{capture:!0,passive:!0}),i&&b("resize",E,{passive:!0}),h((()=>{a&&E()})),{height:d,bottom:p,left:O,right:S,top:j,width:x,x:z,y:A,update:E}}export{b as a,O as u};
-1
View File
@@ -1 +0,0 @@
import{v as e,x as t,i as n,o,n as i,h as l,B as r,Q as u,e as a,R as s}from"./vue-BPkYw06T.js";function v(n){return!!e()&&(t(n),!0)}const c="undefined"!=typeof window&&"undefined"!=typeof document;"undefined"!=typeof WorkerGlobalScope&&(globalThis,WorkerGlobalScope);const f=e=>null!=e,d=Object.prototype.toString,p=e=>"[object Object]"===d.call(e);function m(e){return Array.isArray(e)?e:[e]}function h(e,t=!0,n){l()?o(e,n):t?e():i(e)}const w=c?window:void 0;function b(e){var t;const n=u(e);return null!=(t=null==n?void 0:n.$el)?t:n}function y(...e){const t=[],o=()=>{t.forEach((e=>e())),t.length=0},i=r((()=>{const t=m(u(e[0])).filter((e=>null!=e));return t.every((e=>"string"!=typeof e))?t:void 0})),l=(s=([e,n,i,l])=>{if(o(),!(null==e?void 0:e.length)||!(null==n?void 0:n.length)||!(null==i?void 0:i.length))return;const r=p(l)?{...l}:l;t.push(...e.flatMap((e=>n.flatMap((t=>i.map((n=>((e,t,n,o)=>(e.addEventListener(t,n,o),()=>e.removeEventListener(t,n,o)))(e,t,n,r))))))))},c={flush:"post"},n((()=>{var t,n;return[null!=(n=null==(t=i.value)?void 0:t.map((e=>b(e))))?n:[w].filter((e=>null!=e)),m(u(i.value?e[1]:e[0])),m(a(i.value?e[2]:e[1])),u(i.value?e[3]:e[2])]}),s,{...c,immediate:!0}));var s,c;return v(o),()=>{l(),o()}}function g(e){const t=function(){const e=s(!1),t=l();return t&&o((()=>{e.value=!0}),t),e}();return r((()=>(t.value,Boolean(e()))))}function O(e,t={}){const{reset:o=!0,windowResize:i=!0,windowScroll:l=!0,immediate:a=!0,updateTiming:c="sync"}=t,d=s(0),p=s(0),O=s(0),x=s(0),z=s(0),A=s(0),R=s(0),S=s(0);function j(){const t=b(e);if(!t)return void(o&&(d.value=0,p.value=0,O.value=0,x.value=0,z.value=0,A.value=0,R.value=0,S.value=0));const n=t.getBoundingClientRect();d.value=n.height,p.value=n.bottom,O.value=n.left,x.value=n.right,z.value=n.top,A.value=n.width,R.value=n.x,S.value=n.y}function E(){"sync"===c?j():"next-frame"===c&&requestAnimationFrame((()=>j()))}return function(e,t,o={}){const{window:i=w,...l}=o;let a;const s=g((()=>i&&"ResizeObserver"in i)),c=()=>{a&&(a.disconnect(),a=void 0)},f=r((()=>{const t=u(e);return Array.isArray(t)?t.map((e=>b(e))):[b(t)]})),d=n(f,(e=>{if(c(),s.value&&i){a=new ResizeObserver(t);for(const t of e)t&&a.observe(t,l)}}),{immediate:!0,flush:"post"}),p=()=>{c(),d()};v(p)}(e,E),n((()=>b(e)),(e=>!e&&E())),function(e,t,o={}){const{window:i=w,...l}=o;let a;const s=g((()=>i&&"MutationObserver"in i)),c=()=>{a&&(a.disconnect(),a=void 0)},d=r((()=>{const t=m(u(e)).map(b).filter(f);return new Set(t)})),p=n((()=>d.value),(e=>{c(),s.value&&e.size&&(a=new MutationObserver(t),e.forEach((e=>a.observe(e,l))))}),{immediate:!0,flush:"post"}),h=()=>{p(),c()};v(h)}(e,E,{attributeFilter:["style","class"]}),l&&y("scroll",E,{capture:!0,passive:!0}),i&&y("resize",E,{passive:!0}),h((()=>{a&&E()})),{height:d,bottom:p,left:O,right:x,top:z,width:A,x:R,y:S,update:E}}export{y as a,O as u};
File diff suppressed because one or more lines are too long