Compare commits

...
57 Commits
Author SHA1 Message Date
shadowcz007 5564ee1246 rembgNode update
"briarmbg","u2net","u2netp","u2net_human_seg","u2net_cloth_seg","silueta","isnet-general-use","isnet-anime"
2024-02-08 14:01:36 +08:00
shadowcz007 a6e9251521 add briarmbg to rembgNode 2024-02-08 13:57:01 +08:00
shadowcz007 0bee093916 Update README.md 2024-02-08 11:57:10 +08:00
shadowcz007 13a9878823 Update README.md 2024-02-08 11:56:41 +08:00
shadowcz007 acd35d50f8 Update ChatGPT.py 2024-02-08 11:53:03 +08:00
shadowcz007 5c0d99e72d comfyui-CLIPSeg 2024-02-08 10:16:49 +08:00
shadowcz007 e37af93be3 0.15.1 2024-02-06 23:06:33 +08:00
shadowcz007 244c1700e1 Update ImageNode.py 2024-02-06 22:22:56 +08:00
shadowcz007 a3a15473ba LoadImage 的mask保存到appinfo 2024-02-06 19:01:41 +08:00
shadowcz007 d734b5077c GetImageSize_ 增加最小尺寸 2024-02-05 22:44:04 +08:00
shadowcz007 5430072b19 Update app_mixlab.js 2024-02-05 22:25:25 +08:00
shadowcz007 465aebaed4 Update app_mixlab.js 2024-02-05 22:24:34 +08:00
shadowcz007 ee0b16c2ea 修复appinfo配置的bug 2024-02-05 17:56:25 +08:00
shadowcz007 21d5eacb41 Update index.html 2024-02-04 23:46:25 +08:00
shadowcz007 c06688eb0b comfyui-consistency-decoder 2024-02-02 09:47:51 +08:00
shadowcz007 b74bbcd279 fixbug :SaveImageToLocal 2024-02-01 23:52:23 +08:00
shadowcz007 64d8d9b05d Delete echarts.min.js 2024-02-01 00:48:25 +08:00
shadowcz007 6ab60f281b Update ImageNode.py 2024-01-31 00:24:23 +08:00
shadowcz007 e35be3b2fa Update README.md 2024-01-30 22:57:13 +08:00
shadowcz007 0a0c27ac96 fixbug:Save Group as Template 2024-01-30 22:54:39 +08:00
shadowcz007 fb249e84eb SplitImage增加mask输出 2024-01-30 18:07:21 +08:00
shadowcz007 76ad86fcae Update __init__.py 2024-01-29 23:18:16 +08:00
shadowcz007 037614d227 showText 可以保存txt到本地目录 2024-01-29 14:37:51 +08:00
shadowcz007 4a50e445fd 从本地读取文件-输出文件名 2024-01-29 13:46:02 +08:00
shadowcz007 6f3c1c4393 update 2024-01-29 11:55:10 +08:00
shadowcz007 c6a9b4b592 Update gpt_mixlab.js 2024-01-29 11:45:53 +08:00
shadowcz007 e915ac4eca Update ChatGPT.py 2024-01-29 11:37:49 +08:00
shadowcz007 a857793f63 fixbug 2024-01-29 11:34:06 +08:00
shadowcz007 31914f7510 Update ChatGPT.py 2024-01-29 11:12:02 +08:00
shadowcz007 0be859f0ee fixbug 2024-01-28 21:38:13 +08:00
shadowcz007 14b9c3697b Update ImageNode.py 2024-01-28 20:38:22 +08:00
shadowcz007 96b66a57bb showText can save to local 2024-01-28 18:13:50 +08:00
shadowcz007 f13701c489 v0.15.0 2024-01-28 16:40:31 +08:00
shadow 31515b810e Merge pull request #159 from wfjsw/debloat-init-1
publish routes without having to replicate add_routes
2024-01-28 16:01:17 +08:00
shadowcz007 29e84e08a4 add CenterImage 2024-01-28 15:59:54 +08:00
shadowcz007 0ac9ad9757 修复批量保存本地图片的bug 2024-01-27 22:44:47 +08:00
shadowcz007 9a432e0608 Update Utils.py 2024-01-27 00:58:08 +08:00
shadowcz007 c83ba5fe7f Update PromptNode.py 2024-01-27 00:57:40 +08:00
shadowcz007 1b55c743ea 增加一些seed来控制节点 2024-01-27 00:33:49 +08:00
shadowcz007 a93579376c 不覆盖文件 2024-01-26 22:47:34 +08:00
shadowcz007 eba49f3c68 add SaveImageToLocal 2024-01-26 12:15:01 +08:00
shadowcz007 3a3da49c69 Update ImageNode.py 2024-01-25 20:10:29 +08:00
shadowcz007 3d68e48219 Update ImageNode.py 2024-01-25 20:07:31 +08:00
shadowcz007 4351fa6a0e 增加mask 的resize 2024-01-25 17:54:51 +08:00
shadowcz007 77222d2808 修复ImageCropByAlpha的bug 2024-01-25 15:40:00 +08:00
Jabasukuriputo Wang 36db7e5a9a publish routes without having to replicate add_routes 2024-01-25 00:44:51 -06:00
shadowcz007 3572368f16 fixbug 2024-01-25 10:38:23 +08:00
shadowcz007 0755dc1462 CreateLoraNames 2024-01-24 22:59:49 +08:00
shadowcz007 94f81b7102 fixbug 2024-01-24 17:45:15 +08:00
shadowcz007 08f8fe3d7e Update README.md 2024-01-23 23:44:47 +08:00
shadowcz007 a6cd383d67 add Sampler_names 2024-01-23 14:23:04 +08:00
shadowcz007 10face6ab0 CkptNames 2024-01-23 14:04:26 +08:00
shadowcz007 a6259ff600 add CkptNames 2024-01-23 13:58:59 +08:00
shadowcz007 d7d9e6cbfe add smart_connect_v1 2024-01-23 12:18:38 +08:00
shadowcz007 8263609470 优化LoadImageURL,增加seed,保证图片加载失败后可以继续 2024-01-23 10:03:55 +08:00
shadowcz007 fc2367de76 centerOnNode & fix node (widgets) 2024-01-21 20:31:57 +08:00
shadowcz007 9a4f2ebc70 Update ui_mixlab.js 2024-01-20 22:39:56 +08:00
21 changed files with 2462 additions and 659 deletions
+25 -15
View File
@@ -1,5 +1,15 @@
> 适配了最新版comfyui的py3.11 ,torch 2.1.2+cu121
> [discord](https://discord.gg/cXs9vZSqeK)
####
[comfyui-ultralytics-yolo](https://github.com/shadowcz007/comfyui-ultralytics-yolo)
[comfyui-moondream](https://github.com/shadowcz007/comfyui-moondream)
[comfyui-CLIPSeg](https://github.com/shadowcz007/comfyui-CLIPSeg)
## 🚀🚗🚚🏃 Workflow-to-APP
- 新增AppInfo节点,可以通过简单的配置,把workflow转变为一个Web APP。
- 支持多个web app 切换
@@ -115,6 +125,11 @@ https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43
- [Added DynamicDelayByText, enabling delayed execution based on input text length.](./workflow/audio-chatgpt-workflow.json)
- [使用CkptNames 对比不同的模型效果](./workflow/ckpts-image-workflow.json)
- [CkptNames compare the effects of different models.](./workflow/ckpts-image-workflow.json)
## Other Nodes
@@ -130,15 +145,6 @@ https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43
![TransparentImage](./assets/TransparentImage.png)
> Consistency Decoder
[openai Consistency Decoder]( https://github.com/openai/consistencydecoder)
![Consistency](./assets/consistency.png)
After downloading the OpenAI VAE model, place it in the "model/vae" directory for use.
https://openaipublic.azureedge.net/diff-vae/c9cebd3132dd9c42936d803e33424145a748843c8f716c0814838bdc8a2fe7cb/decoder.pt
> FeatheredMask、SmoothMask
Add edges to an image.
@@ -151,6 +157,12 @@ Add edges to an image.
from [simple-lama-inpainting](https://github.com/enesmsahin/simple-lama-inpainting)
> rembgNode
"briarmbg","u2net","u2netp","u2net_human_seg","u2net_cloth_seg","silueta","isnet-general-use","isnet-anime"
### Improvement
- Add "help" option to the context menu for each node.
@@ -167,8 +179,6 @@ An improvement has been made to directly redirect to GitHub to search for missin
[Download rembg Models](https://github.com/danielgatis/rembg/tree/main#Models),move to:models/rembg
[Download CLIPSeg](https://huggingface.co/CIDAS/clipseg-rd64-refined/tree/main), move to : models/clipseg
[Download lama](https://github.com/enesmsahin/simple-lama-inpainting/releases/download/v0.1.0/big-lama.pt), move to : models/lama
[Download Salesforce/blip-image-captioning-base](https://huggingface.co/Salesforce/blip-image-captioning-base), move to : models/clip_interrogator/Salesforce/blip-image-captioning-base
@@ -206,18 +216,18 @@ If you are using a venv, make sure you have it activated before installation and
pip3 install -r requirements.txt
```
#### Chinese community
访问 [www.mixcomfy.com](https://www.mixcomfy.com),获得更多内测功能,关注微信公众号:Mixlab无界社区
#### Thanks:
[ComfyUI-CLIPSeg](https://github.com/biegert/ComfyUI-CLIPSeg/tree/main)
#### discussions:
[discussions](https://github.com/shadowcz007/comfyui-mixlab-nodes/discussions)
<picture>
<source
media="(prefers-color-scheme: dark)"
+18 -33
View File
@@ -222,7 +222,8 @@ def get_my_workflow_for_app(filename="my_workflow_app.json",category="",is_all=F
# print(item)
try:
x=item["data"]
if i==0:
# 管理员模式,读取全部数据
if i==0 or is_all:
apps.append({
"filename":item["filename"],
# "category":item['category'],
@@ -430,7 +431,7 @@ async def new_start(self, address, port, verbose=True, call_on_start=None):
PromptServer.start=new_start
# 创建路由表
routes = web.RouteTableDef()
routes = PromptServer.instance.routes
@routes.post('/mixlab')
async def mixlab_hander(request):
@@ -516,27 +517,6 @@ async def nodes_map_hander(request):
return web.json_response(result)
# 把插件自定义的路由添加到comfyui server里
def new_add_routes(self):
import nodes
try:
self.user_manager.add_routes(self.routes)
except:
print('pls update')
self.app.add_routes(routes)
self.app.add_routes(self.routes)
for name, dir in nodes.EXTENSION_WEB_DIRS.items():
self.app.add_routes([
web.static('/extensions/' + urllib.parse.quote(name), dir, follow_symlinks=True),
])
self.app.add_routes([
web.static('/', self.web_root, follow_symlinks=True),
])
PromptServer.add_routes=new_add_routes
# 扩展api接口
# from server import PromptServer
@@ -551,13 +531,13 @@ PromptServer.add_routes=new_add_routes
# 导入节点
from .nodes.PromptNode import EmbeddingPrompt,RandomPrompt,PromptSlide,PromptSimplification,PromptImage,JoinWithDelimiter
from .nodes.ImageNode import SplitImage,GridOutput,GetImageSize_,MirroredImage,ImageColorTransfer,NoiseImage,TransparentImage,GradientImage,LoadImagesFromPath,LoadImagesFromURL,ResizeImage,TextImage,SvgImage,Image3D,ShowLayer,NewLayer,MergeLayers,AreaToMask,SmoothMask,SplitLongMask,ImageCropByAlpha,EnhanceImage,FaceToMask
from .nodes.Vae import VAELoader,VAEDecode
from .nodes.ImageNode import SaveImageToLocal,SplitImage,GridOutput,GetImageSize_,MirroredImage,ImageColorTransfer,NoiseImage,TransparentImage,GradientImage,LoadImagesFromPath,LoadImagesFromURL,ResizeImage,TextImage,SvgImage,Image3D,ShowLayer,NewLayer,MergeLayers,CenterImage,AreaToMask,SmoothMask,SplitLongMask,ImageCropByAlpha,EnhanceImage,FaceToMask
# from .nodes.Vae import VAELoader,VAEDecode
from .nodes.ScreenShareNode import ScreenShareNode,FloatingVideo
from .nodes.Clipseg import CLIPSeg,CombineMasks
from .nodes.ChatGPT import ChatGPTNode,ShowTextForGPT,CharacterInText
from .nodes.Audio import GamePal,SpeechRecognition,SpeechSynthesis
from .nodes.Utils import TESTNODE_,AppInfo,IntNumber,FloatSlider,TextInput,ColorInput,FontInput,TextToNumber,DynamicDelayProcessor,LimitNumber,SwitchByIndex,MultiplicationNode
from .nodes.Utils import CreateLoraNames,CreateSampler_names,CreateCkptNames,CreateSeedNode,TESTNODE_,AppInfo,IntNumber,FloatSlider,TextInput,ColorInput,FontInput,TextToNumber,DynamicDelayProcessor,LimitNumber,SwitchByIndex,MultiplicationNode
from .nodes.Mask import OutlineMask,FeatheredMask
@@ -586,6 +566,7 @@ NODE_CLASS_MAPPINGS = {
"ShowLayer":ShowLayer,
"NewLayer":NewLayer,
"SplitImage":SplitImage,
"CenterImage":CenterImage,
"GridOutput":GridOutput,
"MergeLayers":MergeLayers,
"SplitLongMask":SplitLongMask,
@@ -594,12 +575,11 @@ NODE_CLASS_MAPPINGS = {
"FaceToMask":FaceToMask,
"AreaToMask":AreaToMask,
"ImageCropByAlpha":ImageCropByAlpha,
"VAELoaderConsistencyDecoder":VAELoader,
"VAEDecodeConsistencyDecoder":VAEDecode,
# "VAELoaderConsistencyDecoder":VAELoader,
"SaveImageToLocal":SaveImageToLocal,
# "VAEDecodeConsistencyDecoder":VAEDecode,
"ScreenShare":ScreenShareNode,
"FloatingVideo":FloatingVideo,
"CLIPSeg_":CLIPSeg,
"CombineMasks_":CombineMasks,
"ChatGPTOpenAI":ChatGPTNode,
"ShowTextForGPT":ShowTextForGPT,
"CharacterInText":CharacterInText,
@@ -617,7 +597,11 @@ NODE_CLASS_MAPPINGS = {
"SwitchByIndex":SwitchByIndex,
"LimitNumber":LimitNumber,
"OutlineMask":OutlineMask,
"JoinWithDelimiter":JoinWithDelimiter
"JoinWithDelimiter":JoinWithDelimiter,
"Seed_":CreateSeedNode,
"CkptNames_":CreateCkptNames,
"SamplerNames_":CreateSampler_names,
"LoraNames_":CreateLoraNames
# "LaMaInpainting":LaMaInpainting
# "GamePal":GamePal
}
@@ -644,7 +628,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"PromptGenerate_Mix":"PromptGenerate ♾️Mixlab",
"ChinesePrompt_Mix":"ChinesePrompt ♾️Mixlab",
"GamePal":"GamePal ♾️Mixlab",
"RembgNode_Mix":"Removebg"
"RembgNode_Mix":"Removebg",
"LoraNames_":"LoraName_TriggerWords.safetensors"
}
# web ui的节点功能
Binary file not shown.

Before

Width:  |  Height:  |  Size: 784 KiB

+8 -5
View File
@@ -4773,12 +4773,15 @@
"ResizeImage",
"NoiseImage",
"PromptImage",
"SaveImageToLocal",
"AreaToMask",
"CLIPSeg_",
"CharacterInText",
"ChatGPTOpenAI",
"Color",
"CombineMasks_",
"Seed_",
"CkptNames_",
"SamplerNames_",
"LoraNames_",
"EnhanceImage",
"GradientImage",
"FaceToMask",
@@ -4790,6 +4793,7 @@
"LoadImagesFromURL",
"MergeLayers",
"NewLayer",
"CenterImage",
"RandomPrompt",
"PromptSlide",
"PromptSimplification",
@@ -4805,12 +4809,11 @@
"TextImage",
"ResizeImageMixlab",
"TransparentImage",
"VAEDecodeConsistencyDecoder",
"VAELoaderConsistencyDecoder",
"TextToNumber",
"TextInput_",
"DynamicDelayProcessor",
"LaMaInpainting"
"LaMaInpainting",
"Moondream"
],
{
"title_aux": "comfyui-mixlab-nodes"
+81 -6
View File
@@ -1,7 +1,26 @@
import openai
import time
import urllib.error
import re,json
import re,json,os,string,random
import folder_paths
import hashlib
def get_unique_hash(string):
hash_object = hashlib.sha1(string.encode())
unique_hash = hash_object.hexdigest()
return unique_hash
def generate_random_string(length):
letters = string.ascii_letters + string.digits
return ''.join(random.choice(letters) for _ in range(length))
class AnyType(str):
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
def __ne__(self, __value: object) -> bool:
return False
any_type = AnyType("*")
# 判断是否是azure服务
def is_azure_url(url):
@@ -73,8 +92,8 @@ class ChatGPTNode:
def INPUT_TYPES(cls):
return {
"required": {
"api_key":("KEY", {"default": "", "multiline": True}),
"api_url":("URL", {"default": "", "multiline": True}),
"api_key":("KEY", {"default": "", "multiline": True,"dynamicPrompts": False}),
"api_url":("URL", {"default": "", "multiline": True,"dynamicPrompts": False}),
"prompt": ("STRING", {"multiline": True,"dynamicPrompts": False}),
"system_content": ("STRING",
{
@@ -168,7 +187,10 @@ class ShowTextForGPT:
return {
"required": {
"text": ("STRING", {"forceInput": True,"dynamicPrompts": False}),
}
},
"optional":{
"output_dir": ("STRING",{"forceInput": True,"default": "","multiline": True,"dynamicPrompts": False}),
}
}
INPUT_IS_LIST = True
@@ -179,7 +201,60 @@ class ShowTextForGPT:
CATEGORY = "♾️Mixlab/GPT"
def run(self, text):
def run(self, text,output_dir=[""]):
# 类型纠正
texts=[]
for t in text:
if not isinstance(t, str):
t = str(t)
texts.append(t)
text=texts
if len(output_dir)==1 and (output_dir[0]=='' or os.path.dirname(output_dir[0])==''):
t='\n'.join(text)
output_dir=[
os.path.join(folder_paths.get_temp_directory(),
get_unique_hash(t)+'.txt'
)
]
elif len(output_dir)==1:
base=os.path.basename(output_dir[0])
t='\n'.join(text)
if base=='' or os.path.splitext(base)[1]=='':
base=get_unique_hash(t)+'.txt'
output_dir=[
os.path.join(output_dir[0],
base
)
]
# elif len(output_dir)>1:
if len(output_dir)==1 and len(text)>1:
output_dir=[output_dir[0] for _ in range(len(text))]
for i in range(len(text)):
o_fp=output_dir[i]
dirp=os.path.dirname(o_fp)
if dirp=='':
dirp=folder_paths.get_temp_directory()
o_fp=os.path.join(folder_paths.get_temp_directory(),o_fp
)
if not os.path.exists(dirp):
os.mkdir(dirp)
if not os.path.splitext(o_fp)[1].lower()=='.txt':
o_fp=o_fp+'.txt'
t=text[i]
with open(o_fp, 'w') as file:
file.write(t)
# print(text)
return {"ui": {"text": text}, "result": (text,)}
@@ -212,7 +287,7 @@ class CharacterInText:
def run(self, text,character,start_index):
# print(text,character,start_index)
b=1 if character in text else 0
b=1 if character.lower() in text.lower() else 0
return (b+start_index,)
-272
View File
@@ -1,272 +0,0 @@
#### Thanks:
# [ComfyUI-CLIPSeg](https://github.com/biegert/ComfyUI-CLIPSeg/tree/main)
from transformers import CLIPSegProcessor, CLIPSegForImageSegmentation
from PIL import Image
import torch
import torchvision.transforms as T
import numpy as np
from torchvision.transforms.functional import to_pil_image
import matplotlib.pyplot as plt
import matplotlib.cm as cm
import cv2
from scipy.ndimage import gaussian_filter
from typing import Optional, Tuple
import warnings,os
warnings.filterwarnings("ignore", category=UserWarning, module="torch")
warnings.filterwarnings("ignore", category=UserWarning, module="safetensors")
import folder_paths
import logging
logger = logging.getLogger('CLIPSeg nodes')
clipseg_model_dir = os.path.join(folder_paths.models_dir, "clipseg")
if not os.path.exists(clipseg_model_dir):
print(f"## clipseg model not found: {clipseg_model_dir},pls download from https://huggingface.co/CIDAS/clipseg-rd64-refined/tree/main")
clipseg_model_dir='CIDAS/clipseg-rd64-refined'
"""Helper methods for CLIPSeg nodes"""
# Tensor to PIL
def tensor2pil(image):
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
# Convert PIL to Tensor
def pil2tensor(image):
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
def tensor_to_numpy(tensor: torch.Tensor) -> np.ndarray:
"""Convert a tensor to a numpy array and scale its values to 0-255."""
array = tensor.numpy().squeeze()
return (array * 255).astype(np.uint8)
def numpy_to_tensor(array: np.ndarray) -> torch.Tensor:
"""Convert a numpy array to a tensor and scale its values from 0-255 to 0-1."""
array = array.astype(np.float32) / 255.0
return torch.from_numpy(array)[None,]
def apply_colormap(mask: torch.Tensor, colormap) -> np.ndarray:
"""Apply a colormap to a tensor and convert it to a numpy array."""
colored_mask = colormap(mask.numpy())[:, :, :3]
return (colored_mask * 255).astype(np.uint8)
def resize_image(image: np.ndarray, dimensions: Tuple[int, int]) -> np.ndarray:
"""Resize an image to the given dimensions using linear interpolation."""
return cv2.resize(image, dimensions, interpolation=cv2.INTER_LINEAR)
def overlay_image(background: np.ndarray, foreground: np.ndarray, alpha: float) -> np.ndarray:
"""Overlay the foreground image onto the background with a given opacity (alpha)."""
return cv2.addWeighted(background, 1 - alpha, foreground, alpha, 0)
def dilate_mask(mask: torch.Tensor, dilation_factor: float) -> torch.Tensor:
"""Dilate a mask using a square kernel with a given dilation factor."""
kernel_size = int(dilation_factor * 2) + 1
kernel = np.ones((kernel_size, kernel_size), np.uint8)
mask_dilated = cv2.dilate(mask.numpy(), kernel, iterations=1)
return torch.from_numpy(mask_dilated)
class CLIPSeg:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
"""
Return a dictionary which contains config for all input fields.
Some types (string): "MODEL", "VAE", "CLIP", "CONDITIONING", "LATENT", "IMAGE", "INT", "STRING", "FLOAT".
Input types "INT", "STRING" or "FLOAT" are special values for fields on the node.
The type can be a list for selection.
Returns: `dict`:
- Key input_fields_group (`string`): Can be either required, hidden or optional. A node class must have property `required`
- Value input_fields (`dict`): Contains input fields config:
* Key field_name (`string`): Name of a entry-point method's argument
* Value field_config (`tuple`):
+ First value is a string indicate the type of field or a list for selection.
+ Secound value is a config for type "INT", "STRING" or "FLOAT".
"""
return {"required":
{
"image": ("IMAGE",),
"text": ("STRING", {"multiline": False,"dynamicPrompts": False}),
},
"optional":
{
"blur": ("FLOAT", {"min": 0, "max": 15, "step": 0.1, "default": 3}),
"threshold": ("FLOAT", {"min": 0, "max": 1, "step": 0.05, "default": 0.3}),
"dilation_factor": ("INT", {"min": 0, "max": 10, "step": 1, "default": 4}),
}
}
CATEGORY = "♾️Mixlab/Mask"
RETURN_TYPES = ("MASK", "IMAGE", "IMAGE",)
RETURN_NAMES = ("Mask","Heatmap Mask", "BW Mask")
# INPUT_IS_LIST = True
OUTPUT_IS_LIST = (False,False,False,)
FUNCTION = "segment_image"
def segment_image(self, image: torch.Tensor, text: str, blur: float, threshold: float, dilation_factor: int) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Create a segmentation mask from an image and a text prompt using CLIPSeg.
Args:
image (torch.Tensor): The image to segment.
text (str): The text prompt to use for segmentation.
blur (float): How much to blur the segmentation mask.
threshold (float): The threshold to use for binarizing the segmentation mask.
dilation_factor (int): How much to dilate the segmentation mask.
Returns:
Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: The segmentation mask, the heatmap mask, and the binarized mask.
"""
# Convert the Tensor to a PIL image
image_np = image.numpy().squeeze() # Remove the first dimension (batch size of 1)
# Convert the numpy array back to the original range (0-255) and data type (uint8)
image_np = (image_np * 255).astype(np.uint8)
# Create a PIL image from the numpy array
i = Image.fromarray(image_np, mode="RGB")
processor = CLIPSegProcessor.from_pretrained(clipseg_model_dir)
model = CLIPSegForImageSegmentation.from_pretrained(clipseg_model_dir)
prompt = text
input_prc = processor(text=prompt, images=i, padding="max_length", return_tensors="pt")
# Predict the segemntation mask
with torch.no_grad():
outputs = model(**input_prc)
tensor = torch.sigmoid(outputs[0]) # get the mask
# Apply a threshold to the original tensor to cut off low values
thresh = threshold
tensor_thresholded = torch.where(tensor > thresh, tensor, torch.tensor(0, dtype=torch.float))
# Apply Gaussian blur to the thresholded tensor
sigma = blur
tensor_smoothed = gaussian_filter(tensor_thresholded.numpy(), sigma=sigma)
tensor_smoothed = torch.from_numpy(tensor_smoothed)
# Normalize the smoothed tensor to [0, 1]
mask_normalized = (tensor_smoothed - tensor_smoothed.min()) / (tensor_smoothed.max() - tensor_smoothed.min())
# Dilate the normalized mask
mask_dilated = dilate_mask(mask_normalized, dilation_factor)
# Convert the mask to a heatmap and a binary mask
heatmap = apply_colormap(mask_dilated, cm.viridis)
binary_mask = apply_colormap(mask_dilated, cm.Greys_r)
# Overlay the heatmap and binary mask on the original image
dimensions = (image_np.shape[1], image_np.shape[0])
heatmap_resized = resize_image(heatmap, dimensions)
binary_mask_resized = resize_image(binary_mask, dimensions)
alpha_heatmap, alpha_binary = 0.5, 1
overlay_heatmap = overlay_image(image_np, heatmap_resized, alpha_heatmap)
overlay_binary = overlay_image(image_np, binary_mask_resized, alpha_binary)
# Convert the numpy arrays to tensors
image_out_heatmap = numpy_to_tensor(overlay_heatmap)
image_out_binary = numpy_to_tensor(overlay_binary)
# Save or display the resulting binary mask
binary_mask_image = Image.fromarray(binary_mask_resized[..., 0])
# convert PIL image to numpy array
tensor_bw = binary_mask_image.convert("L")
tensor_bw=pil2tensor(tensor_bw)
# tensor_bw = np.array(tensor_bw).astype(np.float32) / 255.0
# tensor_bw = torch.from_numpy(tensor_bw)[None,]
# tensor_bw = tensor_bw.squeeze(0)[..., 0]
return (tensor_bw, image_out_heatmap, image_out_binary,)
#OUTPUT_NODE = False
class CombineMasks:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {"required":
{
"input_image": ("IMAGE", ),
"mask_1": ("MASK", ),
"mask_2": ("MASK", ),
},
"optional":
{
"mask_3": ("MASK",),
},
}
CATEGORY = "♾️Mixlab/Mask"
RETURN_TYPES = ("MASK", "IMAGE", "IMAGE",)
RETURN_NAMES = ("Combined Mask","Heatmap Mask", "BW Mask")
FUNCTION = "combine_masks"
def combine_masks(self, input_image: torch.Tensor, mask_1: torch.Tensor, mask_2: torch.Tensor, mask_3: Optional[torch.Tensor] = None) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""A method that combines two or three masks into one mask. Takes in tensors and returns the mask as a tensor, as well as the heatmap and binary mask as tensors."""
# Combine masks
if mask_1 is not None:
mask_1 = mask_1.squeeze()
if mask_2 is not None:
mask_2 = mask_2.squeeze()
if mask_3 is not None:
mask_3 = mask_3.squeeze()
print(mask_1.shape,mask_2.shape , mask_3.shape)
combined_mask = mask_1 + mask_2 + mask_3 if mask_3 is not None else mask_1 + mask_2
# print(combined_mask)
# Convert image and masks to numpy arrays
image_np = tensor_to_numpy(input_image)
heatmap = apply_colormap(combined_mask, cm.viridis)
binary_mask = apply_colormap(combined_mask, cm.Greys_r)
# Resize heatmap and binary mask to match the original image dimensions
dimensions = (image_np.shape[1], image_np.shape[0])
# print('heatmap',heatmap)
if dimensions is None or dimensions[0] == 0 or dimensions[1] == 0:
raise ValueError("Invalid dimensions")
heatmap_resized = resize_image(heatmap, dimensions)
binary_mask_resized = resize_image(binary_mask, dimensions)
# Overlay the heatmap and binary mask onto the original image
alpha_heatmap, alpha_binary = 0.5, 1
overlay_heatmap = overlay_image(image_np, heatmap_resized, alpha_heatmap)
overlay_binary = overlay_image(image_np, binary_mask_resized, alpha_binary)
# Convert overlays to tensors
image_out_heatmap = numpy_to_tensor(overlay_heatmap)
image_out_binary = numpy_to_tensor(overlay_binary)
return combined_mask, image_out_heatmap, image_out_binary
# A dictionary that contains all nodes you want to export with their names
# NOTE: names should be globally unique
# NODE_CLASS_MAPPINGS = {
# "CLIPSeg": CLIPSeg,
# "CombineSegMasks": CombineMasks,
# }
+284 -28
View File
@@ -9,11 +9,22 @@ from io import BytesIO
import folder_paths
import json,io
from comfy.cli_args import args
import cv2
import math
import cv2
import string
import math,glob
from .Watcher import FolderWatcher
def generate_random_string(length):
letters = string.ascii_letters + string.digits
return ''.join(random.choice(letters) for _ in range(length))
class AnyType(str):
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
def __ne__(self, __value: object) -> bool:
return False
any_type = AnyType("*")
FONT_PATH= os.path.abspath(os.path.join(os.path.dirname(__file__),'../assets/王汉宗颜楷体繁.ttf'))
@@ -129,10 +140,10 @@ def naive_cutout(img, mask,invert=True):
image using the mask.
"""
img=img.convert("RGBA")
# img=img.convert("RGBA")
mask=mask.convert("RGBA")
empty = Image.new("RGBA", (img.size), 0)
empty = Image.new("RGBA", (mask.size), 0)
red, green, blue, alpha = mask.split()
@@ -369,6 +380,7 @@ def get_images_filepath(f,white_bg=False):
for root, dirs, files in os.walk(f):
for file in files:
file_path = os.path.join(root, file)
file_name=os.path.basename(file_path)
try:
imgs=load_image(file_path,white_bg)
for img in imgs:
@@ -376,6 +388,7 @@ def get_images_filepath(f,white_bg=False):
"image":img['image'],
"mask":img['mask'],
"file_path":file_path,
"file_name":file_name,
"psd":len(imgs)>1
})
except:
@@ -383,12 +396,15 @@ def get_images_filepath(f,white_bg=False):
elif os.path.isfile(f):
try:
file_path = os.path.join(root, f)
file_name=os.path.basename(file_path)
imgs=load_image(f,white_bg)
for img in imgs:
images.append({
"image":img['image'],
"mask":img['mask'],
"file_path":file_path,
"file_name":file_name,
"psd":len(imgs)>1
})
except:
@@ -961,6 +977,7 @@ class TransparentImage:
# ui.images 节点里显示图片,和 传参,image_path自定义的数据,需要写节点的自定义ui
# result 里输出给下个节点的数据
# print('TransparentImage',len(images_rgb))
return {"ui":{"images": ui_images,"image_paths":image_paths},"result": (image_paths,images_rgb,images_rgba)}
@@ -1047,14 +1064,15 @@ class LoadImagesFromPath:
}
}
RETURN_TYPES = ('IMAGE','MASK','STRING',)
RETURN_TYPES = ('IMAGE','MASK','STRING','STRING',)
RETURN_NAMES = ("IMAGE","MASK","prompt_for_FloatingVideo","filepaths",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Image"
# INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True,True,False,)
OUTPUT_IS_LIST = (True,True,False,True,)
global watcher_folder
watcher_folder=None
@@ -1090,10 +1108,12 @@ class LoadImagesFromPath:
imgs=[]
masks=[]
file_names=[]
for im in sorted_files:
imgs.append(im['image'])
masks.append(im['mask'])
file_names.append(im['file_name'])
# print('index_variable',index_variable)
@@ -1101,11 +1121,12 @@ class LoadImagesFromPath:
if index_variable!=-1:
imgs=[imgs[index_variable]] if index_variable < len(imgs) else None
masks=[masks[index_variable]] if index_variable < len(masks) else None
file_names=[file_names[index_variable]] if index_variable < len(file_names) else None
except Exception as e:
print("发生了一个未知的错误:", str(e))
# print('#prompt::::',prompt)
return {"ui": {"seed": [1]}, "result":(imgs,masks,prompt,)}
return {"ui": {"seed": [1]}, "result":(imgs,masks,prompt,file_names,)}
# TODO 扩大选区的功能,重新输出mask
@@ -1136,7 +1157,12 @@ class ImageCropByAlpha:
# print(RGBA)
im=tensor2pil(RGBA)
im=naive_cutout(im,im)
# 要把im的alpha通道转为mask
im=im.convert('RGBA')
red, green, blue, alpha = im.split()
im=naive_cutout(bf_im,alpha)
x, y, w, h=get_not_transparent_area(im)
# print('#ForImageCrop:',w, h,x, y,)
@@ -1151,11 +1177,12 @@ class ImageCropByAlpha:
height_1=h
img = image[:,y:to_y, x:to_x, :]
# tensor2pil(img).save('test2.png')
# 原图的mask
ori=RGBA[:,y:to_y, x:to_x, :]
ori=tensor2pil(ori)
# ori.save('test.png')
# 创建一个新的图像对象,大小和模式与原始图像相同
new_image = Image.new("RGBA", ori.size)
@@ -1176,7 +1203,7 @@ class ImageCropByAlpha:
if a != 0:
new_pixel_data[x, y] = (255, 255, 255, 255)
else:
new_pixel_data[x, y] = (r, g, b, a)
new_pixel_data[x, y] = (0,0,0,0)
# 保存修改后的图像
# new_image.save("output.png")
@@ -1250,6 +1277,9 @@ class LoadImagesFromURL:
return {"required": {
"url": ("STRING",{"multiline": True,"default": "https://","dynamicPrompts": False}),
},
"optional":{
"seed": (any_type, {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
}
}
RETURN_TYPES = ("IMAGE","MASK",)
@@ -1266,7 +1296,7 @@ class LoadImagesFromURL:
global urls_image
urls_image={}
def run(self,url):
def run(self,url,seed=0):
global urls_image
print(urls_image)
def filter_http_urls(urls):
@@ -1599,6 +1629,18 @@ class NewLayer:
return (layer_n,)
def createMask(image,x,y,w,h):
mask = Image.new("L", image.size)
pixels = mask.load()
# 遍历指定区域的像素,将其设置为黑色(0 表示黑色)
for i in range(int(x), int(x + w)):
for j in range(int(y), int(y + h)):
pixels[i, j] = 255
# mask.save("mask.png")
return mask
def splitImage(image, num):
width, height = image.size
@@ -1617,6 +1659,19 @@ def splitImage(image, num):
return grid_coordinates
def centerImage(margin,canvas):
w,h=canvas.size
l,t,r,b=margin
x=l
y=t
width=w-r-l
height=h-t-b
return (x,y,width,height)
# # 读取图片
# image = Image.open("path_to_your_image.jpg")
@@ -1654,8 +1709,8 @@ class SplitImage:
}
}
RETURN_TYPES = ("_GRID","_GRID",)
RETURN_NAMES = ("grids","grid")
RETURN_TYPES = ("_GRID","_GRID","MASK",)
RETURN_NAMES = ("grids","grid","mask",)
FUNCTION = "run"
@@ -1666,10 +1721,10 @@ class SplitImage:
def run(self,image,num,seed):
image=tensor2pil(image)
grids=splitImage(image,num)
if seed>=num:
if seed>num:
num=int(seed / 500 * num)-1
else:
num=seed-1
@@ -1678,7 +1733,71 @@ class SplitImage:
g=grids[num]
return (grids,g,)
x,y,w,h=g
mask=createMask(image, x,y,w,h)
mask=pil2tensor(mask)
return (grids,g,mask,)
class CenterImage:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"canvas": ("IMAGE",),
"left": ("INT",{
"default":24,
"min": 0, #Minimum value
"max": 5000, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"top": ("INT",{
"default":24,
"min": 0, #Minimum value
"max": 5000, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"right": ("INT",{
"default": 24,
"min": 0, #Minimum value
"max": 5000, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"bottom": ("INT",{
"default": 24,
"min": 0, #Minimum value
"max": 5000, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
}
}
RETURN_TYPES = ("_GRID","MASK",)
RETURN_NAMES = ("grid","mask",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Layer"
INPUT_IS_LIST = False
# OUTPUT_IS_LIST = (True,)
def run(self,canvas,left,top,right,bottom):
canvas=tensor2pil(canvas)
grid=centerImage((left,top,right,bottom),canvas)
mask=createMask(canvas,left,top,canvas.width-left-right,canvas.height-top-bottom)
return (grid,pil2tensor(mask),)
class GridOutput:
@@ -2040,14 +2159,14 @@ class ResizeImage:
"default": 512,
"min": 1, #Minimum value
"max": 8192, #Maximum value
"step": 1, #Slider's step
"step": 8, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"height": ("INT",{
"default": 512,
"min": 1, #Minimum value
"max": 8192, #Maximum value
"step": 1, #Slider's step
"step": 8, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"scale_option": (["width","height",'overall','center'],),
@@ -2058,20 +2177,21 @@ class ResizeImage:
"image": ("IMAGE",),
"average_color": (["on",'off'],),
"fill_color":("STRING",{"multiline": False,"default": "#FFFFFF","dynamicPrompts": False}),
"mask": ("MASK",),
}
}
RETURN_TYPES = ("IMAGE","IMAGE","STRING",)
RETURN_NAMES = ("image","average_image","average_hex",)
RETURN_TYPES = ("IMAGE","IMAGE","STRING","MASK",)
RETURN_NAMES = ("image","average_image","average_hex","mask",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Image"
INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True,True,True,)
OUTPUT_IS_LIST = (True,True,True,True,)
def run(self,width,height,scale_option,image=None,average_color=['on'],fill_color=["#FFFFFF"]):
def run(self,width,height,scale_option,image=None,average_color=['on'],fill_color=["#FFFFFF"],mask=None):
w=width[0]
h=height[0]
@@ -2080,6 +2200,7 @@ class ResizeImage:
fill_color=fill_color[0]
imgs=[]
masks=[]
average_images=[]
hexs=[]
@@ -2108,8 +2229,20 @@ class ResizeImage:
a_im=pil2tensor(a_im)
average_images.append(a_im)
hexs.append(hex)
try:
for mas in mask:
for ma in mas:
ma=tensor2pil(ma)
ma=ma.convert('RGB')
ma=resize_image(ma,scale_option,w,h,fill_color)
ma=ma.convert('L')
ma=pil2tensor(ma)
masks.append(ma)
except:
print('')
return (imgs,average_images,hexs,)
return (imgs,average_images,hexs,masks,)
class MirroredImage:
@@ -2153,19 +2286,41 @@ class GetImageSize_:
return {
"required": {
"image": ("IMAGE",),
}
},
"optional":{
"min_width":("INT", {
"default": 512,
"min":1, #Minimum value
"max": 2048, #Maximum value
"step": 8, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
})
},
}
RETURN_TYPES = ("INT", "INT")
RETURN_NAMES = ("width", "height")
RETURN_TYPES = ("INT", "INT","INT", "INT",)
RETURN_NAMES = ("width", "height","min_width", "min_height",)
FUNCTION = "get_size"
CATEGORY = "♾️Mixlab/Image"
def get_size(self, image):
def get_size(self, image,min_width):
_, height, width, _ = image.shape
return (width, height)
# 如果比min_widht,还小,则输出 min width
if min_width>width:
im=tensor2pil(image)
im=resize_image(im,'width',min_width,min_width,"white")
im=im.convert('RGB')
min_width,min_height=im.size
else:
min_width=width
min_height=height
return (width, height,min_width,min_height,)
@@ -2212,3 +2367,104 @@ class ImageColorTransfer:
class SaveImageToLocal:
def __init__(self):
self.output_dir = folder_paths.get_output_directory()
self.type = "output"
self.compress_level = 4
@classmethod
def INPUT_TYPES(s):
return {"required":
{"images": ("IMAGE", ),
"file_path": ("STRING",{"multiline": True,"default": "","dynamicPrompts": False}),
},
"hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"},
}
RETURN_TYPES = ()
FUNCTION = "save_images"
OUTPUT_NODE = True
CATEGORY = "♾️Mixlab/Image"
def save_images(self, images,file_path , prompt=None, extra_pnginfo=None):
filename_prefix = os.path.basename(file_path)
if file_path=='':
filename_prefix="ComfyUI"
filename_prefix, _ = os.path.splitext(filename_prefix)
_, extension = os.path.splitext(file_path)
if extension:
# 是文件名,需要处理
file_path=os.path.dirname(file_path)
# filename_prefix=
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, self.output_dir, images[0].shape[1], images[0].shape[0])
if not os.path.exists(file_path) and not extension:
# 使用os.makedirs函数创建新目录
os.makedirs(file_path)
print("目录已创建")
else:
print("目录已存在")
# 使用glob模块获取当前目录下的所有文件
if file_path=="":
files = glob.glob(full_output_folder + '/*')
else:
files = glob.glob(file_path + '/*')
# 统计文件数量
file_count = len(files)
counter+=file_count
print('统计文件数量',file_count,counter)
results = list()
for image in images:
i = 255. * image.cpu().numpy()
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
metadata = None
if not args.disable_metadata:
metadata = PngInfo()
if prompt is not None:
metadata.add_text("prompt", json.dumps(prompt))
if extra_pnginfo is not None:
for x in extra_pnginfo:
metadata.add_text(x, json.dumps(extra_pnginfo[x]))
file = f"{filename}_{counter:05}_.png"
if file_path=="":
fp=os.path.join(full_output_folder, file)
if os.path.exists(fp):
file = f"{filename}_{counter:05}_{generate_random_string(8)}.png"
fp=os.path.join(full_output_folder, file)
img.save(fp, pnginfo=metadata, compress_level=self.compress_level)
results.append({
"filename": file,
"subfolder": subfolder,
"type": self.type
})
else:
fp=os.path.join(file_path, file)
if os.path.exists(fp):
file = f"{filename}_{counter:05}_{generate_random_string(8)}.png"
fp=os.path.join(file_path, file)
img.save(os.path.join(file_path, file), pnginfo=metadata, compress_level=self.compress_level)
results.append({
"filename": file,
"subfolder": file_path,
"type": self.type
})
counter += 1
return ()
+4 -2
View File
@@ -323,7 +323,9 @@ class RandomPrompt:
"random_sample": (["enable", "disable"],),
# "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "step": 1}),
},
"optional":{
"seed": (any_type, {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
}
}
@@ -339,7 +341,7 @@ class RandomPrompt:
# 运行的函数
def run(self,max_count,mutable_prompt,immutable_prompt,random_sample):
def run(self,max_count,mutable_prompt,immutable_prompt,random_sample,seed=0):
# print('#运行的函数',mutable_prompt,immutable_prompt,max_count,random_sample)
# Split the text into an array of words
+533 -3
View File
@@ -8,6 +8,467 @@ import comfy.utils
import numpy as np
import torch
from huggingface_hub import hf_hub_download
import torch.nn as nn
import torch.nn.functional as F
from torchvision.transforms.functional import normalize
# BRIA-RMBG-1.4 / briarmbg.py
class REBNCONV(nn.Module):
def __init__(self,in_ch=3,out_ch=3,dirate=1,stride=1):
super(REBNCONV,self).__init__()
self.conv_s1 = nn.Conv2d(in_ch,out_ch,3,padding=1*dirate,dilation=1*dirate,stride=stride)
self.bn_s1 = nn.BatchNorm2d(out_ch)
self.relu_s1 = nn.ReLU(inplace=True)
def forward(self,x):
hx = x
xout = self.relu_s1(self.bn_s1(self.conv_s1(hx)))
return xout
## upsample tensor 'src' to have the same spatial size with tensor 'tar'
def _upsample_like(src,tar):
src = F.interpolate(src,size=tar.shape[2:],mode='bilinear')
return src
### RSU-7 ###
class RSU7(nn.Module):
def __init__(self, in_ch=3, mid_ch=12, out_ch=3, img_size=512):
super(RSU7,self).__init__()
self.in_ch = in_ch
self.mid_ch = mid_ch
self.out_ch = out_ch
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1) ## 1 -> 1/2
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool4 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool5 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv6 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.rebnconv7 = REBNCONV(mid_ch,mid_ch,dirate=2)
self.rebnconv6d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv5d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
def forward(self,x):
b, c, h, w = x.shape
hx = x
hxin = self.rebnconvin(hx)
hx1 = self.rebnconv1(hxin)
hx = self.pool1(hx1)
hx2 = self.rebnconv2(hx)
hx = self.pool2(hx2)
hx3 = self.rebnconv3(hx)
hx = self.pool3(hx3)
hx4 = self.rebnconv4(hx)
hx = self.pool4(hx4)
hx5 = self.rebnconv5(hx)
hx = self.pool5(hx5)
hx6 = self.rebnconv6(hx)
hx7 = self.rebnconv7(hx6)
hx6d = self.rebnconv6d(torch.cat((hx7,hx6),1))
hx6dup = _upsample_like(hx6d,hx5)
hx5d = self.rebnconv5d(torch.cat((hx6dup,hx5),1))
hx5dup = _upsample_like(hx5d,hx4)
hx4d = self.rebnconv4d(torch.cat((hx5dup,hx4),1))
hx4dup = _upsample_like(hx4d,hx3)
hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))
hx3dup = _upsample_like(hx3d,hx2)
hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
hx2dup = _upsample_like(hx2d,hx1)
hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
return hx1d + hxin
### RSU-6 ###
class RSU6(nn.Module):
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
super(RSU6,self).__init__()
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool4 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.rebnconv6 = REBNCONV(mid_ch,mid_ch,dirate=2)
self.rebnconv5d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
def forward(self,x):
hx = x
hxin = self.rebnconvin(hx)
hx1 = self.rebnconv1(hxin)
hx = self.pool1(hx1)
hx2 = self.rebnconv2(hx)
hx = self.pool2(hx2)
hx3 = self.rebnconv3(hx)
hx = self.pool3(hx3)
hx4 = self.rebnconv4(hx)
hx = self.pool4(hx4)
hx5 = self.rebnconv5(hx)
hx6 = self.rebnconv6(hx5)
hx5d = self.rebnconv5d(torch.cat((hx6,hx5),1))
hx5dup = _upsample_like(hx5d,hx4)
hx4d = self.rebnconv4d(torch.cat((hx5dup,hx4),1))
hx4dup = _upsample_like(hx4d,hx3)
hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))
hx3dup = _upsample_like(hx3d,hx2)
hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
hx2dup = _upsample_like(hx2d,hx1)
hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
return hx1d + hxin
### RSU-5 ###
class RSU5(nn.Module):
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
super(RSU5,self).__init__()
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=2)
self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
def forward(self,x):
hx = x
hxin = self.rebnconvin(hx)
hx1 = self.rebnconv1(hxin)
hx = self.pool1(hx1)
hx2 = self.rebnconv2(hx)
hx = self.pool2(hx2)
hx3 = self.rebnconv3(hx)
hx = self.pool3(hx3)
hx4 = self.rebnconv4(hx)
hx5 = self.rebnconv5(hx4)
hx4d = self.rebnconv4d(torch.cat((hx5,hx4),1))
hx4dup = _upsample_like(hx4d,hx3)
hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))
hx3dup = _upsample_like(hx3d,hx2)
hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
hx2dup = _upsample_like(hx2d,hx1)
hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
return hx1d + hxin
### RSU-4 ###
class RSU4(nn.Module):
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
super(RSU4,self).__init__()
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=2)
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
def forward(self,x):
hx = x
hxin = self.rebnconvin(hx)
hx1 = self.rebnconv1(hxin)
hx = self.pool1(hx1)
hx2 = self.rebnconv2(hx)
hx = self.pool2(hx2)
hx3 = self.rebnconv3(hx)
hx4 = self.rebnconv4(hx3)
hx3d = self.rebnconv3d(torch.cat((hx4,hx3),1))
hx3dup = _upsample_like(hx3d,hx2)
hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
hx2dup = _upsample_like(hx2d,hx1)
hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
return hx1d + hxin
### RSU-4F ###
class RSU4F(nn.Module):
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
super(RSU4F,self).__init__()
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=2)
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=4)
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=8)
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=4)
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=2)
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
def forward(self,x):
hx = x
hxin = self.rebnconvin(hx)
hx1 = self.rebnconv1(hxin)
hx2 = self.rebnconv2(hx1)
hx3 = self.rebnconv3(hx2)
hx4 = self.rebnconv4(hx3)
hx3d = self.rebnconv3d(torch.cat((hx4,hx3),1))
hx2d = self.rebnconv2d(torch.cat((hx3d,hx2),1))
hx1d = self.rebnconv1d(torch.cat((hx2d,hx1),1))
return hx1d + hxin
class myrebnconv(nn.Module):
def __init__(self, in_ch=3,
out_ch=1,
kernel_size=3,
stride=1,
padding=1,
dilation=1,
groups=1):
super(myrebnconv,self).__init__()
self.conv = nn.Conv2d(in_ch,
out_ch,
kernel_size=kernel_size,
stride=stride,
padding=padding,
dilation=dilation,
groups=groups)
self.bn = nn.BatchNorm2d(out_ch)
self.rl = nn.ReLU(inplace=True)
def forward(self,x):
return self.rl(self.bn(self.conv(x)))
class BriaRMBG(nn.Module):
def __init__(self,in_ch=3,out_ch=1):
super(BriaRMBG,self).__init__()
self.conv_in = nn.Conv2d(in_ch,64,3,stride=2,padding=1)
self.pool_in = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.stage1 = RSU7(64,32,64)
self.pool12 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.stage2 = RSU6(64,32,128)
self.pool23 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.stage3 = RSU5(128,64,256)
self.pool34 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.stage4 = RSU4(256,128,512)
self.pool45 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.stage5 = RSU4F(512,256,512)
self.pool56 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
self.stage6 = RSU4F(512,256,512)
# decoder
self.stage5d = RSU4F(1024,256,512)
self.stage4d = RSU4(1024,128,256)
self.stage3d = RSU5(512,64,128)
self.stage2d = RSU6(256,32,64)
self.stage1d = RSU7(128,16,64)
self.side1 = nn.Conv2d(64,out_ch,3,padding=1)
self.side2 = nn.Conv2d(64,out_ch,3,padding=1)
self.side3 = nn.Conv2d(128,out_ch,3,padding=1)
self.side4 = nn.Conv2d(256,out_ch,3,padding=1)
self.side5 = nn.Conv2d(512,out_ch,3,padding=1)
self.side6 = nn.Conv2d(512,out_ch,3,padding=1)
# self.outconv = nn.Conv2d(6*out_ch,out_ch,1)
def forward(self,x):
hx = x
hxin = self.conv_in(hx)
#hx = self.pool_in(hxin)
#stage 1
hx1 = self.stage1(hxin)
hx = self.pool12(hx1)
#stage 2
hx2 = self.stage2(hx)
hx = self.pool23(hx2)
#stage 3
hx3 = self.stage3(hx)
hx = self.pool34(hx3)
#stage 4
hx4 = self.stage4(hx)
hx = self.pool45(hx4)
#stage 5
hx5 = self.stage5(hx)
hx = self.pool56(hx5)
#stage 6
hx6 = self.stage6(hx)
hx6up = _upsample_like(hx6,hx5)
#-------------------- decoder --------------------
hx5d = self.stage5d(torch.cat((hx6up,hx5),1))
hx5dup = _upsample_like(hx5d,hx4)
hx4d = self.stage4d(torch.cat((hx5dup,hx4),1))
hx4dup = _upsample_like(hx4d,hx3)
hx3d = self.stage3d(torch.cat((hx4dup,hx3),1))
hx3dup = _upsample_like(hx3d,hx2)
hx2d = self.stage2d(torch.cat((hx3dup,hx2),1))
hx2dup = _upsample_like(hx2d,hx1)
hx1d = self.stage1d(torch.cat((hx2dup,hx1),1))
#side output
d1 = self.side1(hx1d)
d1 = _upsample_like(d1,x)
d2 = self.side2(hx2d)
d2 = _upsample_like(d2,x)
d3 = self.side3(hx3d)
d3 = _upsample_like(d3,x)
d4 = self.side4(hx4d)
d4 = _upsample_like(d4,x)
d5 = self.side5(hx5d)
d5 = _upsample_like(d5,x)
d6 = self.side6(hx6)
d6 = _upsample_like(d6,x)
return [F.sigmoid(d1), F.sigmoid(d2), F.sigmoid(d3), F.sigmoid(d4), F.sigmoid(d5), F.sigmoid(d6)],[hx1d,hx2d,hx3d,hx4d,hx5d,hx6]
U2NET_HOME=os.path.join(folder_paths.models_dir, "rembg")
os.environ["U2NET_HOME"] = U2NET_HOME
@@ -48,6 +509,70 @@ except:
_available=False
def briarmbg_run(images=[]):
mroot=os.path.join(folder_paths.models_dir, "rembg")
m=os.path.join(mroot,'briarmbg.pth')
if os.path.exists(m)==False:
# 下载
m1=hf_hub_download("briaai/RMBG-1.4",
local_dir=mroot,
filename='model.pth',
local_dir_use_symlinks=False,
endpoint='https://hf-mirror.com')
os.rename(m1, m)
net=BriaRMBG()
if torch.cuda.is_available():
net.load_state_dict(torch.load(m))
net=net.cuda()
else:
net.load_state_dict(torch.load(m,map_location="cpu"))
net.eval()
masks=[]
rgba_images=[]
rgb_images=[]
for orig_image in images:
w,h = orig_im_size = orig_image.size
image = orig_image.convert('RGB')
model_input_size = (1024, 1024)
image = image.resize(model_input_size, Image.BILINEAR)
im_np = np.array(image)
im_tensor = torch.tensor(im_np, dtype=torch.float32).permute(2,0,1)
im_tensor = torch.unsqueeze(im_tensor,0)
im_tensor = torch.divide(im_tensor,255.0)
im_tensor = normalize(im_tensor,[0.5,0.5,0.5],[1.0,1.0,1.0])
if torch.cuda.is_available():
im_tensor=im_tensor.cuda()
result=net(im_tensor)
result = torch.squeeze(F.interpolate(result[0][0], size=(h,w), mode='bilinear') ,0)
ma = torch.max(result)
mi = torch.min(result)
result = (result-mi)/(ma-mi)
im_array = (result*255).cpu().data.numpy().astype(np.uint8)
mask = Image.fromarray(np.squeeze(im_array))
# mask.save('test.png')
# mask=tensor2pil(result)
mask=mask.convert('L')
masks.append(mask)
# rgba图
image_rgba =orig_image.convert("RGBA")
image_rgba.putalpha(mask)
rgba_images.append(image_rgba)
#rgb
rgb_image = Image.new("RGB", image_rgba.size, (0, 0, 0))
rgb_image.paste(image_rgba, mask=image_rgba.split()[3])
rgb_images.append(rgb_image)
return (masks,rgba_images,rgb_images)
def run_bg(model_name= "unet",images=[]):
# model_name = "unet" # "isnet-general-use"
rembg_session = new_session(model_name)
@@ -118,14 +643,16 @@ class RembgNode_:
def INPUT_TYPES(s):
return {"required": {
"image": ("IMAGE",),
"model_name": (["u2net",
"model_name": ([
"briarmbg",
"u2net",
"u2netp",
"u2net_human_seg",
"u2net_cloth_seg",
"silueta",
"isnet-general-use",
"isnet-anime",
# "sam"
],),
},
@@ -153,7 +680,10 @@ class RembgNode_:
im=tensor2pil(im)
images.append(im)
masks,rgba_images,rgb_images=run_bg(model_name,images)
if model_name=='briarmbg':
masks,rgba_images,rgb_images=briarmbg_run(images)
else:
masks,rgba_images,rgb_images=run_bg(model_name,images)
masks=[pil2tensor(m) for m in masks]
+125 -6
View File
@@ -7,6 +7,8 @@ import folder_paths
import matplotlib.font_manager as fm
import torch
def recursive_search(directory, excluded_dir_names=None):
if not os.path.isdir(directory):
return [], {}
@@ -198,14 +200,17 @@ class TextToNumber:
return {"required": {
"text": ("STRING",{"multiline": False,"default": "1"}),
"random_number": (["enable", "disable"],),
"number":("INT", {
"default": 0,
"min": 0, #Minimum value
"max_num":("INT", {
"default": 10,
"min":2, #Minimum value
"max": 10000000000, #Maximum value
"step": 1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
},
"optional":{
"seed": (any_type, {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
}
}
RETURN_TYPES = ("INT",)
@@ -218,7 +223,7 @@ class TextToNumber:
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
def run(self,text,random_number,number):
def run(self,text,random_number,max_num,seed=0):
numbers = re.findall(r'\d+', text)
result=0
@@ -227,7 +232,7 @@ class TextToNumber:
# print(result)
if random_number=='enable' and result>0:
result= random.randint(1, 10000000000)
result= random.randint(1, max_num)
return {"ui": {"text": [text],"num":[result]}, "result": (result,)}
@@ -425,7 +430,7 @@ class DynamicDelayProcessor:
},
"optional":{
"any_input":(any_type,),
"delay_by_text":("STRING",{"multiline":True,}),
"delay_by_text":("STRING",{"multiline":True,"dynamicPrompts": False,}),
"words_per_seconds":("FLOAT",{ "default":1.50,"min": 0.0,"max": 1000.00,"display":"Chars per second?"}),
"replace_output": (["disable","enable"],),
"replace_value":("INT",{ "default":-1,"min": 0,"max": 1000000,"display":"Replacement value"})
@@ -712,3 +717,117 @@ class TESTNODE_:
result = list_stats.count_types(ANY)
return {"ui": {"data": result,"type":[str(type(ANY[0]))]}, "result": (ANY,)}
class CreateSeedNode:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
}
}
RETURN_TYPES = ("INT",)
RETURN_NAMES = ("seed",)
OUTPUT_NODE = True
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Utils"
def run(self, seed):
return (seed,)
class CreateCkptNames:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"ckpt_names": ("STRING",{"multiline": True,"default": "\n".join(folder_paths.get_filename_list("checkpoints")),"dynamicPrompts": False}),
}
}
RETURN_TYPES = (any_type,)
RETURN_NAMES = ("ckpt_names",)
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (True,)
# OUTPUT_NODE = True
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Utils"
def run(self, ckpt_names):
ckpt_names=ckpt_names.split('\n')
ckpt_names = [name for name in ckpt_names if name.strip()]
return (ckpt_names,)
class CreateLoraNames:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"lora_names": ("STRING",{"multiline": True,"default": "\n".join(folder_paths.get_filename_list("loras")),"dynamicPrompts": False}),
}
}
RETURN_TYPES = (any_type,"STRING",)
RETURN_NAMES = ("lora_names","prompt",)
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (True,True,)
# OUTPUT_NODE = True
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Utils"
def run(self, lora_names):
lora_names=lora_names.split('\n')
lora_names = [name for name in lora_names if name.strip()]
prompts=[os.path.splitext(n)[0] for n in lora_names]
return (lora_names,prompts,)
class CreateSampler_names:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"sampler_names": ("STRING",{"multiline": True,"default": "\n".join(comfy.samplers.KSampler.SAMPLERS),"dynamicPrompts": False}),
}
}
RETURN_TYPES = (any_type,)
RETURN_NAMES = ("sampler_names",)
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (True,)
# OUTPUT_NODE = True
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Utils"
def run(self, sampler_names):
sampler_names=sampler_names.split('\n')
sampler_names = [name for name in sampler_names if name.strip()]
return (sampler_names,)
-179
View File
@@ -1,179 +0,0 @@
# https://github.com/openai/consistencydecoder/blob/main/consistencydecoder/__init__.py
import folder_paths
from comfy import model_management
import math
import torch
import numpy as np
from PIL import Image
class ConsistencyDecoderWrapper:
def __init__(self, decoder):
self.decoder = decoder
def decode(self, x):
return self.decoder(x)
def _extract_into_tensor(arr, timesteps, broadcast_shape):
# from: https://github.com/openai/guided-diffusion/blob/22e0df8183507e13a7813f8d38d51b072ca1e67c/guided_diffusion/gaussian_diffusion.py#L895 """
res = arr[timesteps].float()
dims_to_append = len(broadcast_shape) - len(res.shape)
return res[(...,) + (None,) * dims_to_append]
def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999):
# from: https://github.com/openai/guided-diffusion/blob/22e0df8183507e13a7813f8d38d51b072ca1e67c/guided_diffusion/gaussian_diffusion.py#L45
betas = []
for i in range(num_diffusion_timesteps):
t1 = i / num_diffusion_timesteps
t2 = (i + 1) / num_diffusion_timesteps
betas.append(min(1 - alpha_bar(t2) / alpha_bar(t1), max_beta))
return torch.tensor(betas)
class ConsistencyDecoder:
def __init__(self, device="cuda:0", download_target=""):
self.n_distilled_steps = 64
# download_target = _download("https://openaipublic.azureedge.net/diff-vae/c9cebd3132dd9c42936d803e33424145a748843c8f716c0814838bdc8a2fe7cb/decoder.pt", download_root)
self.ckpt = torch.jit.load(download_target).to(device)
self.device = device
sigma_data = 0.5
betas = betas_for_alpha_bar(
1024, lambda t: math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
).to(device)
alphas = 1.0 - betas
alphas_cumprod = torch.cumprod(alphas, dim=0)
self.sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)
self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod)
sqrt_recip_alphas_cumprod = torch.sqrt(1.0 / alphas_cumprod)
sigmas = torch.sqrt(1.0 / alphas_cumprod - 1)
self.c_skip = (
sqrt_recip_alphas_cumprod
* sigma_data**2
/ (sigmas**2 + sigma_data**2)
)
self.c_out = sigmas * sigma_data / (sigmas**2 + sigma_data**2) ** 0.5
self.c_in = sqrt_recip_alphas_cumprod / (sigmas**2 + sigma_data**2) ** 0.5
@staticmethod
def round_timesteps(
timesteps, total_timesteps, n_distilled_steps, truncate_start=True
):
with torch.no_grad():
space = torch.div(total_timesteps, n_distilled_steps, rounding_mode="floor")
rounded_timesteps = (
torch.div(timesteps, space, rounding_mode="floor") + 1
) * space
if truncate_start:
rounded_timesteps[rounded_timesteps == total_timesteps] -= space
else:
rounded_timesteps[rounded_timesteps == total_timesteps] -= space
rounded_timesteps[rounded_timesteps == 0] += space
return rounded_timesteps
@staticmethod
def ldm_transform_latent(z, extra_scale_factor=1):
channel_means = [0.38862467, 0.02253063, 0.07381133, -0.0171294]
channel_stds = [0.9654121, 1.0440036, 0.76147926, 0.77022034]
if len(z.shape) != 4:
raise ValueError()
z = z * 0.18215
channels = [z[:, i] for i in range(z.shape[1])]
channels = [
extra_scale_factor * (c - channel_means[i]) / channel_stds[i]
for i, c in enumerate(channels)
]
return torch.stack(channels, dim=1)
@torch.no_grad()
def __call__(
self,
features: torch.Tensor,
schedule=[1.0, 0.5],
):
features = self.ldm_transform_latent(features)
ts = self.round_timesteps(
torch.arange(0, 1024),
1024,
self.n_distilled_steps,
truncate_start=False,
)
shape = (
features.size(0),
3,
8 * features.size(2),
8 * features.size(3),
)
x_start = torch.zeros(shape, device=features.device, dtype=features.dtype)
schedule_timesteps = [int((1024 - 1) * s) for s in schedule]
for i in schedule_timesteps:
t = ts[i].item()
t_ = torch.tensor([t] * features.shape[0]).to(self.device)
noise = torch.randn_like(x_start)
x_start = (
_extract_into_tensor(self.sqrt_alphas_cumprod, t_, x_start.shape)
* x_start
+ _extract_into_tensor(
self.sqrt_one_minus_alphas_cumprod, t_, x_start.shape
)
* noise
)
c_in = _extract_into_tensor(self.c_in, t_, x_start.shape)
model_output = self.ckpt(c_in * x_start, t_, features=features)
B, C = x_start.shape[:2]
model_output, _ = torch.split(model_output, C, dim=1)
pred_xstart = (
_extract_into_tensor(self.c_out, t_, x_start.shape) * model_output
+ _extract_into_tensor(self.c_skip, t_, x_start.shape) * x_start
).clamp(-1, 1)
x_start = pred_xstart
return x_start
class VAELoader:
@classmethod
def INPUT_TYPES(s):
return {"required": { "vae_name": (folder_paths.get_filename_list("vae"), )}}
RETURN_TYPES = ("VAE",)
FUNCTION = "load_vae"
CATEGORY = "♾️Mixlab/__TEST"
#TODO: scale factor?
def load_vae(self, vae_name):
vae_path = folder_paths.get_full_path("vae", vae_name)
device = 'cuda:0'
# print('device',device)
consistencyDecoder = ConsistencyDecoder(device=device,
download_target=vae_path) # Model size: 2.49 GB
vae = ConsistencyDecoderWrapper(consistencyDecoder)
return (vae,)
class VAEDecode:
@classmethod
def INPUT_TYPES(s):
return {"required": { "samples": ("LATENT", ), "vae": ("VAE", )}}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "decode"
CATEGORY = "♾️Mixlab/__TEST"
def decode(self, vae, samples):
image = vae.decode(samples["samples"].to("cuda:0"))
image = image[0].cpu().numpy()
image = (image + 1.0) * 127.5
image = image.clip(0, 255).astype(np.uint8)
image = Image.fromarray(image.transpose(1, 2, 0))
image = image.convert("RGB")
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
return (image, )
+2 -1
View File
@@ -5,4 +5,5 @@ opencv-python-headless
matplotlib
openai
simple-lama-inpainting
clip-interrogator==0.6.0
clip-interrogator==0.6.0
transformers>=4.36.0
+5 -2
View File
@@ -643,7 +643,7 @@
return true
}
// 种子的处理
function randomSeed(seed, data) {
for (const id in data) {
if (data[id].inputs.seed != undefined
@@ -1316,6 +1316,7 @@
}
}), value);
// 选择事件绑定
selectDom.addEventListener('change', e => {
e.preventDefault();
// console.log(selectDom.value)
@@ -1550,6 +1551,7 @@
}
// 创建下拉选择
function createSelect(options, defaultValue) {
var selectElement = document.createElement("select");
selectElement.className = "select"
@@ -1558,7 +1560,7 @@
for (var i = 0; i < options.length; i++) {
var option = document.createElement("option");
option.value = options[i].value;
option.text = options[i].text;
option.innerText = options[i].text;
selectElement.appendChild(option);
// if(options[i].selected)
}
@@ -1569,6 +1571,7 @@
return selectElement
}
// 创建下拉选择 - 带说明
function createSelectWithOptions(title, options, defaultValue) {
const div = document.createElement("div");
+27 -7
View File
@@ -67,9 +67,13 @@ async function drawImageToCanvas (imageUrl) {
}
function extractInputAndOutputData (jsonData, inputIds = [], outputIds = []) {
const data = jsonData
const input = []
const output = []
// workflow
// const workflow=jsonData.workflow;
// const nodes=workflow.nodes;
const data = jsonData.output
let input = []
let output = []
const seed = {}
for (const id in data) {
@@ -77,7 +81,7 @@ function extractInputAndOutputData (jsonData, inputIds = [], outputIds = []) {
let node = app.graph.getNodeById(id)
if (inputIds.includes(id)) {
// let node = app.graph.getNodeById(id)
let options = []
let options = {}
// 模型
try {
if (node.type === 'CheckpointLoaderSimple') {
@@ -113,6 +117,15 @@ function extractInputAndOutputData (jsonData, inputIds = [], outputIds = []) {
if (node.type == 'Color') {
}
// loadImage的mask支持
if (node.type === 'LoadImage') {
let output = node.outputs.filter(ot => ot.type == 'MASK')[0]
if (output.links) {
// 有输出
options.hasMask = true
}
}
input[inputIds.indexOf(id)] = {
...data[id],
title: node.title,
@@ -138,6 +151,10 @@ function extractInputAndOutputData (jsonData, inputIds = [], outputIds = []) {
}
}
// 修复bug,当节点不存在时
input = input.filter(i => i)
output = output.filter(i => i)
return { input, output, seed }
}
@@ -211,8 +228,8 @@ async function save (json, download = false, showInfo = true) {
try {
let data = await app.graphToPrompt()
const { input, output, seed } = extractInputAndOutputData(
data.output,
let { input, output, seed } = extractInputAndOutputData(
data,
inputIds,
outputIds
)
@@ -399,7 +416,7 @@ app.registerExtension({
}
},
async loadedGraphNode (node, app) {
console.log('#loadedGraphNode1111')
// console.log('#loadedGraphNode1111')
window._mixlab_app_json = null //切换workflow需要清空
if (node.type === 'AppInfo') {
let auto_save = node.widgets.filter(w => w.name == 'auto_save')[0]
@@ -408,6 +425,9 @@ app.registerExtension({
auto_save.value = 'enable'
}
}
// app.canvas.centerOnNode(node)
// app.canvas.setZoom(0.45)
}
}
})
+1 -1
View File
@@ -3,7 +3,7 @@ import { app } from '../../../scripts/app.js'
const repoOwner = 'shadowcz007' // 替换为仓库的所有者
const repoName = 'comfyui-mixlab-nodes' // 替换为仓库的名称
const version = 'v0.14.0'
const version = 'v0.16.0'
fetch(`https://api.github.com/repos/${repoOwner}/${repoName}/releases/latest`)
.then(response => response.json())
+71 -65
View File
@@ -68,7 +68,7 @@ app.registerExtension({
size: [128, 32], // a default size
draw (ctx, node, width, y) {},
computeSize (...args) {
return [128,32] // a method to compute the current size of the widget
return [128, 32] // a method to compute the current size of the widget
},
async serializeValue (nodeId, widgetIndex) {
let data = getLocalData('_mixlab_api_key')
@@ -203,76 +203,82 @@ app.registerExtension({
app.registerExtension({
name: 'Mixlab.GPT.ShowTextForGPT',
async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (nodeData.name === "ShowTextForGPT") {
function populate(text) {
if (this.widgets) {
const pos = this.widgets.findIndex((w) => w.name === "text");
if (pos !== -1) {
for (let i = pos; i < this.widgets.length; i++) {
this.widgets[i].onRemove?.();
}
this.widgets.length = pos;
}
}
// console.log('ShowTextForGPT',text)
for (let list of text) {
const w = ComfyWidgets["STRING"](this, "text", ["STRING", { multiline: true }], app).widget;
w.inputEl.readOnly = true;
w.inputEl.style.opacity = 0.6;
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeData.name === 'ShowTextForGPT') {
function populate (text) {
text = text.filter(t => t && t?.trim())
try {
let data=JSON.parse(list);
data=Array.from(data,d=>{
return {
...d,
content:decodeURIComponent(d.content)
}
})
list=JSON.stringify(data,null,2)
} catch (error) {
// console.log(error)
if (this.widgets) {
// const pos = this.widgets.findIndex(w => w.name === 'text')
for (let i = 0; i < this.widgets.length; i++) {
if (this.widgets[i].name == 'show_text') this.widgets[i].onRemove?.()
}
this.widgets.length = 1
}
// console.log('ShowTextForGPT',text)
for (let list of text) {
if (list) {
// console.log('#####', list)
const w = ComfyWidgets['STRING'](
this,
'show_text',
['STRING', { multiline: true }],
app
).widget
w.inputEl.readOnly = true
w.inputEl.style.opacity = 0.6
w.value =list;
}
try {
if (typeof list != 'string') {
let data = JSON.parse(list)
data = Array.from(data, d => {
return {
...d,
content: decodeURIComponent(d.content)
}
})
list = JSON.stringify(data, null, 2)
}
} catch (error) {
console.log(error)
}
w.value = list
}
}
// console.log('ShowTextForGPT',this.widgets.length)
requestAnimationFrame(() => {
const sz = this.computeSize();
if (sz[0] < this.size[0]) {
sz[0] = this.size[0];
}
if (sz[1] < this.size[1]) {
sz[1] = this.size[1];
}
this.onResize?.(sz);
app.graph.setDirtyCanvas(true, false);
});
}
requestAnimationFrame(() => {
if (this) {
const sz = this.computeSize()
if (sz[0] < this.size[0]) {
sz[0] = this.size[0]
}
if (sz[1] < this.size[1]) {
sz[1] = this.size[1]
}
this.onResize?.(sz)
app.graph.setDirtyCanvas(true, false)
}
})
}
// When the node is executed we will be sent the input text, display this in the widget
const onExecuted = nodeType.prototype.onExecuted;
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments);
console.log('##',message.text)
populate.call(this, message.text);
};
// When the node is executed we will be sent the input text, display this in the widget
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments)
// console.log('##onExecuted', this, message)
if (message.text) populate.call(this, message.text)
}
const onConfigure = nodeType.prototype.onConfigure;
nodeType.prototype.onConfigure = function () {
onConfigure?.apply(this, arguments);
if (this.widgets_values?.length) {
populate.call(this, this.widgets_values);
}
};
const onConfigure = nodeType.prototype.onConfigure
nodeType.prototype.onConfigure = function () {
onConfigure?.apply(this, arguments)
if (this.widgets_values?.length) {
populate.call(this, this.widgets_values)
}
}
this.serialize_widgets = true //需要保存参数
}
},
}
}
})
+38 -6
View File
@@ -156,6 +156,32 @@ const parseSvg = async svgContent => {
return { data, image: base64, svgElement }
}
function findImages(nodeId) {
// 检查当前节点是否有 imgs 字段
const n = app.graph.getNodeById(nodeId)
if (n.imgs) {
return n.imgs;
}
// 检查当前节点的 inputs 是否有 image 字段
if (n.inputs) {
for (let i = 0; i < n.inputs.length; i++) {
if (n.inputs[i].name==='image'||n.inputs[i].name==='images') {
// 获取新的 nodeId,并递归调用 findImages 函数
var linkId = n.inputs[i]?.link;
var origin_id = app.graph.links[linkId].origin_id
return findImages(origin_id);
}
}
}
// 如果没有找到 imgs 字段或者 image 字段,则返回 null
return null;
}
async function setArea (cw, ch, topBase64, base64, data, fn) {
let displayHeight = Math.round(window.screen.availHeight * 0.8)
let div = document.createElement('div')
@@ -571,15 +597,19 @@ app.registerExtension({
}
}
try {
console.log('this.inputs', this.inputs)
let topLinkId = this.inputs[0].link
let topNodeId = app.graph.links[topLinkId].origin_id
let topIm = app.graph.getNodeById(topNodeId).imgs[0]
console.log('this.inputs', this.id)
let imgs=findImages(this.id)
// let topLinkId = this.inputs[0].link
// let topNodeId = app.graph.links[topLinkId].origin_id
let topIm = imgs[0]
let linkId = this.inputs[3].link
let nodeId = app.graph.links[linkId].origin_id
// console.log(linkId,this.inputs)
let im = app.graph.getNodeById(nodeId).imgs[0]
let imgs2=findImages(nodeId)
let im = imgs2[0]
console.log(topIm,im)
// let src = im.src
setArea(
im.naturalWidth,
@@ -589,7 +619,9 @@ app.registerExtension({
data,
updateValue
)
} catch (error) {}
} catch (error) {
console.log(error)
}
})
}
}
+302
View File
@@ -0,0 +1,302 @@
const smart_connect_config_input = [
{
node_type: 'CLIPTextEncode',
node_widget_name: 'text',
inputNodeName: 'RandomPrompt',
inputNode_output_name: 'STRING'
},
{
node_type: 'CLIPTextEncode',
node_widget_name: 'text',
inputNodeName: 'EmbeddingPrompt',
inputNode_output_name: 'STRING'
},
{
node_type: 'CLIPTextEncode',
node_widget_name: 'text',
inputNodeName: 'ChinesePrompt_Mix',
inputNode_output_name: 'prompt'
},
{
node_type: 'CheckpointLoaderSimple',
node_widget_name: 'ckpt_name',
inputNodeName: 'CkptNames_',
inputNode_output_name: 'ckpt_names'
},
{
node_type: 'KSampler',
node_widget_name: 'sampler_name',
inputNodeName: 'SamplerNames_',
inputNode_output_name: 'sampler_names'
},
{
node_type: 'LoraLoaderModelOnly',
node_widget_name: 'lora_name',
inputNodeName: 'LoraNames_',
inputNode_output_name: 'lora_names'
},
{
node_type: 'LoadLoRA',
node_widget_name: 'lora_name',
inputNodeName: 'LoraNames_',
inputNode_output_name: 'lora_names'
},
{
node_type: 'Moondream',
node_widget_name: 'image',
inputNodeName: 'LoadImage',
inputNode_output_name: 'IMAGE'
}
]
const smart_connect_config_output = [
{
node_type: 'LoadImage',
node_output_name: 'IMAGE',
outputNodeName: 'ClipInterrogator',
outputNode_input_name: 'image'
},
{
node_type: 'VAEDecode',
node_output_name: 'IMAGE',
outputNodeName: 'PromptImage',
outputNode_input_name: 'images'
},
{
node_type: 'VAEDecode',
node_output_name: 'IMAGE',
outputNodeName: 'PreviewImage',
outputNode_input_name: 'images'
},
{
node_type: 'VAEDecode',
node_output_name: 'IMAGE',
outputNodeName: 'SaveImage',
outputNode_input_name: 'images'
},
{
node_type: 'Moondream',
node_output_name: 'STRING',
outputNodeName: 'ShowTextForGPT',
outputNode_input_name: 'text'
}
]
// import {
// convertToInput,
// getConfig,
// isConvertableWidget
// } from '../../../extensions/core/widgetInputs.js'
const CONVERTED_TYPE = 'converted-widget'
const GET_CONFIG = Symbol()
function getConfig (widgetName) {
const { nodeData } = this.constructor
return (
nodeData?.input?.required[widgetName] ??
nodeData?.input?.optional?.[widgetName]
)
}
function hideWidget (node, widget, suffix = '') {
widget.origType = widget.type
widget.origComputeSize = widget.computeSize
widget.origSerializeValue = widget.serializeValue
widget.computeSize = () => [0, -4] // -4 is due to the gap litegraph adds between widgets automatically
widget.type = CONVERTED_TYPE + suffix
widget.serializeValue = () => {
// Prevent serializing the widget if we have no input linked
if (!node.inputs) {
return undefined
}
let node_input = node.inputs.find(i => i.widget?.name === widget.name)
if (!node_input || !node_input.link) {
return undefined
}
return widget.origSerializeValue
? widget.origSerializeValue()
: widget.value
}
// Hide any linked widgets, e.g. seed+seedControl
if (widget.linkedWidgets) {
for (const w of widget.linkedWidgets) {
hideWidget(node, w, ':' + widget.name)
}
}
}
function convertToInput (node, widget, config) {
hideWidget(node, widget)
const type = config[0]
// Add input and store widget config for creating on primitive node
const sz = node.size
node.addInput(widget.name, type, {
widget: { name: widget.name, [GET_CONFIG]: () => config }
})
for (const widget of node.widgets) {
widget.last_y += LiteGraph.NODE_SLOT_HEIGHT
}
// Restore original size but grow if needed
node.setSize([Math.max(sz[0], node.size[0]), Math.max(sz[1], node.size[1])])
}
export function smart_init () {
LGraphCanvas.prototype._createNodeForInput = function (
node,
widget,
inputNodeName,
inputNode_slot
) {
// console.log(node.pos)
// var widget = node.widgets.filter(w => w.name === node_widget_name)[0]
if (widget) {
// 如果有存在的,没有连线输出的,自动连,不新建
let input_node = null
Array.from(app.graph.findNodesByType(inputNodeName), n => {
var links = n.outputs.filter(o => o.name === inputNode_slot)[0].links
// console.log(links)
if (!links || links?.length === 0) input_node = n
})
// 新建
if (!input_node) {
input_node = LiteGraph.createNode(inputNodeName)
input_node.pos = [node.pos[0] - node.size[0] - 24, node.pos[1] - 48]
app.canvas.graph.add(input_node, false)
} else {
input_node.pos = [node.pos[0] - node.size[0] - 24, node.pos[1] - 48]
}
const config = getConfig.call(node, widget.name) ?? [
widget.type,
widget.options || {}
]
let node_slotType = config[0]
// 如果input没有,则创建
if (!node.inputs?.filter(inp => inp.name === widget.name)[0]||!node.inputs)
convertToInput(node, widget, config)
input_node.connectByType(inputNode_slot, node, node_slotType)
}
}
LGraphCanvas.prototype._createNodeForOutput = function (
node,
widget,
outputNodeName,
outputNode_slot
) {
if (widget) {
let output_node
Array.from(app.graph.findNodesByType(outputNodeName), n => {
var links = n.inputs.filter(o => o.name === outputNode_slot)[0].links
// console.log(links)
if (!links || links?.length === 0) output_node = n
})
console.log('output_node', output_node, widget.name)
if (!output_node) {
// 新建
output_node = LiteGraph.createNode(outputNodeName)
output_node.pos = [node.pos[0] + node.size[0] + 24, node.pos[1] - 48]
app.canvas.graph.add(output_node, false)
} else {
output_node.pos = [node.pos[0] + node.size[0] + 24, node.pos[1] - 48]
}
const config = getConfig.call(node, widget.name) ?? [
widget.type,
widget.options || {}
]
let node_slotType = config[0]
console.log(node_slotType, output_node, outputNode_slot)
let type = output_node.inputs.filter(
inp => inp.name == outputNode_slot
)[0].type
node.connectByType(node_slotType, output_node, type)
}
}
}
export function addSmartMenu (options, node) {
let sopts = []
for (const sc of smart_connect_config_input) {
// 有智能推荐,则出现
if (node.type === sc.node_type) {
// console.log('smart',node)
// 则出现 randomPrompt
// CLIPTextEncode 的widget ,name== 'text'
let node_widget_name = sc.node_widget_name
let widget = node.widgets.filter(w => w.name === node_widget_name)[0]
if (!widget) {
// 控件没有,则查找inputs
widget = node.inputs.filter(w => w.name === node_widget_name)[0]
}
let isLinkNull = true
// 如果input里已经有,但是link为空
if (node.inputs?.filter(inp => inp.name === node_widget_name)[0]) {
isLinkNull =
node.inputs.filter(inp => inp.name === node_widget_name)[0].link ===
null
}
if (widget && isLinkNull) {
sopts.push({
content: sc.inputNodeName.split('_')[0] + '➡️',
callback: () => {
LGraphCanvas.prototype._createNodeForInput(
node, //当前node
widget, //当前node里需要自动连线的widget
sc.inputNodeName, //作为input的node type
sc.inputNode_output_name // 作为input的node的outputs的name. the input slot type of the target node
)
}
})
}
}
}
for (const sc of smart_connect_config_output) {
if (node.type === sc.node_type) {
let node_output_name = sc.node_output_name
const widget = node.outputs.filter(w => w.name === node_output_name)[0]
let isLinkNull = true
// 如果output里 link为空
if (node.outputs?.filter(inp => inp.name === node_output_name)[0]) {
isLinkNull =
node.outputs.filter(inp => inp.name === node_output_name)[0].links
?.length === 0
if (!node.outputs.filter(inp => inp.name === node_output_name)[0].links)
isLinkNull = true
}
if (widget && isLinkNull) {
sopts.push({
content: '➡️' + sc.outputNodeName.split('_')[0],
callback: () => {
LGraphCanvas.prototype._createNodeForOutput(
node, //当前node
widget, //当前node里需要自动连线的widget
sc.outputNodeName, //作为input的node type
sc.outputNode_input_name // 作为input的node的outputs的name. the input slot type of the target node
)
}
})
}
}
}
if (sopts.length > 0) options = [...sopts, null, ...options]
return options
}
+379 -24
View File
@@ -1,14 +1,247 @@
import { app } from '../../../scripts/app.js'
import { api } from '../../../scripts/api.js'
import { ComfyWidgets } from '../../../scripts/widgets.js'
import { $el } from '../../../scripts/ui.js'
import { closeIcon } from './svg_icons.js'
import { api } from '../../../scripts/api.js'
import {
GroupNodeConfig,
GroupNodeHandler
} from '../../../extensions/core/groupNode.js'
import { smart_init, addSmartMenu } from './smart_connect.js'
let isScriptLoaded = {};
function loadExternalScript(url) {
return new Promise((resolve, reject) => {
if (isScriptLoaded[url]) {
resolve();
return;
}
const script = document.createElement('script');
script.src = url;
script.onload = () => {
isScriptLoaded[url]= true;
resolve();
};
script.onerror = reject;
document.head.appendChild(script);
});
}
//
function createChart(chartDom,nodes){
var myChart = echarts.init(chartDom);
var option;
console.log(nodes)
option = {
series: [
{
type: 'treemap',
data: [
{
name: 'nodeA',
value: 10,
children: Array.from(nodes,n=>{
return {
name:n.type,
value:n.count
}
})
},
]
}
]
};
option && myChart.setOption(option);
}
async function createNodesCharts () {
await loadExternalScript('/extensions/comfyui-mixlab-nodes/lib/echarts.min.js')
const templates = await loadTemplate()
var nodes = {}
Array.from(templates, t => {
let j = JSON.parse(t.data)
for (let node of j.nodes) {
if (!nodes[node.type]) nodes[node.type] = { type: node.type, count: 0 }
nodes[node.type].count++
}
})
nodes = Object.values(nodes).sort((a, b) => b.count - a.count)
const menu = document.querySelector('.comfy-menu')
const separator = document.createElement('div')
separator.style = `margin: 20px 0px;
width: 100%;
height: 1px;
background: var(--border-color);
`
menu.append(separator)
const appsButton = document.createElement('button')
appsButton.textContent = 'Nodes';
appsButton.onclick = () => {
let div = document.querySelector('#mixlab_apps')
if (!div) {
div = document.createElement('div')
div.id = 'mixlab_apps'
document.body.appendChild(div)
let btn = document.createElement('div')
btn.style = `display: flex;
width: calc(100% - 24px);
justify-content: space-between;
align-items: center;
padding: 0 12px;
height: 44px;`
let btnB = document.createElement('button')
let textB = document.createElement('p')
btn.appendChild(textB)
btn.appendChild(btnB)
textB.style.fontSize = '12px'
textB.innerText = `Nodes`
btnB.style = `float: right; border: none; color: var(--input-text);
background-color: var(--comfy-input-bg); border-color: var(--border-color);cursor: pointer;`
btnB.addEventListener('click', () => {
div.style.display = 'none'
})
btnB.innerText = 'X'
// 悬浮框拖动事件
div.addEventListener('mousedown', function (e) {
var startX = e.clientX
var startY = e.clientY
var offsetX = div.offsetLeft
var offsetY = div.offsetTop
function moveBox (e) {
var newX = e.clientX
var newY = e.clientY
var deltaX = newX - startX
var deltaY = newY - startY
div.style.left = offsetX + deltaX + 'px'
div.style.top = offsetY + deltaY + 'px'
localStorage.setItem(
'mixlab_app_pannel',
JSON.stringify({ x: div.style.left, y: div.style.top })
)
}
function stopMoving () {
document.removeEventListener('mousemove', moveBox)
document.removeEventListener('mouseup', stopMoving)
}
document.addEventListener('mousemove', moveBox)
document.addEventListener('mouseup', stopMoving)
})
div.appendChild(btn)
let chartDom = document.createElement('div')
chartDom.style=`height:80vh;width:450px`
chartDom.className='chart'
div.appendChild(chartDom)
}
if (div.style.display == 'flex') {
div.style.display = 'none'
} else {
let pos = JSON.parse(
localStorage.getItem('mixlab_app_pannel') ||
JSON.stringify({ x: 0, y: 0 })
)
div.style = `
flex-direction: column;
align-items: end;
display:flex;
position: absolute;
top: ${pos.y}; left: ${pos.x}; width: 450px;
color: var(--descrip-text);
background-color: var(--comfy-menu-bg);
padding: 10px;
border: 1px solid black;z-index: 999999999;padding-top: 0;`
};
createChart(div.querySelector('.chart'),nodes)
}
menu.append(appsButton)
}
function copyNodeValues (src, dest) {
// title
dest.title = src.title
// copy input connections
for (let i in src.inputs) {
let input = src.inputs[i]
if (input.link) {
let link = app.graph.links[input.link]
let src_node = app.graph.getNodeById(link.origin_id)
if (dest.inputs.filter(inp => inp.name === input.name).length === 0) {
// 没有,name换了
let dInp = dest.inputs.filter(inp => inp.type === input.type)
if (dInp.length === 1) {
src_node.connect(link.origin_slot, dest.id, dInp[0].name)
}
} else {
src_node.connect(link.origin_slot, dest.id, input.name)
}
}
}
// copy output connections
let output_links = {}
for (let i in src.outputs) {
let output = src.outputs[i]
if (output.links) {
let links = []
for (let j in output.links) {
links.push(app.graph.links[output.links[j]])
}
output_links[output.name] = links
}
}
for (let i in dest.outputs) {
let links = output_links[dest.outputs[i].name]
if (links) {
for (let j in links) {
let link = links[j]
let target_node = app.graph.getNodeById(link.target_id)
dest.connect(parseInt(i), target_node, link.target_slot)
}
}
}
// copy widgets
for (const w of src.widgets) {
for (const d of dest.widgets) {
if (w.name === d.name) {
d.value = w.value
}
}
}
app.graph.afterChange()
}
function deepEqual (obj1, obj2) {
if (typeof obj1 !== typeof obj2) {
return false
@@ -558,6 +791,39 @@ function createModal (url, markdown, title) {
div.appendChild(bgElement)
}
const loadTemplate = async () => {
const id = 'Comfy.NodeTemplates'
const file = 'comfy.templates.json'
let templates = []
if (app.storageLocation === 'server') {
if (app.isNewUserSession) {
// New user so migrate existing templates
const json = localStorage.getItem(id)
if (json) {
templates = JSON.parse(json)
}
await api.storeUserData(file, json, { stringify: false })
} else {
const res = await api.getUserData(file)
if (res.status === 200) {
try {
templates = await res.json()
} catch (error) {}
} else if (res.status !== 404) {
console.error(res.status + ' ' + res.statusText)
}
}
} else {
const json = localStorage.getItem(id)
if (json) {
templates = JSON.parse(json)
}
}
return templates ?? []
}
app.registerExtension({
name: 'Comfy.Mixlab.ui',
init () {
@@ -575,22 +841,71 @@ app.registerExtension({
}
}
LGraphCanvas.prototype.fixTheNode = function (node) {
let new_node = LiteGraph.createNode(node.comfyClass)
new_node.pos = [node.pos[0], node.pos[1]]
app.canvas.graph.add(new_node, false)
copyNodeValues(node, new_node)
app.canvas.graph.remove(node)
}
smart_init()
const getNodeMenuOptions = LGraphCanvas.prototype.getNodeMenuOptions // store the existing method
LGraphCanvas.prototype.getNodeMenuOptions = function (node) {
// replace it
const options = getNodeMenuOptions.apply(this, arguments) // start by calling the stored one
node.setDirtyCanvas(true, true) // force a redraw of (foreground, background)
console.log('getNodeMenuOptions', node.type == 'CLIPTextEncode')
return [
let opts = [
{
content: 'Help ♾️Mixlab', // with a name
callback: () => {
LGraphCanvas.prototype.helpAboutNode(node)
} // and the callback
},
null,
...options
] // and return the options
{
content: 'Fix node v2', // with a name
callback: () => {
LGraphCanvas.prototype.fixTheNode(node)
}
}
]
opts = addSmartMenu(opts, node)
// if (node.type == 'CLIPTextEncode') {
// // 则出现 randomPrompt
// // CLIPTextEncode 的widget ,name== 'text'
// let node_widget_name = 'text'
// const widget = node.widgets.filter(w => w.name === node_widget_name)[0]
// let mixlab_nodes_smart_connect= [{node_type:'CLIPTextEncode',
// node_widget_name:'text',
// inputNodeName:'RandomPrompt',
// inputNode_output_type:'STRING'}]
// if (widget) {
// opts = [
// {
// content: 'RandomPrompt',
// callback: () => {
// LGraphCanvas.prototype._createNodeForInput(
// node, //当前node
// widget,//当前node里需要自动连线的widget
// 'RandomPrompt',//作为input的node type
// 'STRING'// 作为input的node的outputs的type. the input slot type of the target node
// )
// }
// },
// null,
// ...opts
// ]
// }
// }
return [...opts, null, ...options] // and return the options
}
const getGroupMenuOptions = LGraphCanvas.prototype.getGroupMenuOptions // store the existing method
@@ -599,16 +914,6 @@ app.registerExtension({
const options = getGroupMenuOptions.apply(this, arguments) // start by calling the stored one
node.setDirtyCanvas(true, true) // force a redraw of (foreground, background)
// templete
const key = 'Comfy.NodeTemplates'
let templates = localStorage.getItem(key)
if (templates) {
templates = JSON.parse(templates)
} else {
templates = []
}
const store = () => localStorage.setItem(key, JSON.stringify(templates))
return [
{
content: 'Clone Group ♾️Mixlab', // with a name
@@ -666,7 +971,7 @@ app.registerExtension({
localStorage.setItem('litegrapheditor_clipboard', old)
}
clipboardAction(() => {
clipboardAction(async () => {
let name = group.title + ' ♾️Mixlab'
let nodes = group._nodes
@@ -679,6 +984,8 @@ app.registerExtension({
const nodeData = node.serialize()
let groupData = GroupNodeHandler.getGroupData(node)
// console.log('groupData',GroupNodeHandler.isGroupNode(node),groupData)
if (groupData) {
groupData = groupData.nodeData
if (!data.groupNodes) {
@@ -689,14 +996,46 @@ app.registerExtension({
}
}
templates.push({
// templete
const store = async nt => {
const id = 'Comfy.NodeTemplates'
const file = 'comfy.templates.json'
let templates = await loadTemplate()
templates.push(nt)
if (app.storageLocation === 'server') {
const ts = JSON.stringify(templates, undefined, 4)
localStorage.setItem(id, ts) // Backwards compatibility
try {
await api.storeUserData(file, ts, {
stringify: false
})
} catch (error) {
console.error(error)
alert(error.message)
}
} else {
localStorage.setItem(id, JSON.stringify(templates))
}
}
console.log('data', data)
store({
name,
data: JSON.stringify(data)
})
store()
})
} // and the callback
},
{
content: `Remove Group&Nodes ♾️Mixlab`, // with a name
callback: async (value, opts, e, menu, group) => {
// console.log(group)
let nodes = group._nodes
for (const node of nodes) {
app.graph.remove(node)
}
app.graph.remove(group)
} // and the callback
},
null,
...options
] // and return the options
@@ -722,7 +1061,7 @@ app.registerExtension({
const apps = await get_my_app()
let apps_map = { '0': [] }
let apps_map = { 0: [] }
for (const app of apps) {
if (app.category) {
@@ -735,7 +1074,7 @@ app.registerExtension({
let apps_opts = []
for (const category in apps_map) {
console.log('category',typeof(category))
console.log('category', typeof category)
if (category === '0') {
apps_opts.push(
...Array.from(apps_map[category], a => {
@@ -764,7 +1103,7 @@ app.registerExtension({
} else {
// 二级
apps_opts.push({
content: '🚀 '+category,
content: '🚀 ' + category,
has_submenu: true,
disabled: false,
submenu: {
@@ -1047,7 +1386,7 @@ app.registerExtension({
has_submenu: true,
disabled: false,
submenu: {
options:apps_opts
options: apps_opts
}
}
)
@@ -1055,5 +1394,21 @@ app.registerExtension({
return options
}
}, 1000)
// createNodesCharts()
},
async loadedGraphNode (node, app) {
// console.log(
// '#ui init',
// app.graph._nodes[app.graph._nodes.length - 1].id,
// node.id
// )
try {
// 用来居中显示节点
if ((app.graph._nodes[app.graph._nodes.length - 1].id, node.id)) {
app.canvas.centerOnNode(node)
app.canvas.setZoom(0.45)
}
} catch (error) {}
}
})
+1 -4
View File
@@ -1,8 +1,5 @@
import { app } from '../../../scripts/app.js'
import { api } from '../../../scripts/api.js'
import { ComfyWidgets } from '../../../scripts/widgets.js'
import { $el } from '../../../scripts/ui.js'
import { addValueControlWidget } from '../../../scripts/widgets.js'
import { $el } from '../../../scripts/ui.js'
const getLocalData = key => {
let data = {}
+558
View File
@@ -0,0 +1,558 @@
{
"last_node_id": 69,
"last_link_id": 72,
"nodes": [
{
"id": 37,
"type": "CLIPTextEncode",
"pos": [
6705,
-216
],
"size": {
"0": 400,
"1": 200
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 33
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
34
],
"shape": 3
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"beautiful scenery nature glass bottle landscape, , purple galaxy bottle,"
]
},
{
"id": 5,
"type": "CLIPTextEncode",
"pos": [
6693,
61
],
"size": {
"0": 425.27801513671875,
"1": 180.6060791015625
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 6
},
{
"name": "text",
"type": "STRING",
"link": 25,
"widget": {
"name": "text"
}
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
3
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"text, watermark"
]
},
{
"id": 3,
"type": "EmptyLatentImage",
"pos": [
6689,
306
],
"size": {
"0": 315,
"1": 106
},
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
4
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EmptyLatentImage"
},
"widgets_values": [
512,
512,
1
]
},
{
"id": 27,
"type": "EmbeddingPrompt",
"pos": [
6104,
23
],
"size": {
"0": 399.6408996582031,
"1": 82
},
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "STRING",
"type": "STRING",
"links": [
25
],
"shape": 3
}
],
"properties": {
"Node name for S&R": "EmbeddingPrompt"
},
"widgets_values": [
"negative-embed-verybadimagenegative_v1.3",
1
]
},
{
"id": 67,
"type": "VAEDecode",
"pos": [
7551,
-184
],
"size": {
"0": 210,
"1": 46
},
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "samples",
"type": "LATENT",
"link": 70
},
{
"name": "vae",
"type": "VAE",
"link": 69
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
71
],
"shape": 3
}
],
"properties": {
"Node name for S&R": "VAEDecode"
}
},
{
"id": 1,
"type": "KSampler",
"pos": [
7187,
-145
],
"size": {
"0": 315,
"1": 262
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 1
},
{
"name": "positive",
"type": "CONDITIONING",
"link": 34
},
{
"name": "negative",
"type": "CONDITIONING",
"link": 3
},
{
"name": "latent_image",
"type": "LATENT",
"link": 4
}
],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
70
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "KSampler"
},
"widgets_values": [
408562451564429,
"fixed",
20,
8,
"euler",
"normal",
1
]
},
{
"id": 61,
"type": "PromptImage",
"pos": [
7853,
-238
],
"size": [
465.6378949342379,
760.4568424013569
],
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 71
},
{
"name": "prompts",
"type": "STRING",
"link": 72,
"widget": {
"name": "prompts"
}
}
],
"properties": {
"Node name for S&R": "PromptImage"
},
"widgets_values": [
"",
"disable",
{
"_images": [
[
{
"filename": "mixlab_PromptImage_0_00027_.png",
"subfolder": "",
"type": "output"
}
],
[
{
"filename": "mixlab_PromptImage_1_00028_.png",
"subfolder": "",
"type": "output"
}
],
[
{
"filename": "mixlab_PromptImage_2_00029_.png",
"subfolder": "",
"type": "output"
}
],
[
{
"filename": "mixlab_PromptImage_3_00030_.png",
"subfolder": "",
"type": "output"
}
],
[
{
"filename": "mixlab_PromptImage_4_00031_.png",
"subfolder": "",
"type": "output"
}
],
[
{
"filename": "mixlab_PromptImage_5_00032_.png",
"subfolder": "",
"type": "output"
}
]
],
"prompts": [
"512-inpainting-ema.safetensors",
"SSD-1B.safetensors",
"awportrait_v12.safetensors",
"cardosAnime_v20.safetensors",
"deliberate_v2.safetensors",
"gameIconInstitute_v40.safetensors"
]
}
]
},
{
"id": 2,
"type": "CheckpointLoaderSimple",
"pos": [
6105,
-139
],
"size": {
"0": 315,
"1": 98
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "ckpt_name",
"type": [
"512-inpainting-ema.safetensors",
"SSD-1B.safetensors",
"awportrait_v12.safetensors",
"cardosAnime_v20.safetensors",
"deliberate_v2.safetensors",
"gameIconInstitute_v40.safetensors",
"illuminatiDiffusionV1_v11-unclip-h-fp16.safetensors",
"sd_xl_turbo_1.0_fp16.safetensors",
"svd.safetensors"
],
"link": 64,
"widget": {
"name": "ckpt_name"
}
}
],
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
1
],
"slot_index": 0
},
{
"name": "CLIP",
"type": "CLIP",
"links": [
6,
33
],
"slot_index": 1
},
{
"name": "VAE",
"type": "VAE",
"links": [
69
],
"slot_index": 2
}
],
"properties": {
"Node name for S&R": "CheckpointLoaderSimple"
},
"widgets_values": [
"deliberate_v2.safetensors"
]
},
{
"id": 56,
"type": "CkptNames_",
"pos": [
7860,
-494
],
"size": {
"0": 400,
"1": 200
},
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "ckpt_names",
"type": "*",
"links": [
64,
72
],
"shape": 6,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CkptNames_"
},
"widgets_values": [
"512-inpainting-ema.safetensors\nSSD-1B.safetensors\nawportrait_v12.safetensors\ncardosAnime_v20.safetensors\ndeliberate_v2.safetensors\ngameIconInstitute_v40.safetensors"
]
}
],
"links": [
[
1,
2,
0,
1,
0,
"MODEL"
],
[
3,
5,
0,
1,
2,
"CONDITIONING"
],
[
4,
3,
0,
1,
3,
"LATENT"
],
[
6,
2,
1,
5,
0,
"CLIP"
],
[
25,
27,
0,
5,
1,
"STRING"
],
[
33,
2,
1,
37,
0,
"CLIP"
],
[
34,
37,
0,
1,
1,
"CONDITIONING"
],
[
64,
56,
0,
2,
0,
[
"512-inpainting-ema.safetensors",
"SSD-1B.safetensors",
"awportrait_v12.safetensors",
"cardosAnime_v20.safetensors",
"deliberate_v2.safetensors",
"gameIconInstitute_v40.safetensors",
"illuminatiDiffusionV1_v11-unclip-h-fp16.safetensors",
"sd_xl_turbo_1.0_fp16.safetensors",
"svd.safetensors"
]
],
[
69,
2,
2,
67,
1,
"VAE"
],
[
70,
1,
0,
67,
0,
"LATENT"
],
[
71,
67,
0,
61,
0,
"IMAGE"
],
[
72,
56,
0,
61,
1,
"STRING"
]
],
"groups": [],
"config": {},
"extra": {},
"version": 0.4
}