Compare commits

..
71 Commits
Author SHA1 Message Date
shadowcz007 117d58c58e v0.24.0 2024-05-02 12:06:19 +08:00
shadowcz007 ff6626ed89 add llama.cpp & Local LLM -Phi-3 & llama3
Phi-3
llama3
2024-05-02 12:02:47 +08:00
shadowcz007 10c798a440 Update index.html 2024-05-02 10:54:20 +08:00
shadowcz007 9097d87819 Update index.html 2024-05-02 10:46:09 +08:00
shadowcz007 d901f503d1 fixbug 2024-05-02 10:06:23 +08:00
shadowcz007 65f2b6ce6f update https 2024-05-02 09:07:28 +08:00
shadowcz007 d949fe8bf1 Update __init__.py 2024-05-01 22:26:46 +08:00
shadowcz007 15cfb48550 Update __init__.py 2024-05-01 20:52:50 +08:00
shadowcz007 447dc6d4c3 fixbug 2024-05-01 15:55:43 +08:00
shadowcz007 c92e43b920 Update checkVersion_mixlab.js 2024-04-29 00:16:55 +08:00
shadowcz007 be6c32b0e0 fixbug 2024-04-29 00:15:47 +08:00
shadowcz007 3d2062e810 add TripoSRModel 2024-04-29 00:15:32 +08:00
shadowcz007 dd816e95cd Update README.md 2024-04-27 23:43:40 +08:00
shadowcz007 6d1b51890d Update checkVersion_mixlab.js 2024-04-27 23:38:00 +08:00
shadowcz007 11f03ec99a 支持正片叠底 2024-04-25 21:23:40 +08:00
shadowcz007 bd192f43e7 优化 2024-04-24 23:10:03 +08:00
shadowcz007 36e4b11983 gridout can export mask 2024-04-23 12:14:50 +08:00
shadowcz007 97397ba8c2 Update ImageNode.py 2024-04-22 12:32:48 +08:00
shadowcz007 052eee4111 Update __init__.py 2024-04-22 12:31:40 +08:00
shadowcz007 e319496044 v0.22.0
- 优化ImageColorTransfer
- 支持动态提示
- 添加更多节点支持,AppInfo自动填充id
- LoadImagesToBatch加载
- zhipuai 按需安装
2024-04-22 09:52:45 +08:00
shadowcz007 b83b63c362 添加支持的节点自动填充id 2024-04-21 21:51:07 +08:00
shadowcz007 4d6b1675bb Load Images to Batch加载文件最大宽度1024 2024-04-21 21:50:56 +08:00
shadowcz007 13110fab39 支持动态提示 2024-04-21 20:58:44 +08:00
shadowcz007 74ea509848 fixbug 2024-04-20 23:05:59 +08:00
shadowcz007 192bff9d2c optimize ImageColorTransfer and support batching, 2024-04-20 22:23:34 +08:00
shadowcz007 44ed8812dc output add TransparentImage 2024-04-19 22:27:53 +08:00
shadowcz007 200696ba21 Update ui_mixlab.js 2024-04-19 20:44:27 +08:00
shadowcz007 51cf3b0c04 install Zhipuai as needed 2024-04-19 11:41:56 +08:00
shadowcz007 6a4831c83b add nodes map for appinfo 2024-04-18 18:24:58 +08:00
shadowcz007 c5e7ed95a3 Update index.html 2024-04-18 16:10:02 +08:00
shadowcz007 42a97fa4d9 Update image_mixlab.js 2024-04-18 16:00:43 +08:00
shadowcz007 45240d0012 Update image_mixlab.js 2024-04-18 16:00:18 +08:00
shadowcz007 8ed085febd Update image_mixlab.js 2024-04-18 15:57:34 +08:00
shadowcz007 37803ea61b Update smart_connect.js 2024-04-18 09:06:42 +08:00
shadowcz007 acd416952c Update image_mixlab.js 2024-04-18 07:48:50 +08:00
shadowcz007 6ec46cbc44 Update ui_mixlab.js 2024-04-17 22:55:16 +08:00
shadowcz007 a9d971e476 Update __init__.py 2024-04-17 22:55:12 +08:00
shadowcz007 9fe064675d paste appinfo data / 支持粘贴appinfo导出的数据 2024-04-17 16:28:32 +08:00
shadowcz007 c84fa467d0 Update requirements.txt 2024-04-17 15:36:49 +08:00
shadowcz007 3d7a55f6d3 SaveImageAndMetadata支持格式化文件名@bakkhos8 2024-04-17 15:01:05 +08:00
shadowcz007 d49baa1540 Update checkVersion_mixlab.js 2024-04-15 23:40:01 +08:00
shadowcz007 5c686af842 Update ImageNode.py 2024-04-15 09:02:49 +08:00
shadowcz007 22425b5bc6 Update index.html 2024-04-13 00:12:05 +08:00
shadowcz007 b6d9b338d2 Update index.html 2024-04-12 21:55:56 +08:00
shadowcz007 a191a13751 Update index.html 2024-04-12 21:48:31 +08:00
shadowcz007 3eccdbcc9b Update index.html 2024-04-12 21:43:21 +08:00
shadowcz007 50063903f9 mixlab app add 3D 2024-04-12 16:10:05 +08:00
shadowcz007 95a1b70533 Update app_mixlab.js 2024-04-08 17:14:30 +08:00
shadowcz007 c3679ac90b Update image_mixlab.js 2024-04-07 11:49:26 +08:00
shadowcz007 29e48eb6a2 Update ui_mixlab.js 2024-04-07 10:32:13 +08:00
shadowcz007 71d02e9651 v0.20.0 2024-04-06 21:47:07 +08:00
shadowcz007 f5193f3eec add LoadImagesToBatch 2024-04-06 21:39:55 +08:00
shadowcz007 480c4d6919 fixbug 2024-04-06 19:56:43 +08:00
shadowcz007 1c5e030540 Update index.html 2024-04-06 18:41:06 +08:00
shadowcz007 27e83a5908 Incrementing List 2024-03-30 23:46:32 +08:00
shadowcz007 d3cbf8fa8d Update Video.py 2024-03-30 22:46:50 +08:00
shadowcz007 74b1f8129b Update PromptNode.py 2024-03-29 15:35:41 +08:00
shadowcz007 56ed513cfd update 2024-03-27 17:03:28 +08:00
shadowcz007 5f412371c4 Update Utils.py 2024-03-27 13:10:30 +08:00
shadowcz007 c6b0b67585 update 2024-03-26 21:44:25 +08:00
shadowcz007 16d18e681a add composite_images 2024-03-26 18:23:11 +08:00
shadowcz007 1fe99f33b2 add VAEEncodeForInpaint_Frames 2024-03-26 16:48:45 +08:00
shadowcz007 a2ece25ac0 update 2024-03-26 15:43:43 +08:00
shadowcz007 4865f4d148 fixbug 2024-03-26 15:14:01 +08:00
shadowcz007 41e88824cf Update Video.py 2024-03-26 13:38:00 +08:00
shadowcz007 0961ab138e Update videoupload.js 2024-03-26 13:35:43 +08:00
shadowcz007 fe8271a12f fixbug 2024-03-26 13:34:52 +08:00
shadowcz007 f3866ede89 add ImageListReplace 2024-03-26 00:29:21 +08:00
shadowcz007 d938adf3cc update TextSplitByDelimiter 2024-03-25 20:24:54 +08:00
shadowcz007 9908cff64b Update Utils.py 2024-03-25 13:21:47 +08:00
shadowcz007 1b9b0bb4e6 fixbug 2024-03-25 13:20:30 +08:00
39 changed files with 5362 additions and 832 deletions
+52 -12
View File
@@ -3,13 +3,27 @@
> [Mixlab nodes discord](https://discord.gg/cXs9vZSqeK)
####
[comfyui-sd-prompt-mixlab](https://github.com/shadowcz007/comfyui-sd-prompt-mixlab)
[comfyui-Image-reward](https://github.com/shadowcz007/comfyui-Image-reward)
[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)
<!-- [comfyui-CLIPSeg](https://github.com/shadowcz007/comfyui-CLIPSeg) -->
####
最新:ChatGPT节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。
Model download,move to :```models/llamafile/```
强烈推荐:[Phi-3-mini-4k-instruct-GGUF](https://huggingface.co/lmstudio-community/Phi-3-mini-4k-instruct-GGUF/tree/main)
备选:[llama3_if_ai_sdpromptmkr_q2k](https://hf-mirror.com/impactframes/llama3_if_ai_sdpromptmkr_q2k/tree/main)
## 🚀🚗🚚🏃 Workflow-to-APP
@@ -17,7 +31,9 @@
- 支持多个web app 切换
- 发布为app的workflow,可以在右键里再次编辑了
- web app可以设置分类,在comfyui右键菜单可以编辑更新web app
- 支持动态提示
![](./assets/微信图片_20240421205440.png)
- Support multiple web app switching.
- Add the AppInfo node, which allows you to transform the workflow into a web app by simple configuration.
@@ -75,13 +91,23 @@ https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43
[Voice + Real-time Face Swap Workflow](./workflow/语音+实时换脸workflow.json)
### GPT
> Support for calling multiple GPTs.ChatGPT、ChatGLM3 、ChatGLM4 , Some code provided by rui. If you are using OpenAI's service, fill in https://api.openai.com/v1 . If you are using a local LLM service, fill in http://127.0.0.1:xxxx/v1 . Azure OpenAI:https://xxxx.openai.azure.com
> Support for calling multiple GPTs.Local LLM(llama.cpp)、 ChatGPT、ChatGLM3 、ChatGLM4 , Some code provided by rui. If you are using OpenAI's service, fill in https://api.openai.com/v1 . If you are using a local LLM service, fill in http://127.0.0.1:xxxx/v1 . Azure OpenAI:https://xxxx.openai.azure.com
![gpt-workflow.svg](./assets/gpt-workflow.svg)
[workflow-5](./workflow/5-gpt-workflow.json)
最新:ChatGPT节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。
Model download,move to :```models/llamafile/```
强烈推荐:[Phi-3-mini-4k-instruct-GGUF](https://huggingface.co/lmstudio-community/Phi-3-mini-4k-instruct-GGUF/tree/main)
备选:[llama3_if_ai_sdpromptmkr_q2k](https://hf-mirror.com/impactframes/llama3_if_ai_sdpromptmkr_q2k/tree/main)
## Prompt
> PromptSlide
![](./assets/prompt_weight.png)
@@ -112,22 +138,31 @@ https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43
### 3D
![](./assets/3d-workflow.png)
![](./assets/3d_app.png)
[workflow](./assets/Image-to-3D_1.json)
![](./assets/3dimage.png)
[workflow](./workflow/3D-workflow.json)
### LoadImagesFromLocal
### Image
#### LoadImagesToBatch
> Upload multiple images for batch input into the IP adapter.
#### LoadImagesFromLocal
> Monitor changes to images in a local folder, and trigger real-time execution of workflows, supporting common image formats, especially PSD format, in conjunction with Photoshop.
![watch](./assets/4-loadfromlocal-watcher-workflow.svg)
[workflow-4](./workflow/4-loadfromlocal-watcher-workflow.json)
### LoadImagesFromURL
#### LoadImagesFromURL
> Conveniently load images from a fixed address on the internet to ensure that default images in the workflow can be executed.
## Style
### Style
> Apply VisualStyle Prompting , Modified from [ComfyUI_VisualStylePrompting](https://github.com/ExponentialML/ComfyUI_VisualStylePrompting)
![](./assets/VisualStylePrompting.png)
@@ -135,7 +170,7 @@ https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43
> StyleAligned , Modified from [style_aligned_comfy](https://github.com/brianfitzgerald/style_aligned_comfy)
## Utils
### Utils
> The Color node provides a color picker for easy color selection, the Font node offers built-in font selection for use with TextImage to generate text images, and the DynamicDelayByText node allows delayed execution based on the length of the input text.
- [添加了DynamicDelayByText功能,可以根据输入文本的长度进行延迟执行。](./workflow/audio-chatgpt-workflow.json)
@@ -148,7 +183,7 @@ https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43
## Other Nodes
### Other Nodes
![main](./assets/all-workflow.svg)
![main2](./assets/detect-face-all.png)
@@ -195,15 +230,20 @@ An improvement has been made to directly redirect to GitHub to search for missin
### Models
[Download rembg Models](https://github.com/danielgatis/rembg/tree/main#Models),move to:models/rembg
* [Download TripoSR](https://huggingface.co/stabilityai/TripoSR/blob/main/model.ckpt) and place it in ```models/triposr```
[Download lama](https://github.com/enesmsahin/simple-lama-inpainting/releases/download/v0.1.0/big-lama.pt), move to : models/lama
* [Download facebook/dino-vitb16](https://huggingface.co/facebook/dino-vitb16/tree/main) and place it in ```models/triposr/facebook/dino-vitb16```
[Download Salesforce/blip-image-captioning-base](https://huggingface.co/Salesforce/blip-image-captioning-base), move to : models/clip_interrogator/Salesforce/blip-image-captioning-base
[Download succinctly/text2image-prompt-generator](https://huggingface.co/succinctly/text2image-prompt-generator/tree/main),move to:prompt_generator/text2image-prompt-generator
[Download rembg Models](https://github.com/danielgatis/rembg/tree/main#Models),move to:```models/rembg```
[Download Helsinki-NLP/opus-mt-zh-en](https://huggingface.co/Helsinki-NLP/opus-mt-zh-en/tree/main),move to:prompt_generator/opus-mt-zh-en
[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```
[Download succinctly/text2image-prompt-generator](https://huggingface.co/succinctly/text2image-prompt-generator/tree/main),move to:```models/prompt_generator/text2image-prompt-generator```
[Download Helsinki-NLP/opus-mt-zh-en](https://huggingface.co/Helsinki-NLP/opus-mt-zh-en/tree/main),move to:```models/prompt_generator/opus-mt-zh-en```
## Installation
+119 -58
View File
@@ -7,7 +7,8 @@ import urllib
import hashlib
import datetime
import folder_paths
import logging
from comfy.cli_args import args
python = sys.executable
@@ -263,7 +264,7 @@ def get_my_workflow_for_app(filename="my_workflow_app.json",category="",is_all=F
print("发生异常:", str(e))
else:
app_workflow_path=os.path.join(category_path, filename)
# print('app_workflow_path: ',app_workflow_path)
print('app_workflow_path: ',app_workflow_path)
try:
with open(app_workflow_path) as json_file:
apps = [{
@@ -273,7 +274,8 @@ def get_my_workflow_for_app(filename="my_workflow_app.json",category="",is_all=F
except Exception as e:
print("发生异常:", str(e))
if len(apps)==1 and category!='' and category!=None:
# 这个代码不需要
# if len(apps)==1 and category!='' and category!=None:
data=read_workflow_json_files(category_path)
for item in data:
@@ -413,37 +415,75 @@ async def new_start(self, address, port, verbose=True, call_on_start=None):
runner = web.AppRunner(self.app, access_log=None)
await runner.setup()
if not await check_port_available(address, port):
raise RuntimeError(f"Port {port} is already in use.")
# if not await check_port_available(address, port):
# raise RuntimeError(f"Port {port} is already in use.")
http_success = False
http_port=port
for i in range(11): # 尝试最多11次
if await check_port_available(address, port + i):
http_port = port + i
site = web.TCPSite(runner, address, http_port)
await site.start()
http_success = True
break
site = web.TCPSite(runner, address, port)
await site.start()
if not http_success:
raise RuntimeError(f"Ports {port} to {port + 10} are all in use.")
import ssl
crt, key = create_for_https()
ssl_context = ssl.create_default_context(ssl.Purpose.CLIENT_AUTH)
ssl_context.load_cert_chain(crt, key)
# site = web.TCPSite(runner, address, port)
# await site.start()
ssl_context = None
scheme = "http"
# 跟着本体修改
if args.tls_keyfile and args.tls_certfile:
scheme = "https"
ssl_context = ssl.SSLContext(protocol=ssl.PROTOCOL_TLS_SERVER, verify_mode=ssl.CERT_NONE)
ssl_context.load_cert_chain(certfile=args.tls_certfile,
keyfile=args.tls_keyfile)
else:
# 如果没传,则自动创建
import ssl
crt, key = create_for_https()
ssl_context = ssl.create_default_context(ssl.Purpose.CLIENT_AUTH)
ssl_context.load_cert_chain(crt, key)
success = False
for i in range(10): # 尝试最多10次
if await check_port_available(address, port + 1 + i):
https_port = port + 1 + i
for i in range(11): # 尝试最多11次
if await check_port_available(address, http_port + 1 + i):
https_port = http_port + 1 + i
site2 = web.TCPSite(runner, address, https_port, ssl_context=ssl_context)
await site2.start()
success = True
break
if not success:
raise RuntimeError(f"Ports {port + 1} to {port + 10} are all in use.")
raise RuntimeError(f"Ports {http_port + 1} to {http_port + 10} are all in use.")
if address == '':
address = '0.0.0.0'
address = '127.0.0.1'
if address=='0.0.0.0':
address = '127.0.0.1'
if verbose:
print("\033[93mStarting server\n")
print("\033[93mTo see the GUI go to: http://{}:{}".format(address, port))
print("\033[93mTo see the GUI go to: https://{}:{}\033[0m".format(address, https_port))
logging.info("\n")
logging.info("\n\nStarting server")
# print("\033[93mStarting server\n")
logging.info("\033[93mTo see the GUI go to: http://{}:{}".format(address, http_port))
logging.info("\033[93mTo see the GUI go to: https://{}:{}\033[0m".format(address, https_port))
# print("\033[93mTo see the GUI go to: http://{}:{}".format(address, http_port))
# print("\033[93mTo see the GUI go to: https://{}:{}\033[0m".format(address, https_port))
if call_on_start is not None:
call_on_start(address, port)
if scheme=='https':
call_on_start(scheme,address, https_port)
else:
call_on_start(scheme,address, http_port)
except Exception as e:
print(f"Error starting the server: {e}")
@@ -455,6 +495,7 @@ async def new_start(self, address, port, verbose=True, call_on_start=None):
# webbrowser.open(f"http://{address}:{port}")
PromptServer.start=new_start
# 创建路由表
@@ -576,9 +617,6 @@ async def post_prompt_result(request):
return web.json_response({"result":res})
# 扩展api接口
# from server import PromptServer
# from aiohttp import web
@@ -592,18 +630,20 @@ async def post_prompt_result(request):
# 导入节点
from .nodes.PromptNode import GLIGENTextBoxApply_Advanced,EmbeddingPrompt,RandomPrompt,PromptSlide,PromptSimplification,PromptImage,JoinWithDelimiter
from .nodes.ImageNode import GridDisplayAndSave,GridInput,ImagesPrompt,SaveImageAndMetadata,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.ImageNode import LoadImages_,CompositeImages,GridDisplayAndSave,GridInput,ImagesPrompt,SaveImageAndMetadata,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.ChatGPT import ChatGPTNode,ShowTextForGPT,CharacterInText,TextSplitByDelimiter
from .nodes.Audio import GamePal,SpeechRecognition,SpeechSynthesis
from .nodes.Utils import ListSplit,CreateLoraNames,CreateSampler_names,CreateCkptNames,CreateSeedNode,TESTNODE_,TESTNODE_TOKEN,AppInfo,IntNumber,FloatSlider,TextInput,ColorInput,FontInput,TextToNumber,DynamicDelayProcessor,LimitNumber,SwitchByIndex,MultiplicationNode
from .nodes.Mask import MaskListReplace,MaskListMerge,OutlineMask,FeatheredMask
from .nodes.Utils import IncrementingListNode,ListSplit,CreateLoraNames,CreateSampler_names,CreateCkptNames,CreateSeedNode,TESTNODE_,TESTNODE_TOKEN,AppInfo,IntNumber,FloatSlider,TextInput,ColorInput,FontInput,TextToNumber,DynamicDelayProcessor,LimitNumber,SwitchByIndex,MultiplicationNode
from .nodes.Mask import PreviewMask_,MaskListReplace,MaskListMerge,OutlineMask,FeatheredMask
from .nodes.Style import ApplyVisualStylePrompting,StyleAlignedReferenceSampler,StyleAlignedBatchAlign,StyleAlignedSampleReferenceLatents
from .nodes.Video import LoadVideoAndSegment
from .nodes.Video import VideoCombine_Adv,LoadVideoAndSegment,ImageListReplace,VAEEncodeForInpaint_Frames
from .nodes.TripoSR import LoadTripoSRModel,TripoSRSampler,SaveTripoSRMesh
# 要导出的所有节点及其名称的字典
@@ -626,6 +666,7 @@ NODE_CLASS_MAPPINGS = {
"ResizeImageMixlab":ResizeImage,
"LoadImagesFromPath":LoadImagesFromPath,
"LoadImagesFromURL":LoadImagesFromURL,
"LoadImagesToBatch":LoadImages_,
"TextImage":TextImage,
"EnhanceImage":EnhanceImage,
"SvgImage":SvgImage,
@@ -633,6 +674,7 @@ NODE_CLASS_MAPPINGS = {
"ImageColorTransfer":ImageColorTransfer,
"ShowLayer":ShowLayer,
"NewLayer":NewLayer,
"CompositeImages_":CompositeImages,
"SplitImage":SplitImage,
"CenterImage":CenterImage,
"GridOutput":GridOutput,
@@ -681,9 +723,16 @@ NODE_CLASS_MAPPINGS = {
"StyleAlignedSampleReferenceLatents_": StyleAlignedSampleReferenceLatents,
"StyleAlignedBatchAlign_": StyleAlignedBatchAlign,
"LoadVideoAndSegment_":LoadVideoAndSegment,
"VideoCombine_Adv":VideoCombine_Adv,
"ListSplit_":ListSplit,
"MaskListReplace_":MaskListReplace
# "LaMaInpainting":LaMaInpainting
"MaskListReplace_":MaskListReplace,
"ImageListReplace_":ImageListReplace,
"VAEEncodeForInpaint_Frames":VAEEncodeForInpaint_Frames,
"IncrementingListNode_":IncrementingListNode,
"PreviewMask_":PreviewMask_,
"LoadTripoSRModel_": LoadTripoSRModel,
"TripoSRSampler_": TripoSRSampler,
"SaveTripoSRMesh": SaveTripoSRMesh
# "GamePal":GamePal
}
@@ -698,81 +747,93 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"SaveImageAndMetadata_":"Save Image Output ♾️MixlabApp",
"ResizeImageMixlab":"Resize Image ♾️Mixlab",
"RandomPrompt": "Random Prompt ♾️Mixlab",
"PromptImage":"Output Prompt and Image",
"PromptImage":"Output Prompt and Image ♾️Mixlab",
"SplitLongMask":"Splitting a long image into sections",
"VAELoaderConsistencyDecoder":"Consistency Decoder Loader",
"VAEDecodeConsistencyDecoder":"Consistency Decoder Decode",
"ScreenShare":"Screen Share ♾️Mixlab",
"FloatingVideo":"FloatingVideo ♾️Mixlab",
"ChatGPTOpenAI":"ChatGPT ♾️Mixlab",
"ChatGPTOpenAI":"ChatGPT & Local LLM ♾️Mixlab",
"ShowTextForGPT":"Show Text ♾️MixlabApp",
"MergeLayers":"Merge Layers ♾️Mixlab",
"SpeechSynthesis":"SpeechSynthesis ♾️Mixlab",
"SpeechRecognition":"SpeechRecognition ♾️Mixlab",
"3DImage":"3DImage ♾️Mixlab",
"CompositeImages_":"Composite Images ♾️Mixlab",
"DynamicDelayProcessor":"DynamicDelayByText ♾️Mixlab",
"LaMaInpainting":"LaMaInpainting ♾️Mixlab",
"PromptSlide":"Prompt Slide ♾️Mixlab",
"PromptGenerate_Mix":"Prompt Generate ♾️Mixlab",
"ChinesePrompt_Mix":"Chinese Prompt ♾️Mixlab",
"GamePal":"GamePal ♾️Mixlab",
"RembgNode_Mix":"Remove Background",
"LoraNames_":"LoraName",
"ApplyVisualStylePrompting_":"Apply VisualStyle Prompting",
"StyleAlignedReferenceSampler_": "StyleAligned Reference Sampler",
"StyleAlignedSampleReferenceLatents_": "StyleAligned Sample Reference Latents",
"StyleAlignedBatchAlign_": "StyleAligned Batch Align",
"LoadVideoAndSegment_":"Load Video And Segment",
"MaskListMerge_":"MaskList to Mask",
"ListSplit_":"Split List",
"MaskListReplace_":"MaskList Replace",
"SwitchByIndex":"List Switch By Index",
"RembgNode_Mix":"Remove Background ♾️Mixlab",
"LoraNames_":"LoraName ♾️Mixlab",
"ApplyVisualStylePrompting_":"Apply VisualStyle Prompting ♾️Mixlab",
"StyleAlignedReferenceSampler_": "StyleAligned Reference Sampler ♾️Mixlab",
"StyleAlignedSampleReferenceLatents_": "StyleAligned Sample Reference Latents ♾️Mixlab",
"StyleAlignedBatchAlign_": "StyleAligned Batch Align ♾️Mixlab",
"LoadVideoAndSegment_":"Load Video And Segment ♾️Mixlab",
"VideoCombine_Adv":"Video Combine ♾️Mixlab",
"MaskListMerge_":"MaskList to Mask ♾️Mixlab",
"ListSplit_":"Split List ♾️Mixlab",
"MaskListReplace_":"MaskList Replace ♾️Mixlab",
"ImageListReplace_":"ImageList Replace ♾️Mixlab",
"SwitchByIndex":"List Switch By Index ♾️Mixlab",
"GLIGENTextBoxApply_Advanced":"GLIGEN TextBox Apply ♾️Mixlab",
"GridDisplayAndSave":"Grid Display And Save",
"GridInput":"Grid Input",
"GridOutput":"Grid Output",
"GetImageSize_":"Get Image Size"
"GridDisplayAndSave":"Grid Display And Save ♾️Mixlab",
"GridInput":"Grid Input ♾️Mixlab",
"GridOutput":"Grid Output ♾️Mixlab",
"GetImageSize_":"Get Image Size ♾️Mixlab",
"VAEEncodeForInpaint_Frames":"VAE Encode For Inpaint Frames ♾️Mixlab",
"IncrementingListNode_":"Create Incrementing Number List ♾️Mixlab",
"LoadImagesToBatch":"Load Images(base64) ♾️Mixlab",
"PreviewMask_":"Preview Mask",
"LoadTripoSRModel_": "Load TripoSR Model",
"TripoSRSampler_": "TripoSR Sampler",
"SaveTripoSRMesh": "Save TripoSR Mesh"
}
# web ui的节点功能
WEB_DIRECTORY = "./web"
print('--------------')
print('\033[91m ### Mixlab Nodes: \033[93mLoaded')
logging.info('--------------')
logging.info('\033[91m ### Mixlab Nodes: \033[93mLoaded')
# print('\033[91m ### Mixlab Nodes: \033[93mLoaded')
try:
from .nodes.Lama import LaMaInpainting
print('LaMaInpainting.available',LaMaInpainting.available)
logging.info('LaMaInpainting.available {}'.format(LaMaInpainting.available))
if LaMaInpainting.available:
NODE_CLASS_MAPPINGS['LaMaInpainting']=LaMaInpainting
except Exception as e:
print('LaMaInpainting.available',False,e)
logging.info('LaMaInpainting.available False')
try:
from .nodes.ClipInterrogator import ClipInterrogator
print('ClipInterrogator.available',ClipInterrogator.available)
logging.info('ClipInterrogator.available {}'.format(ClipInterrogator.available))
if ClipInterrogator.available:
NODE_CLASS_MAPPINGS['ClipInterrogator']=ClipInterrogator
except Exception as e:
print('ClipInterrogator.available',False,e)
logging.info('ClipInterrogator.available False')
try:
from .nodes.TextGenerateNode import PromptGenerate,ChinesePrompt
print('PromptGenerate.available',PromptGenerate.available)
logging.info('PromptGenerate.available {}'.format(PromptGenerate.available))
if PromptGenerate.available:
NODE_CLASS_MAPPINGS['PromptGenerate_Mix']=PromptGenerate
print('ChinesePrompt.available',ChinesePrompt.available)
logging.info('ChinesePrompt.available {}'.format(ChinesePrompt.available))
if ChinesePrompt.available:
NODE_CLASS_MAPPINGS['ChinesePrompt_Mix']=ChinesePrompt
except Exception as e:
print('TextGenerateNode.available',False,e)
logging.info('TextGenerateNode.available False')
try:
from .nodes.RembgNode import RembgNode_
print('RembgNode_.available',RembgNode_.available)
logging.info('RembgNode_.available {}'.format(RembgNode_.available))
if RembgNode_.available:
NODE_CLASS_MAPPINGS['RembgNode_Mix']=RembgNode_
except Exception as e:
print('RembgNode_.available',False,e)
logging.info('RembgNode_.available False' )
print('\033[93m -------------- \033[0m')
logging.info('\033[93m -------------- \033[0m')
Binary file not shown.

After

Width:  |  Height:  |  Size: 135 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 210 KiB

File diff suppressed because one or more lines are too long
Binary file not shown.

After

Width:  |  Height:  |  Size: 965 KiB

+139 -26
View File
@@ -4,7 +4,18 @@ import urllib.error
import re,json,os,string,random
import folder_paths
import hashlib
from zhipuai import ZhipuAI
import codecs,sys
import importlib.util
def is_installed(package):
try:
spec = importlib.util.find_spec(package)
except ModuleNotFoundError:
return False
return spec is not None
def get_unique_hash(string):
hash_object = hashlib.sha1(string.encode())
unique_hash = hash_object.hexdigest()
@@ -46,13 +57,97 @@ def openai_client(key,url):
base_url=url
)
return client
def ZhipuAI_client(key):
try:
if is_installed('zhipuai')==False:
import subprocess
# 安装
print('#pip install zhipuai')
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'zhipuai'], capture_output=True, text=True)
#检查命令执行结果
if result.returncode == 0:
print("#install success")
from zhipuai import ZhipuAI
else:
print("#install error")
else:
from zhipuai import ZhipuAI
except:
print("#install zhipuai error")
client = ZhipuAI(
api_key=key, # 填写您的 APIKey
)
return client
# 优先使用phi
def phi_sort(lst):
return sorted(lst, key=lambda x: x.lower().count('phi'), reverse=True)
def get_llama_models():
res=[]
model_path=os.path.join(folder_paths.models_dir, "llamafile")
if os.path.exists(model_path):
files = os.listdir(model_path)
for file in files:
if os.path.isfile(os.path.join(model_path, file)):
res.append(file)
res=phi_sort(res)
return res
llama_modes_list=get_llama_models()
def get_llama_model_path(file_name):
model_path=os.path.join(folder_paths.models_dir, "llamafile")
mp=os.path.join(model_path,file_name)
return mp
def llama_cpp_client(file_name):
try:
if is_installed('llama_cpp')==False:
import subprocess
# 安装
print('#pip install llama-cpp-python')
result = subprocess.run([sys.executable, '-s', '-m', 'pip',
'install',
'llama-cpp-python',
'--extra-index-url',
'https://abetlen.github.io/llama-cpp-python/whl/cu121'
], capture_output=True, text=True)
#检查命令执行结果
if result.returncode == 0:
print("#install success")
from llama_cpp import Llama
else:
print("#install error")
else:
from llama_cpp import Llama
except:
print("#install llama-cpp-python error")
mp=get_llama_model_path(file_name)
# file_name=get_llama_models()[0]
# model_path=os.path.join(folder_paths.models_dir, "llamafile")
# mp=os.path.join(model_path,file_name)
llm = Llama(model_path=mp, chat_format="chatml")
return llm
def chat(client, model_name,messages ):
@@ -60,10 +155,21 @@ def chat(client, model_name,messages ):
while True:
try_count += 1
try:
response = client.chat.completions.create(
model=model_name,
messages=messages
)
if hasattr(client, "chat"):
response = client.chat.completions.create(
model=model_name,
messages=messages
)
else:
# 是llama的
response = client.create_chat_completion_openai_v1(
messages=messages,
# response_format={
# "type": "json_object",
# },
# temperature=0.7,
)
break
except openai.AuthenticationError as ex:
raise ex
@@ -72,7 +178,8 @@ def chat(client, model_name,messages ):
raise ex
time.sleep(3)
continue
# print(response.keys())
finish_reason = response.choices[0].finish_reason
if finish_reason != "stop":
raise RuntimeError("API finished with unexpected reason: " + finish_reason)
@@ -95,6 +202,16 @@ class ChatGPTNode:
@classmethod
def INPUT_TYPES(cls):
model_list=llama_modes_list+[
"gpt-3.5-turbo",
"gpt-3.5-turbo-0125",
"gpt-35-turbo",
"gpt-3.5-turbo-16k",
"gpt-3.5-turbo-16k-0613",
"gpt-4-0613",
"gpt-4-1106-preview",
"glm-4"
]
return {
"required": {
"api_key":("KEY", {"default": "", "multiline": True,"dynamicPrompts": False}),
@@ -105,16 +222,8 @@ class ChatGPTNode:
"default": "You are ChatGPT, a large language model trained by OpenAI. Answer as concisely as possible.",
"multiline": True,"dynamicPrompts": False
}),
"model": ([
"gpt-3.5-turbo",
"gpt-3.5-turbo-0125",
"gpt-35-turbo",
"gpt-3.5-turbo-16k",
"gpt-3.5-turbo-16k-0613",
"gpt-4-0613",
"gpt-4-1106-preview",
"glm-4"],
{"default": "gpt-3.5-turbo"}),
"model": ( model_list,
{"default": model_list[0]}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "step": 1}),
"context_size":("INT", {"default": 1, "min": 0, "max":30, "step": 1}),
},
@@ -137,8 +246,8 @@ class ChatGPTNode:
api_url,
prompt,
system_content,
model,
seed,context_size,unique_id = None, extra_pnginfo=None):
model,
seed,context_size,unique_id = None, extra_pnginfo=None):
# print(api_key!='',api_url,prompt,system_content,model,seed)
# 可以选择保留会话历史以维持上下文记忆
# 或者在此处清除会话历史 self.session_history.clear()
@@ -160,6 +269,9 @@ class ChatGPTNode:
if model == "glm-4" :
client = ZhipuAI_client(api_key) # 使用 Zhipuai 的接口
print('using Zhipuai interface')
elif model in llama_modes_list:
#
client=llama_cpp_client(model)
else :
client = openai_client(api_key,api_url) # 使用 ChatGPT 的接口
print('using ChatGPT interface')
@@ -314,8 +426,8 @@ class TextSplitByDelimiter:
def INPUT_TYPES(s):
return {
"required": {
"text": ("STRING", {"multiline": True,"dynamicPrompts": False}),
"delimiter":(["newline","comma"],),
"text": ("STRING", {"multiline": True,"dynamicPrompts": False}),
"delimiter":("STRING", {"multiline": False,"default":",","dynamicPrompts": False}),
"start_index": ("INT", {
"default": 0,
"min": 0, #Minimum value
@@ -349,12 +461,13 @@ class TextSplitByDelimiter:
CATEGORY = "♾️Mixlab/Text"
def run(self, text,delimiter,start_index,skip_every,max_count):
arr=[]
if delimiter=='newline':
arr = [line for line in text.split('\n') if line.strip()]
elif delimiter=='comma':
arr = [line for line in text.split(',') if line.strip()]
if delimiter=="":
arr=[text.strip()]
else:
delimiter=codecs.decode(delimiter, 'unicode_escape')
arr= [line for line in text.split(delimiter) if line.strip()]
arr= arr[start_index:start_index + max_count * (skip_every+1):(skip_every+1)]
return (arr,)
+282 -36
View File
@@ -8,11 +8,61 @@ import base64,os,random
from io import BytesIO
import folder_paths
import json,io
import comfy.utils
from comfy.cli_args import args
import cv2
import string
import math,glob
from .Watcher import FolderWatcher
import hashlib
# 将PIL图片转换为OpenCV格式
def pil_to_opencv(image):
open_cv_image = cv2.cvtColor(np.array(image), cv2.COLOR_RGB2BGR)
return open_cv_image
# 将OpenCV格式图片转换为PIL格式
def opencv_to_pil(image):
pil_image = Image.fromarray(cv2.cvtColor(image, cv2.COLOR_BGR2RGB))
return pil_image
def composite_images(foreground, background, mask,is_multiply_blend=False):
width,height=foreground.size
bg_image=background
# 按z-index排序
layer = {
"x":0,
"y":0,
"width":width,
"height":height,
"z_index":88,
"scale_option":'overall',
"image":foreground,
"mask":mask
}
width, height = bg_image.size
layer_image=layer['image']
layer_mask=layer['mask']
bg_image=merge_images(bg_image,
layer_image,
layer_mask,
layer['x'],
layer['y'],
layer['width'],
layer['height'],
layer['scale_option'],
is_multiply_blend )
bg_image=bg_image.convert('RGB')
return bg_image
def count_files_in_directory(directory):
@@ -583,7 +633,7 @@ def detect_faces(image):
def areaToMask(x,y,w,h,image):
# 创建一个与原图片大小相同的空白图片
mask = Image.new('1', image.size)
mask = Image.new('L', image.size)
# 创建一个可用于绘制的对象
draw = ImageDraw.Draw(mask)
@@ -616,7 +666,45 @@ def areaToMask(x,y,w,h,image):
# # bg_image.save("output.jpg")
# return bg_image
def merge_images(bg_image, layer_image, mask, x, y, width, height, scale_option):
import cv2
import numpy as np
# ps的正片叠底
# 可以基于https://www.cnblogs.com/jsxyhelu/p/16947810.html ,用gpt写python代码
def multiply_blend(image1, image2):
image1=pil_to_opencv(image1)
image2=pil_to_opencv(image2)
# 将图像转换为浮点型
image1 = image1.astype(float)
image2 = image2.astype(float)
if image1.shape != image2.shape:
image1 = cv2.resize(image1, (image2.shape[1], image2.shape[0]))
# 归一化图像
image1 /= 255.0
image2 /= 255.0
# 正片叠底混合
blended = image1 * image2
# 将图像还原为8位无符号整数
blended = (blended * 255).astype(np.uint8)
blended=opencv_to_pil(blended)
return blended
# # 读取图像
# image1 = cv2.imread('1.png')
# image2 = cv2.imread('3.png')
# # 进行正片叠底混合
# result = multiply_blend(image1, image2)
# cv2.imwrite('result.jpg', result)
def merge_images(bg_image, layer_image, mask, x, y, width, height, scale_option,is_multiply_blend=False):
# 打开底图
bg_image = bg_image.convert("RGBA")
@@ -645,13 +733,45 @@ def merge_images(bg_image, layer_image, mask, x, y, width, height, scale_option)
nw, nh = layer_image.size
mask = mask.resize((nw, nh))
# 在底图上粘贴图层
bg_image.paste(layer_image, (x, y), mask=mask)
# # 分离出a通道
# r, g, b, alpha = layer_image.split()
# alpha = ImageOps.invert(alpha)
# # 创建一个新的RGB图像
# new_rgb_image = Image.new("RGB", layer_image.size)
# # 将透明通道粘贴到新的RGB图像上
# new_rgb_image.paste(layer_image, (0, 0), mask=alpha)
# new_rgb_image.paste(layer_image, (x, y), mask=mask)
# mask=new_rgb_image.convert('L')
# mask = ImageOps.invert(mask)
if is_multiply_blend:
bg_image_white=Image.new("RGB", bg_image.size,(255, 255, 255))
bg_image_white.paste(layer_image, (x, y), mask=mask)
bg_image=multiply_blend(bg_image_white,bg_image)
bg_image=bg_image.convert("RGBA")
else:
transparent_img = Image.new("RGBA",layer_image.size, (0, 0, 0, 0))
transparent_img.paste(layer_image,(0, 0), mask)
bg_image.paste(transparent_img, (x, y), transparent_img)
# 输出合成后的图片
return bg_image
def resize_2(img):
# 检查图像的高度是否是2的倍数,如果不是,则调整高度
if img.height % 2 != 0:
img = img.resize((img.width, img.height + 1))
# 检查图像的宽度是否是2的倍数,如果不是,则调整宽度
if img.width % 2 != 0:
img = img.resize((img.width + 1, img.height))
return img
# TODO 几个像素点的底
def resize_image(layer_image, scale_option, width, height,color="white"):
layer_image = layer_image.convert("RGB")
@@ -681,8 +801,10 @@ def resize_image(layer_image, scale_option, width, height,color="white"):
resized_image = Image.new("RGB", (width, height), color=color)
resized_image.paste(layer_image.resize((new_width, new_height)), ((width - new_width) // 2, (height - new_height) // 2))
resized_image = resized_image.convert("RGB")
resized_image=resize_2(resized_image)
return resized_image
layer_image=resize_2(layer_image)
return layer_image
@@ -1082,6 +1204,42 @@ class EnhanceImage:
class LoadImages_:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"images": ("IMAGEBASE64",),
},
}
CATEGORY = "♾️Mixlab/Image"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,)
RETURN_TYPES = ("IMAGE",)
FUNCTION = "load_image"
def load_image(self, images):
# print(images)
ims=[]
for im in images['base64']:
image = base64_to_image(im)
image=image.convert('RGB')
image=pil2tensor(image)
ims.append(image)
image1 = ims[0]
for image2 in ims[1:]:
if image1.shape[1:] != image2.shape[1:]:
image2 = comfy.utils.common_upscale(image2.movedim(-1, 1), image1.shape[2], image1.shape[1], "bilinear", "center").movedim(1, -1)
image1 = torch.cat((image1, image2), dim=0)
return (image1,)
'''
("STRING",{"multiline": False,"default": "Hello World!"})
对应 widgets.js 里:
@@ -1431,7 +1589,7 @@ class Image3D:
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Image"
CATEGORY = "♾️Mixlab/3D"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,False,False,False,)
@@ -1541,6 +1699,42 @@ class FaceToMask:
return (mask,)
class CompositeImages:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"foreground": ("IMAGE",),
"mask":("MASK",),
"background": ("IMAGE",),
},
"optional":{
"is_multiply_blend": ("BOOLEAN", {"default": False}),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("IMAGE",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Layer"
# OUTPUT_IS_LIST = (True,)
def run(self, foreground,mask,background,is_multiply_blend):
foreground= tensor2pil(foreground)
mask= tensor2pil(mask)
background= tensor2pil(background)
res=composite_images(foreground,background,mask,is_multiply_blend)
return (pil2tensor(res),)
class EmptyLayer:
@classmethod
def INPUT_TYPES(s):
@@ -1636,7 +1830,7 @@ class NewLayer:
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"scale_option": (["width","height",'overall'],),
"image": ("IMAGE",),
"image": (any_type,),
},
"optional":{
"mask": ("MASK",{"default": None}),
@@ -1693,21 +1887,20 @@ def createMask(image,x,y,w,h):
# mask.save("mask.png")
return mask
def splitImage(image, num):
width, height = image.size
num_rows = int(num ** 0.5)
num_cols = int(num / num_rows)
grid_width = width // num_cols
grid_height = height // num_rows
grid_width = int(width // num_cols)
grid_height = int(height // num_rows)
grid_coordinates = []
for i in range(num_rows):
for j in range(num_cols):
x = j * grid_width
y = i * grid_height
x = int(j * grid_width)
y = int(i * grid_height)
grid_coordinates.append((x, y, grid_width, grid_height))
return grid_coordinates
@@ -2039,12 +2232,16 @@ class GridOutput:
def INPUT_TYPES(s):
return {
"required": {
"grid": ("_GRID",)
}
"grid": ("_GRID",),
},
"optional":{
"bg_image":("IMAGE",)
}
}
RETURN_TYPES = ("INT","INT","INT","INT",)
RETURN_NAMES = ("x","y","width","height",)
RETURN_TYPES = ("INT","INT","INT","INT","MASK",)
RETURN_NAMES = ("x","y","width","height","mask",)
FUNCTION = "run"
@@ -2053,9 +2250,29 @@ class GridOutput:
INPUT_IS_LIST = False
# OUTPUT_IS_LIST = (True,)
def run(self,grid):
def run(self,grid,bg_image=None):
x,y,w,h=grid
return (x,y,w,h,)
x=int(x)
y=int(y)
w=int(w)
h=int(h)
masks=[]
if bg_image!=None:
for i in range(len(bg_image)):
im=bg_image[i]
#增加输出mask
im=tensor2pil(im)
mask=areaToMask(x,y,w,h,im)
mask=pil2tensor(mask)
masks.append(mask)
out=None
if len(masks)>0:
out = torch.cat(masks, dim=0)
return (x,y,w,h,out,)
class ShowLayer:
@classmethod
@@ -2147,10 +2364,14 @@ class MergeLayers:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"layers": ("LAYER",),
"images": ("IMAGE",),
},
"layers": ("LAYER",),
"images": ("IMAGE",),
},
"optional":{
"is_multiply_blend": ("BOOLEAN", {"default": False}),
}
}
RETURN_TYPES = ("IMAGE","MASK",)
@@ -2163,11 +2384,12 @@ class MergeLayers:
INPUT_IS_LIST = True
# OUTPUT_IS_LIST = (False,)
def run(self,layers,images):
def run(self,layers,images,is_multiply_blend):
bg_images=[]
masks=[]
is_multiply_blend=is_multiply_blend[0]
# print(len(images),images[0].shape)
# 1 torch.Size([2, 512, 512, 3])
# 4 torch.Size([1, 1024, 768, 3])
@@ -2198,6 +2420,8 @@ class MergeLayers:
layer_image=tensor2pil(image)
layer_mask=tensor2pil(mask)
# t=layer_image.convert("RGBA")
# t.save('test.png') 如果layerimage传入的是rgba,则是透明的
bg_image=merge_images(bg_image,
layer_image,
layer_mask,
@@ -2205,7 +2429,8 @@ class MergeLayers:
layer['y'],
layer['width'],
layer['height'],
layer['scale_option']
layer['scale_option'],
is_multiply_blend
)
final_mask=merge_images(final_mask,
@@ -2618,6 +2843,7 @@ class ImageColorTransfer:
return {"required": {
"source": ("IMAGE",),
"target": ("IMAGE",),
"weight": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
},
}
@@ -2631,25 +2857,45 @@ class ImageColorTransfer:
CATEGORY = "♾️Mixlab/Color"
# 输入是否为列表
INPUT_IS_LIST = True
# INPUT_IS_LIST = True
# 输出是否为列表
OUTPUT_IS_LIST = (True,)
# OUTPUT_IS_LIST = (True,)
def run(self,source,target):
def run(self,source,target,weight):
res=[]
target=target[0][0]
print(target.shape)
target=tensor2pil(target)
#batch-list
source_list = [source[i:i + 1, ...] for i in range(source.shape[0])]
target_list = [target[i:i + 1, ...] for i in range(target.shape[0])]
for ims in source:
for im in ims:
image=tensor2pil(im)
image=color_transfer(image,target)
image=pil2tensor(image)
res.append(image)
# 长度纠正为相等
if len(target_list) != len(source_list):
target_list = target_list * (len(source_list) // len(target_list)) + target_list[:len(source_list) % len(target_list)]
for i in range(len(source_list)):
target=target_list[i]
source=source_list[i]
target=tensor2pil(target)
image=tensor2pil(source)
image_res=color_transfer(image,target)
# weight Blend image # contributors:@ning
blend_mask = Image.new(mode="L", size=image.size,
color=(round(weight * 255)))
blend_mask = ImageOps.invert(blend_mask)
img_result = Image.composite(image, image_res, blend_mask)
del image, image_res, blend_mask
img_result=pil2tensor(img_result)
res.append(img_result)
# list - batch
res=torch.cat(res, dim=0)
return (res,)
+38 -11
View File
@@ -2,16 +2,14 @@
import scipy.ndimage
import torch
from nodes import MAX_RESOLUTION
import numpy as np
# from PIL import Image, ImageDraw
from PIL import Image, ImageOps
from comfy.cli_args import args
import cv2
import cv2,os
from nodes import MAX_RESOLUTION, SaveImage, common_ksampler
import folder_paths,random
# Tensor to PIL
def tensor2pil(image):
@@ -71,6 +69,35 @@ def combine(destination, source, x, y):
return output
class PreviewMask_(SaveImage):
def __init__(self):
self.output_dir = folder_paths.get_temp_directory()
self.type = "temp"
self.prefix_append =''.join(random.choice("abcdehijklmnopqrstupvxyzfg") for x in range(5))
self.compress_level = 4
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mask": ("MASK",),
}
}
RETURN_TYPES = ()
OUTPUT_NODE = True
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Mask"
# 运行的函数
def run(self, mask ):
img=tensor2pil(mask)
img=img.convert('RGB')
img=pil2tensor(img)
return self.save_images(img, 'temp_', None, None)
class OutlineMask:
@classmethod
@@ -108,32 +135,32 @@ class MaskListReplace:
"mask_replace": ("MASK",),
"start_index":("INT", {"default": 0, "min": 0, "step": 1}),
"end_index":("INT", {"default": 0, "min": 0, "step": 1}),
"reverse": ("BOOLEAN", {"default": False}),
"invert": ("BOOLEAN", {"default": False}),
}
}
RETURN_TYPES = ("MASK",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Mask"
CATEGORY = "♾️Mixlab/Video"
INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True,)
def run(self, masks,mask_replace,start_index,end_index,reverse):
def run(self, masks,mask_replace,start_index,end_index,invert):
mask_replace=mask_replace[0]
start_index=start_index[0]
end_index=end_index[0]
reverse=reverse[0]
invert=invert[0]
new_masks=[]
for i in range(len(masks)):
if i>=start_index and i<=end_index:
if reverse:
if invert:
new_masks.append(masks[i])
else:
new_masks.append(mask_replace)
else:
if reverse:
if invert:
new_masks.append(mask_replace)
else:
new_masks.append(masks[i])
+10 -8
View File
@@ -580,16 +580,17 @@ class GLIGENTextBoxApply_Advanced:
RETURN_NAMES = ("CONDITIONING","label",)
FUNCTION = "run"
INPUT_IS_LIST = True
# INPUT_IS_LIST = True
CATEGORY = "♾️Mixlab/Prompt"
def run(self, conditioning, clip, gligen_textbox_model, grids, labels, index,max_size,random_shuffle,seed=0):
conditioning=conditioning[0]
clip=clip[0]
gligen_textbox_model=gligen_textbox_model[0]
index=index[0]
max_size=max_size[0]
random_shuffle=random_shuffle[0]
# print('grids',grids)
# conditioning=conditioning[0]
# clip=clip[0]
# gligen_textbox_model=gligen_textbox_model[0]
# index=index[0]
# max_size=max_size[0]
# random_shuffle=random_shuffle[0]
texts=labels
@@ -618,7 +619,7 @@ class GLIGENTextBoxApply_Advanced:
text=texts[i]
grid=grids[i]
x,y,width,height=grid
print(text)
# print(text)
cond, cond_pooled = clip.encode_from_tokens(clip.tokenize(text), return_pooled=True)
position_params =position_params+ [(cond_pooled, height // 8, width // 8, y // 8, x // 8)]
@@ -628,6 +629,7 @@ class GLIGENTextBoxApply_Advanced:
prev = n[1]['gligen'][2]
n[1]['gligen'] = ("position", gligen_textbox_model, prev + position_params)
# print('gligen',n)
c.append(n)
# 下面这个写法有bug
+173
View File
@@ -0,0 +1,173 @@
import sys
from os import path
sys.path.insert(0, path.dirname(__file__))
from PIL import Image
import numpy as np
import torch
from folder_paths import get_filename_list, get_full_path, get_save_image_path, get_output_directory,models_dir
from comfy.model_management import get_torch_device
from .tsr.system import TSR
import comfy.utils
triposr_model_path=path.join(models_dir,'triposr/model.ckpt')
# 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 fill_background(image):
im = np.array(image).astype(np.float32) / 255.0
im = im[:, :, :3] * im[:, :, 3:4] + (1 - im[:, :, 3:4]) * 0.5
im = Image.fromarray((im * 255.0).astype(np.uint8))
return im
class LoadTripoSRModel:
def __init__(self):
self.initialized_model = None
@classmethod
def INPUT_TYPES(s):
return {
"required": {
# "model": (get_filename_list("checkpoints"),),
"chunk_size": ("INT", {"default": 8192, "min": 0, "max": 10000})
}
}
RETURN_TYPES = ("TRIPOSR_MODEL",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/3D/TripoSR"
def run(self, chunk_size):
device = get_torch_device()
if not torch.cuda.is_available():
device = "cpu"
if not self.initialized_model:
# triposr_model_path
print("#Loading TripoSR model",triposr_model_path)
self.initialized_model = TSR.from_pretrained_custom(
weight_path=triposr_model_path,
config_path=path.join(path.dirname(__file__), "tsr/config.yaml")
)
self.initialized_model.renderer.set_chunk_size(chunk_size)
self.initialized_model.to(device)
return (self.initialized_model,)
class TripoSRSampler:
def __init__(self):
self.initialized_model = None
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("TRIPOSR_MODEL",),
"image": ("IMAGE",),
"resolution": ("INT", {"default": 256, "min": 128, "max": 12288}),
"threshold": ("FLOAT", {"default": 25.0, "min": 0.0, "step": 0.01}),
"device":(["auto","cpu"],),
},
"optional": {
"mask": ("MASK",)
}
}
RETURN_TYPES = ("MESH",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/3D/TripoSR"
def run(self, model, image, resolution, threshold,device='auto', mask=None):
reference_image=image
reference_mask=mask
device = get_torch_device()
if not torch.cuda.is_available():
device = "cpu"
if device=='cpu':
device = "cpu"
print('#TripoSRSampler device',device)
to_images=[]
for i in range(len(reference_image)):
image = reference_image[i]
if reference_mask is not None:
mask = reference_mask[i].unsqueeze(2)
image = torch.cat((image, mask), dim=2).detach().cpu().numpy()
image = Image.fromarray(np.clip(255. * image, 0, 255).astype(np.uint8))
image = fill_background(image)
else:
image = tensor2pil(image)
image = image.convert('RGB')
to_images.append(image)
# 进度条
pbar = comfy.utils.ProgressBar(len(to_images))
def callback(c):
pbar.update(1)
scene_codes = model(to_images, device)
meshes = model.extract_mesh(scene_codes, resolution=resolution, threshold=threshold,callback=callback)
del model
return (meshes,)
class SaveTripoSRMesh:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mesh": ("MESH",),
# "format":(["glb","obj"],),
"filename_prefix":("STRING", {"multiline": False,"default": "TripoSR_"})
}
}
RETURN_TYPES = ()
OUTPUT_NODE = True
FUNCTION = "run"
CATEGORY = "♾️Mixlab/3D/TripoSR"
def run(self, mesh,filename_prefix):
format='glb'
saved = list()
full_output_folder, filename, counter, subfolder, filename_prefix = get_save_image_path(filename_prefix,
get_output_directory())
for (index, single_mesh) in enumerate(mesh):
filename_with_batch_num = filename.replace("%batch_num%", str(index))
file = f"{filename_with_batch_num}_{counter:05}_.{format}"
single_mesh.apply_transform(np.array([[1, 0, 0, 0], [0, 0, 1, 0], [0, -1, 0, 0], [0, 0, 0, 1]]))
single_mesh.export(path.join(full_output_folder, file))
saved.append({
"filename": file,
"type": "output",
"subfolder": subfolder
})
return {"ui": {"mesh": saved}}
+71 -9
View File
@@ -8,6 +8,10 @@ import matplotlib.font_manager as fm
import torch
import importlib.util
def create_incrementing_list(min_value, max_value, step, count):
l1 = [int(min_value + i * step) for i in range(count) if min_value + i * step <= max_value]
l2 = [float(min_value + i * step) for i in range(count) if min_value + i * step <= max_value]
return (l1,l2)
def split_list(lst, chunk_size, transition_size):
result = []
@@ -154,7 +158,6 @@ class ColorInput:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"color":("TCOLOR",),
},
}
@@ -281,7 +284,7 @@ class FloatSlider:
}
RETURN_TYPES = ("FLOAT",)
RETURN_NAMES = ('weight(0-1)',)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Input"
@@ -406,6 +409,61 @@ class TextInput:
return (text,)
class IncrementingListNode:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"min_value": ("FLOAT", {
"default": 0,
"min": -2000, #Minimum value
"max": 0xffffffffffffffff,
"step": 0.01, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"max_value": ("FLOAT", {
"default": 10,
"min": -2000, #Minimum value
"max": 0xffffffffffffffff,
"step": 0.01, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"step": ("FLOAT", {
"default": 0,
"min": -2000, #Minimum value
"max": 0xffffffffffffffff,
"step": 0.01, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
}),
"count": ("INT", {
"default": 1,
"min": 1, #Minimum value
"max": 0xffffffffffffffff,
"step":1, #Slider's step
"display": "number" # Cosmetic only: display as "number" or "slider"
})
},
"optional":{
"seed":("INT", {"default": -1, "min": -1, "max": 1000000}),
},
}
RETURN_TYPES = ("INT","FLOAT",)
RETURN_NAMES = ('int_list','float_list',)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Video"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (True,True,)
def run(self,min_value,max_value,step,count,seed):
print('create_incrementing_list',seed)
l1,l2=create_incrementing_list(min_value,max_value,step,count)
return (l1,l2,)
# 接收一个值,然后根据字符串或数值长度计算延迟时间,用户可以自定义延迟"字/s",延迟之后将转化
import comfy.samplers
@@ -585,14 +643,14 @@ class SwitchByIndex:
}
RETURN_TYPES = (any_type,"INT",)
RETURN_NAMES = ("C","count",)
RETURN_NAMES = ("list", "count",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Utils"
INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True,False,)
OUTPUT_IS_LIST = (True, False,)
def run(self, A=[],B=[],index=-1,flat='on'):
@@ -613,9 +671,9 @@ class SwitchByIndex:
try:
C=[C[index]]
except Exception as e:
C=[]
C=[C[-1]] #最后一个
return (C,len(C),)
return (C, len(C),)
class ListSplit:
@classmethod
@@ -755,9 +813,13 @@ class TESTNODE_:
def run(self,ANY):
print(type(ANY))
print(ANY[0].shape)
img= tensor2pil(ANY[0])
print(img.size)
try:
print(ANY[0].shape)
img= tensor2pil(ANY[0])
print(img.size)
except:
print('')
# data=ANY
list_stats = ListStatistics()
+537 -59
View File
@@ -4,13 +4,13 @@ import json
import subprocess
import shutil
import re
import time
import time,math
import numpy as np
from typing import List
import torch
from PIL import Image, ImageOps
from PIL.PngImagePlugin import PngInfo
import cv2
import cv2,random,string
from pathlib import Path
import folder_paths
@@ -18,6 +18,73 @@ from comfy.k_diffusion.utils import FolderOfImages
from comfy.utils import common_upscale
def generate_folder_name(directory,video_path):
# Get the directory and filename from the video path
_, filename = os.path.split(video_path)
# Generate a random string of lowercase letters and digits
random_string = ''.join(random.choices(string.ascii_lowercase + string.digits, k=8))
# Create the folder name by combining the random string and the filename
folder_name = random_string + '_' + filename
# Create the full folder path by joining the directory and the folder name
folder_path = os.path.join(directory, folder_name)
return folder_path
def create_folder(directory,video_path):
folder_path = generate_folder_name(directory,video_path)
os.makedirs(folder_path)
return folder_path
def split_video(video_path, video_segment_frames, transition_frames, output_dir):
# 读取视频文件
video_capture = cv2.VideoCapture(video_path)
# 获取视频的总帧数和帧率
total_frames = int(video_capture.get(cv2.CAP_PROP_FRAME_COUNT))
fps = video_capture.get(cv2.CAP_PROP_FPS)
# 计算每个视频片段的总帧数,包括过渡帧
segment_total_frames = video_segment_frames + transition_frames
# 计算可以分割的片段数量,向上取整
num_segments = (total_frames + transition_frames - 1) // segment_total_frames
vs=[]
# 计算每个片段的起始帧和结束帧
start_frame = 0
for i in range(num_segments):
# 计算当前片段的结束帧,注意最后一个片段可能没有过渡帧
end_frame = min(start_frame + segment_total_frames, total_frames)
# 打印当前片段的起始帧和结束帧
print(f"Segment {i+1}: Start Frame {start_frame}, End Frame {end_frame}")
# 保存当前片段为一个视频文件
segment_video_path = f"{output_dir}/segment_{i+1}.avi"
fourcc = cv2.VideoWriter_fourcc(*'XVID')
segment_video = cv2.VideoWriter(segment_video_path, fourcc, fps, (int(video_capture.get(cv2.CAP_PROP_FRAME_WIDTH)),
int(video_capture.get(cv2.CAP_PROP_FRAME_HEIGHT))))
for frame_num in range(start_frame, end_frame):
ret, frame = video_capture.read()
if ret:
segment_video.write(frame)
else:
break # 如果读取失败,则退出循环
# 更新起始帧为下一个片段的起始位置
start_frame = end_frame + transition_frames
vs.append(segment_video_path)
# 释放视频捕获对象
video_capture.release()
# print(vs)
return (vs,total_frames,fps)
folder_paths.folder_names_and_paths["video_formats"] = (
[
os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "video_formats"),
@@ -34,6 +101,48 @@ if ffmpeg_path is None:
except:
print("ffmpeg could not be found. Outputs that require it have been disabled")
# 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 count_files(directory):
count = 0
for root, dirs, files in os.walk(directory):
count += len(files)
return count
def create_temp_file(image):
output_dir = folder_paths.get_temp_directory()
c=count_files(output_dir)
(
full_output_folder,
filename,
counter,
subfolder,
_,
) = folder_paths.get_save_image_path('temp_', output_dir)
image=tensor2pil(image)
image_file = f"{filename}_{c}_{counter:05}.png"
image_path=os.path.join(full_output_folder, image_file)
image.save(image_path,compress_level=4)
return [{
"filename": image_file,
"subfolder": subfolder,
"type": "temp"
}]
def split_list(lst, chunk_size, transition_size):
result = []
@@ -50,7 +159,96 @@ def split_list(lst, chunk_size, transition_size):
# result = split_list(images, chunk_size, transition_size)
# print(result)
class ImageListReplace:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"images": ("IMAGE",),
"start_index":("INT", {"default": 0, "min": 0, "step": 1}),
"end_index":("INT", {"default": 0, "min": 0, "step": 1}),
"invert": ("BOOLEAN", {"default": False}),
},
"optional":{
"image_replace": ("IMAGE",),
"images_replace": ("IMAGE",),
}
}
RETURN_TYPES = ("IMAGE","IMAGE",)
RETURN_NAMES = ("images","select_images",)
FUNCTION = "run"
CATEGORY = "♾️Mixlab/Video"
OUTPUT_NODE = True
INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True,True,)
def run(self, images,start_index=[0],end_index=[0],invert=[False],image_replace=None,images_replace=None):
start_index=start_index[0]
end_index=end_index[0]
invert=invert[0]
image_rs=[]
if image_replace!=None:
for i in range(end_index-start_index+1):
image_rs.append(image_replace[0])
if images_replace!=None:
image_rs=images_replace
# 如果image replace 为空
if image_replace==None and images_replace==None:
# print('如果image replace 为空',images[0])
# [[tensor(
# tensor([[[[0.
first_image=tensor2pil(images[0][0])
width, height = first_image.size
image_replace=Image.new("RGB", (width, height), (0, 0, 0))
image_replace=pil2tensor(image_replace)
for i in range(end_index-start_index+1):
image_rs.append(image_replace)
new_images=[]
select_images=[]
k=0
for i in range(len(images)):
if i>=start_index and i<=end_index:
if invert:
new_images.append(images[i])
else:
new_images.append(image_rs[k])
select_images.append(images[i])
k+=1
else:
if invert:
new_images.append(image_rs[k])
select_images.append(images[i])
k+=1
else:
new_images.append(images[i])
imss=[]
# print(len(images))
for i in range(len(images)):
t=images[i][0]
t=tensor2pil(t)
t = t.convert("RGB")
original_width, original_height = t.size
scale = 300 / original_width
new_height = int(original_height * scale)
t = t.resize((300, new_height))
ims=create_temp_file(pil2tensor(t))
imss.append(ims[0])
# image_replace=create_temp_file(image_replace)
return {"ui":{"_images": imss},"result": (new_images,select_images,)}
# The code is based on ComfyUI-VideoHelperSuite modification.
class LoadVideoAndSegment:
@classmethod
def INPUT_TYPES(s):
@@ -70,11 +268,11 @@ class LoadVideoAndSegment:
CATEGORY = "♾️Mixlab/Video"
RETURN_TYPES = ("IMAGE", "INT",)
RETURN_NAMES = ("image_batch", "frame_count",)
RETURN_TYPES = ("SCENE_VIDEO","INT", "INT","INT",)
RETURN_NAMES = ("scenes_video","scenes_count","frame_count","fps",)
FUNCTION = "load_video"
OUTPUT_NODE = True
OUTPUT_IS_LIST = (True,False,)
OUTPUT_IS_LIST = (True,False,False,False,)
def is_gif(self, filename):
@@ -131,70 +329,84 @@ class LoadVideoAndSegment:
return (images, frames_added)
def load_video(self, video,video_segment_frames,transition_frames ):
frame_load_cap=0
skip_first_frames=0
video_path = folder_paths.get_annotated_filepath(video)
# check if video is a gif - will need to use cv fallback to read frames
# use cv fallback if ffmpeg not installed or gif
if ffmpeg_path is None:
return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
# if ffmpeg_path is None:
# return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
# otherwise, continue with ffmpeg
video_path = folder_paths.get_annotated_filepath(video)
args_dummy = [ffmpeg_path, "-i", video_path, "-f", "null", "-"]
try:
with subprocess.Popen(args_dummy, stdout=subprocess.DEVNULL, stderr=subprocess.PIPE) as proc:
for line in proc.stderr.readlines():
match = re.search(", ([1-9]|\\d{2,})x(\\d+)",line.decode('utf-8'))
if match is not None:
size = [int(match.group(1)), int(match.group(2))]
break
except Exception as e:
print(f"Retrying with opencv due to ffmpeg error: {e}")
return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
args_all_frames = [ffmpeg_path, "-i", video_path, "-v", "error",
"-pix_fmt", "rgb24"]
vfilters = []
if skip_first_frames > 0:
vfilters.append(f"select=gt(n\\,{skip_first_frames-1})")
if frame_load_cap > 0:
vfilters.append(f"select=gt({frame_load_cap}\\,n)")
#manually calculate aspect ratio to ensure reads remain aligned
if len(vfilters) > 0:
args_all_frames += ["-vf", ",".join(vfilters)]
# args_dummy = [ffmpeg_path, "-i", video_path, "-f", "null", "-"]
# try:
# with subprocess.Popen(args_dummy, stdout=subprocess.DEVNULL, stderr=subprocess.PIPE) as proc:
# for line in proc.stderr.readlines():
# match = re.search(", ([1-9]|\\d{2,})x(\\d+)",line.decode('utf-8'))
# if match is not None:
# size = [int(match.group(1)), int(match.group(2))]
# break
# except Exception as e:
# print(f"Retrying with opencv due to ffmpeg error: {e}")
# return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
# args_all_frames = [ffmpeg_path, "-i", video_path, "-v", "error",
# "-pix_fmt", "rgb24"]
args_all_frames += ["-f", "rawvideo", "-"]
images = []
try:
with subprocess.Popen(args_all_frames, stdout=subprocess.PIPE) as proc:
#Manually buffer enough bytes for an image
bpi = size[0]*size[1]*3
current_bytes = bytearray(bpi)
current_offset=0
while True:
bytes_read = proc.stdout.read(bpi - current_offset)
if bytes_read is None:#sleep to wait for more data
time.sleep(.2)
continue
if len(bytes_read) == 0:#EOF
break
current_bytes[current_offset:len(bytes_read)] = bytes_read
current_offset+=len(bytes_read)
if current_offset == bpi:
images.append(np.array(current_bytes, dtype=np.float32).reshape(size[1], size[0], 3) / 255.0)
current_offset = 0
except Exception as e:
print(f"Retrying with opencv due to ffmpeg error: {e}")
return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
# vfilters = []
# if skip_first_frames > 0:
# vfilters.append(f"select=gt(n\\,{skip_first_frames-1})")
# if frame_load_cap > 0:
# vfilters.append(f"select=gt({frame_load_cap}\\,n)")
# #manually calculate aspect ratio to ensure reads remain aligned
# if len(vfilters) > 0:
# args_all_frames += ["-vf", ",".join(vfilters)]
imgs=split_list(images,video_segment_frames,transition_frames)
# args_all_frames += ["-f", "rawvideo", "-"]
# images = []
# try:
# with subprocess.Popen(args_all_frames, stdout=subprocess.PIPE) as proc:
# #Manually buffer enough bytes for an image
# bpi = size[0]*size[1]*3
# current_bytes = bytearray(bpi)
# current_offset=0
# while True:
# bytes_read = proc.stdout.read(bpi - current_offset)
# if bytes_read is None:#sleep to wait for more data
# time.sleep(.2)
# continue
# if len(bytes_read) == 0:#EOF
# break
# current_bytes[current_offset:len(bytes_read)] = bytes_read
# current_offset+=len(bytes_read)
# if current_offset == bpi:
# images.append(np.array(current_bytes, dtype=np.float32).reshape(size[1], size[0], 3) / 255.0)
# current_offset = 0
# except Exception as e:
# print(f"Retrying with opencv due to ffmpeg error: {e}")
# return self.load_video_cv_fallback(video, frame_load_cap, skip_first_frames)
imgs=[torch.from_numpy(np.stack(im)) for im in imgs]
# imgs=split_list(images,video_segment_frames,transition_frames)
# temp path
tp=folder_paths.get_temp_directory()
basename = os.path.basename(video_path) # 获取文件名
name_without_extension = os.path.splitext(basename)[0] # 去掉文件后缀
folder_path = create_folder(tp,name_without_extension)
# 导出的数据
scenes_video,total_frames,fps=split_video(video_path,video_segment_frames,
transition_frames,folder_path)
# imgs=[torch.from_numpy(np.stack(im)) for im in imgs]
# images = torch.from_numpy(np.stack(images))
return (imgs, len(imgs))
return (scenes_video,len(scenes_video), total_frames,fps,)
@classmethod
def IS_CHANGED(s, video, **kwargs):
@@ -212,3 +424,269 @@ class LoadVideoAndSegment:
return True
# The code is based on ComfyUI-VideoHelperSuite modification.
class VideoCombine_Adv:
@classmethod
def INPUT_TYPES(s):
#Hide ffmpeg formats if ffmpeg isn't available
if ffmpeg_path is not None:
ffmpeg_formats = ["video/"+x[:-5] for x in folder_paths.get_filename_list("video_formats")]
else:
ffmpeg_formats = []
return {
"required": {
"image_batch": ("IMAGE",),
"frame_rate": (
"INT",
{"default": 8, "min": 1, "step": 1},
),
"loop_count": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1}),
"filename_prefix": ("STRING", {"default": "Comfyui"}),
"format": (["image/gif", "image/webp"] + ffmpeg_formats,),
"pingpong": ("BOOLEAN", {"default": False}),
"save_image": ("BOOLEAN", {"default": True}),
"metadata": ("BOOLEAN", {"default": False}),
},
"hidden": {
"prompt": "PROMPT",
"extra_pnginfo": "EXTRA_PNGINFO",
},
}
RETURN_TYPES = ()
OUTPUT_NODE = True
CATEGORY = "♾️Mixlab/Video"
FUNCTION = "run"
def save_with_tempfile(self, args, metadata, file_path, frames, env):
#Ensure temp directory exists
os.makedirs(folder_paths.get_temp_directory(), exist_ok=True)
metadata_path = os.path.join(folder_paths.get_temp_directory(), "metadata.txt")
#metadata from file should escape = ; # \ and newline
#From my testing, though, only backslashes need escapes and = in particular causes problems
#It is likely better to prioritize future compatibility with containers that don't support
#or shouldn't use the comment tag for embedding metadata
metadata = metadata.replace("\\","\\\\")
metadata = metadata.replace(";","\\;")
metadata = metadata.replace("#","\\#")
#metadata = metadata.replace("=","\\=")
metadata = metadata.replace("\n","\\\n")
with open(metadata_path, "w") as f:
f.write(";FFMETADATA1\n")
f.write(metadata)
args = args[:1] + ["-i", metadata_path] + args[1:] + [file_path]
with subprocess.Popen(args, stdin=subprocess.PIPE, env=env) as proc:
for frame in frames:
proc.stdin.write(frame.tobytes())
def run(
self,
image_batch,
frame_rate: int,
loop_count: int,
filename_prefix="AnimateDiff",
format="image/gif",
pingpong=False,
save_image=True,
metadata=False,
prompt=None,
extra_pnginfo=None,
):
images=image_batch
frames: List[Image.Image] = []
for image in images:
img = 255.0 * image.cpu().numpy()
img = Image.fromarray(np.clip(img, 0, 255).astype(np.uint8))
# resize 保证
# 检查图像的高度是否是2的倍数,如果不是,则调整高度
if img.height % 2 != 0:
img = img.resize((img.width, img.height + 1))
# 检查图像的宽度是否是2的倍数,如果不是,则调整宽度
if img.width % 2 != 0:
img = img.resize((img.width + 1, img.height))
frames.append(img)
# get output information
output_dir = (
folder_paths.get_output_directory()
if save_image
else folder_paths.get_temp_directory()
)
(
full_output_folder,
filename,
counter,
subfolder,
_,
) = folder_paths.get_save_image_path(filename_prefix, output_dir)
metadata = PngInfo()
video_metadata = {}
if prompt is not None:
metadata.add_text("prompt", json.dumps(prompt))
video_metadata["prompt"] = prompt
if extra_pnginfo is not None:
for x in extra_pnginfo:
metadata.add_text(x, json.dumps(extra_pnginfo[x]))
video_metadata[x] = extra_pnginfo[x]
# 取消保存metadata
if metadata==False:
metadata = PngInfo()
# save first frame as png to keep metadata
file = f"{filename}_{counter:05}_.png"
file_path = os.path.join(full_output_folder, file)
frames[0].save(
file_path,
pnginfo=metadata,
compress_level=4,
)
if pingpong:
frames = frames + frames[-2:0:-1]
format_type, format_ext = format.split("/")
file = f"{filename}_{counter:05}_.{format_ext}"
file_path = os.path.join(full_output_folder, file)
if format_type == "image":
# Use pillow directly to save an animated image
frames[0].save(
file_path,
format=format_ext.upper(),
save_all=True,
append_images=frames[1:],
duration=round(1000 / frame_rate),
loop=loop_count,
compress_level=4,
)
else:
# Use ffmpeg to save a video
if ffmpeg_path is None:
#Should never be reachable
raise ProcessLookupError("Could not find ffmpeg")
video_format_path = folder_paths.get_full_path("video_formats", format_ext + ".json")
with open(video_format_path, 'r') as stream:
video_format = json.load(stream)
file = f"{filename}_{counter:05}_.{video_format['extension']}"
file_path = os.path.join(full_output_folder, file)
dimensions = f"{frames[0].width}x{frames[0].height}"
metadata_args = ["-metadata", "comment=" + json.dumps(video_metadata)]
args = [ffmpeg_path, "-v", "error", "-f", "rawvideo", "-pix_fmt", "rgb24",
"-s", dimensions, "-r", str(frame_rate), "-i", "-"] \
+ video_format['main_pass']
# On linux, max arg length is Pagesize * 32 -> 131072
# On windows, this around 32767 but seems to vary wildly by > 500
# in a manor not solely related to other arguments
if os.name == 'posix':
max_arg_length = 4096*32
else:
max_arg_length = 32767 - len(" ".join(args + [metadata_args[0]] + [file_path])) - 1
#test max limit
#metadata_args[1] = metadata_args[1] + "a"*(max_arg_length - len(metadata_args[1])-1)
env=os.environ.copy()
if "environment" in video_format:
env.update(video_format["environment"])
if len(metadata_args[1]) >= max_arg_length:
print(f"Using fallback file for extremely long metadata: {len(metadata_args[1])}/{max_arg_length}")
self.save_with_tempfile(args, metadata_args[1], file_path, frames, env)
else:
try:
with subprocess.Popen(args + metadata_args + [file_path],
stdin=subprocess.PIPE, env=env) as proc:
for frame in frames:
proc.stdin.write(frame.tobytes())
except FileNotFoundError as e:
if "winerror" in dir(e) and e.winerror == 206:
print("Metadata was too long. Retrying with fallback file")
self.save_with_tempfile(args, metadata_args[1], file_path, frames, env)
else:
raise
except OSError as e:
if "errno" in dir(e) and e.errno == 7:
print("Metadata was too long. Retrying with fallback file")
self.save_with_tempfile(args, metadata_args[1], file_path, frames, env)
else:
raise
previews = [
{
"filename": file,
"subfolder": subfolder,
"type": "output" if save_image else "temp",
"format": format,
}
]
return {"ui": {"gifs": previews}}
class VAEEncodeForInpaint_Frames:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"vae": ("VAE", ),
"images": ("IMAGE", ),
"masks": ("MASK", ),
"grow_mask_by": ("INT", {"default": 6, "min": 0, "max": 64, "step": 1}),
}}
FUNCTION = "encode"
RETURN_TYPES = ("LATENT",)
RETURN_NAMES = ("LATENT",)
CATEGORY = "♾️Mixlab/Video"
OUTPUT_NODE = True
INPUT_IS_LIST = True
OUTPUT_IS_LIST = (True,)
def encode(self, vae, images, masks, grow_mask_by=[6]):
vae=vae[0]
grow_mask_by=grow_mask_by[0]
result=[]
for i in range(len(images)):
pixels=images[i]
mask=masks[i]
x = (pixels.shape[1] // 8) * 8
y = (pixels.shape[2] // 8) * 8
mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(pixels.shape[1], pixels.shape[2]), mode="bilinear")
pixels = pixels.clone()
if pixels.shape[1] != x or pixels.shape[2] != y:
x_offset = (pixels.shape[1] % 8) // 2
y_offset = (pixels.shape[2] % 8) // 2
pixels = pixels[:,x_offset:x + x_offset, y_offset:y + y_offset,:]
mask = mask[:,:,x_offset:x + x_offset, y_offset:y + y_offset]
#grow mask by a few pixels to keep things seamless in latent space
if grow_mask_by == 0:
mask_erosion = mask
else:
kernel_tensor = torch.ones((1, 1, grow_mask_by, grow_mask_by))
padding = math.ceil((grow_mask_by - 1) / 2)
mask_erosion = torch.clamp(torch.nn.functional.conv2d(mask.round(), kernel_tensor, padding=padding), 0, 1)
m = (1.0 - mask.round()).squeeze(1)
for i in range(3):
pixels[:,:,:,i] -= 0.5
pixels[:,:,:,i] *= m
pixels[:,:,:,i] += 0.5
t = vae.encode(pixels)
result.append({"samples":t, "noise_mask": (mask_erosion[:,:,:x,:y].round())})
return (result, )
+38
View File
@@ -0,0 +1,38 @@
cond_image_size: 512
image_tokenizer_cls: tsr.models.tokenizers.image.DINOSingleImageTokenizer
image_tokenizer:
pretrained_model_name_or_path: "facebook/dino-vitb16"
tokenizer_cls: tsr.models.tokenizers.triplane.Triplane1DTokenizer
tokenizer:
plane_size: 32
num_channels: 1024
backbone_cls: tsr.models.transformer.transformer_1d.Transformer1D
backbone:
in_channels: ${tokenizer.num_channels}
num_attention_heads: 16
attention_head_dim: 64
num_layers: 16
cross_attention_dim: 768
post_processor_cls: tsr.models.network_utils.TriplaneUpsampleNetwork
post_processor:
in_channels: 1024
out_channels: 40
decoder_cls: tsr.models.network_utils.NeRFMLP
decoder:
in_channels: 120 # 3 * 40
n_neurons: 64
n_hidden_layers: 9
activation: silu
renderer_cls: tsr.models.nerf_renderer.TriplaneNeRFRenderer
renderer:
radius: 0.87 # slightly larger than 0.5 * sqrt(3)
feature_reduction: concat
density_activation: exp
density_bias: -1.0
num_samples_per_ray: 128
+51
View File
@@ -0,0 +1,51 @@
from typing import Callable, Optional, Tuple
import numpy as np
import torch
import torch.nn as nn
from skimage import measure
class IsosurfaceHelper(nn.Module):
points_range: Tuple[float, float] = (0, 1)
@property
def grid_vertices(self) -> torch.FloatTensor:
raise NotImplementedError
class MarchingCubeHelper(IsosurfaceHelper):
def __init__(self, resolution: int) -> None:
super().__init__()
self.resolution = resolution
#self.mc_func: Callable = marching_cubes
self._grid_vertices: Optional[torch.FloatTensor] = None
@property
def grid_vertices(self) -> torch.FloatTensor:
if self._grid_vertices is None:
# keep the vertices on CPU so that we can support very large resolution
x, y, z = (
torch.linspace(*self.points_range, self.resolution),
torch.linspace(*self.points_range, self.resolution),
torch.linspace(*self.points_range, self.resolution),
)
x, y, z = torch.meshgrid(x, y, z, indexing="ij")
verts = torch.cat(
[x.reshape(-1, 1), y.reshape(-1, 1), z.reshape(-1, 1)], dim=-1
).reshape(-1, 3)
self._grid_vertices = verts
return self._grid_vertices
def forward(
self,
level: torch.FloatTensor,
) -> Tuple[torch.FloatTensor, torch.LongTensor]:
level = -level.view(self.resolution, self.resolution, self.resolution)
v_pos, t_pos_idx, _, __ = measure.marching_cubes((level.detach().cpu() if level.is_cuda else level.detach()).numpy(), 0.0) #self.mc_func(level.detach(), 0.0)
v_pos = torch.from_numpy(v_pos.copy()).type(torch.FloatTensor).to(level.device)
t_pos_idx = torch.from_numpy(t_pos_idx.copy()).type(torch.LongTensor).to(level.device)
v_pos = v_pos[..., [0, 1, 2]]
t_pos_idx = t_pos_idx[..., [1, 0, 2]]
v_pos = v_pos / (self.resolution - 1.0)
return v_pos, t_pos_idx
+180
View File
@@ -0,0 +1,180 @@
from dataclasses import dataclass
from typing import Dict
import torch
import torch.nn.functional as F
from einops import rearrange, reduce
from ..utils import (
BaseModule,
chunk_batch,
get_activation,
rays_intersect_bbox,
scale_tensor,
)
class TriplaneNeRFRenderer(BaseModule):
@dataclass
class Config(BaseModule.Config):
radius: float
feature_reduction: str = "concat"
density_activation: str = "trunc_exp"
density_bias: float = -1.0
color_activation: str = "sigmoid"
num_samples_per_ray: int = 128
randomized: bool = False
cfg: Config
def configure(self) -> None:
assert self.cfg.feature_reduction in ["concat", "mean"]
self.chunk_size = 0
def set_chunk_size(self, chunk_size: int):
assert (
chunk_size >= 0
), "chunk_size must be a non-negative integer (0 for no chunking)."
self.chunk_size = chunk_size
def query_triplane(
self,
decoder: torch.nn.Module,
positions: torch.Tensor,
triplane: torch.Tensor,
) -> Dict[str, torch.Tensor]:
input_shape = positions.shape[:-1]
positions = positions.view(-1, 3)
# positions in (-radius, radius)
# normalized to (-1, 1) for grid sample
positions = scale_tensor(
positions, (-self.cfg.radius, self.cfg.radius), (-1, 1)
)
def _query_chunk(x):
indices2D: torch.Tensor = torch.stack(
(x[..., [0, 1]], x[..., [0, 2]], x[..., [1, 2]]),
dim=-3,
)
out: torch.Tensor = F.grid_sample(
rearrange(triplane, "Np Cp Hp Wp -> Np Cp Hp Wp", Np=3),
rearrange(indices2D, "Np N Nd -> Np () N Nd", Np=3),
align_corners=False,
mode="bilinear",
)
if self.cfg.feature_reduction == "concat":
out = rearrange(out, "Np Cp () N -> N (Np Cp)", Np=3)
elif self.cfg.feature_reduction == "mean":
out = reduce(out, "Np Cp () N -> N Cp", Np=3, reduction="mean")
else:
raise NotImplementedError
net_out: Dict[str, torch.Tensor] = decoder(out)
return net_out
if self.chunk_size > 0:
net_out = chunk_batch(_query_chunk, self.chunk_size, positions)
else:
net_out = _query_chunk(positions)
net_out["density_act"] = get_activation(self.cfg.density_activation)(
net_out["density"] + self.cfg.density_bias
)
net_out["color"] = get_activation(self.cfg.color_activation)(
net_out["features"]
)
net_out = {k: v.view(*input_shape, -1) for k, v in net_out.items()}
return net_out
def _forward(
self,
decoder: torch.nn.Module,
triplane: torch.Tensor,
rays_o: torch.Tensor,
rays_d: torch.Tensor,
**kwargs,
):
rays_shape = rays_o.shape[:-1]
rays_o = rays_o.view(-1, 3)
rays_d = rays_d.view(-1, 3)
n_rays = rays_o.shape[0]
t_near, t_far, rays_valid = rays_intersect_bbox(rays_o, rays_d, self.cfg.radius)
t_near, t_far = t_near[rays_valid], t_far[rays_valid]
t_vals = torch.linspace(
0, 1, self.cfg.num_samples_per_ray + 1, device=triplane.device
)
t_mid = (t_vals[:-1] + t_vals[1:]) / 2.0
z_vals = t_near * (1 - t_mid[None]) + t_far * t_mid[None] # (N_rays, N_samples)
xyz = (
rays_o[:, None, :] + z_vals[..., None] * rays_d[..., None, :]
) # (N_rays, N_sample, 3)
mlp_out = self.query_triplane(
decoder=decoder,
positions=xyz,
triplane=triplane,
)
eps = 1e-10
# deltas = z_vals[:, 1:] - z_vals[:, :-1] # (N_rays, N_samples)
deltas = t_vals[1:] - t_vals[:-1] # (N_rays, N_samples)
alpha = 1 - torch.exp(
-deltas * mlp_out["density_act"][..., 0]
) # (N_rays, N_samples)
accum_prod = torch.cat(
[
torch.ones_like(alpha[:, :1]),
torch.cumprod(1 - alpha[:, :-1] + eps, dim=-1),
],
dim=-1,
)
weights = alpha * accum_prod # (N_rays, N_samples)
comp_rgb_ = (weights[..., None] * mlp_out["color"]).sum(dim=-2) # (N_rays, 3)
opacity_ = weights.sum(dim=-1) # (N_rays)
comp_rgb = torch.zeros(
n_rays, 3, dtype=comp_rgb_.dtype, device=comp_rgb_.device
)
opacity = torch.zeros(n_rays, dtype=opacity_.dtype, device=opacity_.device)
comp_rgb[rays_valid] = comp_rgb_
opacity[rays_valid] = opacity_
comp_rgb += 1 - opacity[..., None]
comp_rgb = comp_rgb.view(*rays_shape, 3)
return comp_rgb
def forward(
self,
decoder: torch.nn.Module,
triplane: torch.Tensor,
rays_o: torch.Tensor,
rays_d: torch.Tensor,
) -> Dict[str, torch.Tensor]:
if triplane.ndim == 4:
comp_rgb = self._forward(decoder, triplane, rays_o, rays_d)
else:
comp_rgb = torch.stack(
[
self._forward(decoder, triplane[i], rays_o[i], rays_d[i])
for i in range(triplane.shape[0])
],
dim=0,
)
return comp_rgb
def train(self, mode=True):
self.randomized = mode and self.cfg.randomized
return super().train(mode=mode)
def eval(self):
self.randomized = False
return super().eval()
+124
View File
@@ -0,0 +1,124 @@
from dataclasses import dataclass
from typing import Optional
import torch
import torch.nn as nn
from einops import rearrange
from ..utils import BaseModule
class TriplaneUpsampleNetwork(BaseModule):
@dataclass
class Config(BaseModule.Config):
in_channels: int
out_channels: int
cfg: Config
def configure(self) -> None:
self.upsample = nn.ConvTranspose2d(
self.cfg.in_channels, self.cfg.out_channels, kernel_size=2, stride=2
)
def forward(self, triplanes: torch.Tensor) -> torch.Tensor:
triplanes_up = rearrange(
self.upsample(
rearrange(triplanes, "B Np Ci Hp Wp -> (B Np) Ci Hp Wp", Np=3)
),
"(B Np) Co Hp Wp -> B Np Co Hp Wp",
Np=3,
)
return triplanes_up
class NeRFMLP(BaseModule):
@dataclass
class Config(BaseModule.Config):
in_channels: int
n_neurons: int
n_hidden_layers: int
activation: str = "relu"
bias: bool = True
weight_init: Optional[str] = "kaiming_uniform"
bias_init: Optional[str] = None
cfg: Config
def configure(self) -> None:
layers = [
self.make_linear(
self.cfg.in_channels,
self.cfg.n_neurons,
bias=self.cfg.bias,
weight_init=self.cfg.weight_init,
bias_init=self.cfg.bias_init,
),
self.make_activation(self.cfg.activation),
]
for i in range(self.cfg.n_hidden_layers - 1):
layers += [
self.make_linear(
self.cfg.n_neurons,
self.cfg.n_neurons,
bias=self.cfg.bias,
weight_init=self.cfg.weight_init,
bias_init=self.cfg.bias_init,
),
self.make_activation(self.cfg.activation),
]
layers += [
self.make_linear(
self.cfg.n_neurons,
4, # density 1 + features 3
bias=self.cfg.bias,
weight_init=self.cfg.weight_init,
bias_init=self.cfg.bias_init,
)
]
self.layers = nn.Sequential(*layers)
def make_linear(
self,
dim_in,
dim_out,
bias=True,
weight_init=None,
bias_init=None,
):
layer = nn.Linear(dim_in, dim_out, bias=bias)
if weight_init is None:
pass
elif weight_init == "kaiming_uniform":
torch.nn.init.kaiming_uniform_(layer.weight, nonlinearity="relu")
else:
raise NotImplementedError
if bias:
if bias_init is None:
pass
elif bias_init == "zero":
torch.nn.init.zeros_(layer.bias)
else:
raise NotImplementedError
return layer
def make_activation(self, activation):
if activation == "relu":
return nn.ReLU(inplace=True)
elif activation == "silu":
return nn.SiLU(inplace=True)
else:
raise NotImplementedError
def forward(self, x):
inp_shape = x.shape[:-1]
x = x.reshape(-1, x.shape[-1])
features = self.layers(x)
features = features.reshape(*inp_shape, -1)
out = {"density": features[..., 0:1], "features": features[..., 1:4]}
return out
+72
View File
@@ -0,0 +1,72 @@
from dataclasses import dataclass
import torch
import torch.nn as nn
from einops import rearrange
from huggingface_hub import hf_hub_download
from transformers.models.vit.modeling_vit import ViTModel
from ...utils import BaseModule
import os
import folder_paths
model_path=os.path.join(folder_paths.models_dir,'triposr')
class DINOSingleImageTokenizer(BaseModule):
@dataclass
class Config(BaseModule.Config):
pretrained_model_name_or_path: str = "facebook/dino-vitb16"
enable_gradient_checkpointing: bool = False
cfg: Config
def configure(self) -> None:
print('#Loading ViTModel:',os.path.join(model_path,self.cfg.pretrained_model_name_or_path))
self.model: ViTModel = ViTModel(
ViTModel.config_class.from_pretrained(
hf_hub_download(
repo_id=self.cfg.pretrained_model_name_or_path,
filename="config.json",
local_dir=model_path,
endpoint='https://hf-mirror.com'
)
)
)
if self.cfg.enable_gradient_checkpointing:
self.model.encoder.gradient_checkpointing = True
self.register_buffer(
"image_mean",
torch.as_tensor([0.485, 0.456, 0.406]).reshape(1, 1, 3, 1, 1),
persistent=False,
)
self.register_buffer(
"image_std",
torch.as_tensor([0.229, 0.224, 0.225]).reshape(1, 1, 3, 1, 1),
persistent=False,
)
def forward(self, images: torch.FloatTensor, **kwargs) -> torch.FloatTensor:
packed = False
if images.ndim == 4:
packed = True
images = images.unsqueeze(1)
batch_size, n_input_views = images.shape[:2]
images = (images - self.image_mean) / self.image_std
out = self.model(
rearrange(images, "B N C H W -> (B N) C H W"), interpolate_pos_encoding=True
)
local_features, global_features = out.last_hidden_state, out.pooler_output
local_features = local_features.permute(0, 2, 1)
local_features = rearrange(
local_features, "(B N) Ct Nt -> B N Ct Nt", B=batch_size
)
if packed:
local_features = local_features.squeeze(1)
return local_features
def detokenize(self, *args, **kwargs):
raise NotImplementedError
+45
View File
@@ -0,0 +1,45 @@
import math
from dataclasses import dataclass
import torch
import torch.nn as nn
from einops import rearrange, repeat
from ...utils import BaseModule
class Triplane1DTokenizer(BaseModule):
@dataclass
class Config(BaseModule.Config):
plane_size: int
num_channels: int
cfg: Config
def configure(self) -> None:
self.embeddings = nn.Parameter(
torch.randn(
(3, self.cfg.num_channels, self.cfg.plane_size, self.cfg.plane_size),
dtype=torch.float32,
)
* 1
/ math.sqrt(self.cfg.num_channels)
)
def forward(self, batch_size: int) -> torch.Tensor:
return rearrange(
repeat(self.embeddings, "Np Ct Hp Wp -> B Np Ct Hp Wp", B=batch_size),
"B Np Ct Hp Wp -> B Ct (Np Hp Wp)",
)
def detokenize(self, tokens: torch.Tensor) -> torch.Tensor:
batch_size, Ct, Nt = tokens.shape
assert Nt == self.cfg.plane_size**2 * 3
assert Ct == self.cfg.num_channels
return rearrange(
tokens,
"B Ct (Np Hp Wp) -> B Np Ct Hp Wp",
Np=3,
Hp=self.cfg.plane_size,
Wp=self.cfg.plane_size,
)
+653
View File
@@ -0,0 +1,653 @@
# Copyright 2023 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
# --------
#
# Modified 2024 by the Tripo AI and Stability AI Team.
#
# Copyright (c) 2024 Tripo AI & Stability AI
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
#
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
# SOFTWARE.
from typing import Optional
import torch
import torch.nn.functional as F
from torch import nn
class Attention(nn.Module):
r"""
A cross attention layer.
Parameters:
query_dim (`int`):
The number of channels in the query.
cross_attention_dim (`int`, *optional*):
The number of channels in the encoder_hidden_states. If not given, defaults to `query_dim`.
heads (`int`, *optional*, defaults to 8):
The number of heads to use for multi-head attention.
dim_head (`int`, *optional*, defaults to 64):
The number of channels in each head.
dropout (`float`, *optional*, defaults to 0.0):
The dropout probability to use.
bias (`bool`, *optional*, defaults to False):
Set to `True` for the query, key, and value linear layers to contain a bias parameter.
upcast_attention (`bool`, *optional*, defaults to False):
Set to `True` to upcast the attention computation to `float32`.
upcast_softmax (`bool`, *optional*, defaults to False):
Set to `True` to upcast the softmax computation to `float32`.
cross_attention_norm (`str`, *optional*, defaults to `None`):
The type of normalization to use for the cross attention. Can be `None`, `layer_norm`, or `group_norm`.
cross_attention_norm_num_groups (`int`, *optional*, defaults to 32):
The number of groups to use for the group norm in the cross attention.
added_kv_proj_dim (`int`, *optional*, defaults to `None`):
The number of channels to use for the added key and value projections. If `None`, no projection is used.
norm_num_groups (`int`, *optional*, defaults to `None`):
The number of groups to use for the group norm in the attention.
spatial_norm_dim (`int`, *optional*, defaults to `None`):
The number of channels to use for the spatial normalization.
out_bias (`bool`, *optional*, defaults to `True`):
Set to `True` to use a bias in the output linear layer.
scale_qk (`bool`, *optional*, defaults to `True`):
Set to `True` to scale the query and key by `1 / sqrt(dim_head)`.
only_cross_attention (`bool`, *optional*, defaults to `False`):
Set to `True` to only use cross attention and not added_kv_proj_dim. Can only be set to `True` if
`added_kv_proj_dim` is not `None`.
eps (`float`, *optional*, defaults to 1e-5):
An additional value added to the denominator in group normalization that is used for numerical stability.
rescale_output_factor (`float`, *optional*, defaults to 1.0):
A factor to rescale the output by dividing it with this value.
residual_connection (`bool`, *optional*, defaults to `False`):
Set to `True` to add the residual connection to the output.
_from_deprecated_attn_block (`bool`, *optional*, defaults to `False`):
Set to `True` if the attention block is loaded from a deprecated state dict.
processor (`AttnProcessor`, *optional*, defaults to `None`):
The attention processor to use. If `None`, defaults to `AttnProcessor2_0` if `torch 2.x` is used and
`AttnProcessor` otherwise.
"""
def __init__(
self,
query_dim: int,
cross_attention_dim: Optional[int] = None,
heads: int = 8,
dim_head: int = 64,
dropout: float = 0.0,
bias: bool = False,
upcast_attention: bool = False,
upcast_softmax: bool = False,
cross_attention_norm: Optional[str] = None,
cross_attention_norm_num_groups: int = 32,
added_kv_proj_dim: Optional[int] = None,
norm_num_groups: Optional[int] = None,
out_bias: bool = True,
scale_qk: bool = True,
only_cross_attention: bool = False,
eps: float = 1e-5,
rescale_output_factor: float = 1.0,
residual_connection: bool = False,
_from_deprecated_attn_block: bool = False,
processor: Optional["AttnProcessor"] = None,
out_dim: int = None,
):
super().__init__()
self.inner_dim = out_dim if out_dim is not None else dim_head * heads
self.query_dim = query_dim
self.cross_attention_dim = (
cross_attention_dim if cross_attention_dim is not None else query_dim
)
self.upcast_attention = upcast_attention
self.upcast_softmax = upcast_softmax
self.rescale_output_factor = rescale_output_factor
self.residual_connection = residual_connection
self.dropout = dropout
self.fused_projections = False
self.out_dim = out_dim if out_dim is not None else query_dim
# we make use of this private variable to know whether this class is loaded
# with an deprecated state dict so that we can convert it on the fly
self._from_deprecated_attn_block = _from_deprecated_attn_block
self.scale_qk = scale_qk
self.scale = dim_head**-0.5 if self.scale_qk else 1.0
self.heads = out_dim // dim_head if out_dim is not None else heads
# for slice_size > 0 the attention score computation
# is split across the batch axis to save memory
# You can set slice_size with `set_attention_slice`
self.sliceable_head_dim = heads
self.added_kv_proj_dim = added_kv_proj_dim
self.only_cross_attention = only_cross_attention
if self.added_kv_proj_dim is None and self.only_cross_attention:
raise ValueError(
"`only_cross_attention` can only be set to True if `added_kv_proj_dim` is not None. Make sure to set either `only_cross_attention=False` or define `added_kv_proj_dim`."
)
if norm_num_groups is not None:
self.group_norm = nn.GroupNorm(
num_channels=query_dim, num_groups=norm_num_groups, eps=eps, affine=True
)
else:
self.group_norm = None
self.spatial_norm = None
if cross_attention_norm is None:
self.norm_cross = None
elif cross_attention_norm == "layer_norm":
self.norm_cross = nn.LayerNorm(self.cross_attention_dim)
elif cross_attention_norm == "group_norm":
if self.added_kv_proj_dim is not None:
# The given `encoder_hidden_states` are initially of shape
# (batch_size, seq_len, added_kv_proj_dim) before being projected
# to (batch_size, seq_len, cross_attention_dim). The norm is applied
# before the projection, so we need to use `added_kv_proj_dim` as
# the number of channels for the group norm.
norm_cross_num_channels = added_kv_proj_dim
else:
norm_cross_num_channels = self.cross_attention_dim
self.norm_cross = nn.GroupNorm(
num_channels=norm_cross_num_channels,
num_groups=cross_attention_norm_num_groups,
eps=1e-5,
affine=True,
)
else:
raise ValueError(
f"unknown cross_attention_norm: {cross_attention_norm}. Should be None, 'layer_norm' or 'group_norm'"
)
linear_cls = nn.Linear
self.linear_cls = linear_cls
self.to_q = linear_cls(query_dim, self.inner_dim, bias=bias)
if not self.only_cross_attention:
# only relevant for the `AddedKVProcessor` classes
self.to_k = linear_cls(self.cross_attention_dim, self.inner_dim, bias=bias)
self.to_v = linear_cls(self.cross_attention_dim, self.inner_dim, bias=bias)
else:
self.to_k = None
self.to_v = None
if self.added_kv_proj_dim is not None:
self.add_k_proj = linear_cls(added_kv_proj_dim, self.inner_dim)
self.add_v_proj = linear_cls(added_kv_proj_dim, self.inner_dim)
self.to_out = nn.ModuleList([])
self.to_out.append(linear_cls(self.inner_dim, self.out_dim, bias=out_bias))
self.to_out.append(nn.Dropout(dropout))
# set attention processor
# We use the AttnProcessor2_0 by default when torch 2.x is used which uses
# torch.nn.functional.scaled_dot_product_attention for native Flash/memory_efficient_attention
# but only if it has the default `scale` argument. TODO remove scale_qk check when we move to torch 2.1
if processor is None:
processor = (
AttnProcessor2_0()
if hasattr(F, "scaled_dot_product_attention") and self.scale_qk
else AttnProcessor()
)
self.set_processor(processor)
def set_processor(self, processor: "AttnProcessor") -> None:
self.processor = processor
def forward(
self,
hidden_states: torch.FloatTensor,
encoder_hidden_states: Optional[torch.FloatTensor] = None,
attention_mask: Optional[torch.FloatTensor] = None,
**cross_attention_kwargs,
) -> torch.Tensor:
r"""
The forward method of the `Attention` class.
Args:
hidden_states (`torch.Tensor`):
The hidden states of the query.
encoder_hidden_states (`torch.Tensor`, *optional*):
The hidden states of the encoder.
attention_mask (`torch.Tensor`, *optional*):
The attention mask to use. If `None`, no mask is applied.
**cross_attention_kwargs:
Additional keyword arguments to pass along to the cross attention.
Returns:
`torch.Tensor`: The output of the attention layer.
"""
# The `Attention` class can call different attention processors / attention functions
# here we simply pass along all tensors to the selected processor class
# For standard processors that are defined here, `**cross_attention_kwargs` is empty
return self.processor(
self,
hidden_states,
encoder_hidden_states=encoder_hidden_states,
attention_mask=attention_mask,
**cross_attention_kwargs,
)
def batch_to_head_dim(self, tensor: torch.Tensor) -> torch.Tensor:
r"""
Reshape the tensor from `[batch_size, seq_len, dim]` to `[batch_size // heads, seq_len, dim * heads]`. `heads`
is the number of heads initialized while constructing the `Attention` class.
Args:
tensor (`torch.Tensor`): The tensor to reshape.
Returns:
`torch.Tensor`: The reshaped tensor.
"""
head_size = self.heads
batch_size, seq_len, dim = tensor.shape
tensor = tensor.reshape(batch_size // head_size, head_size, seq_len, dim)
tensor = tensor.permute(0, 2, 1, 3).reshape(
batch_size // head_size, seq_len, dim * head_size
)
return tensor
def head_to_batch_dim(self, tensor: torch.Tensor, out_dim: int = 3) -> torch.Tensor:
r"""
Reshape the tensor from `[batch_size, seq_len, dim]` to `[batch_size, seq_len, heads, dim // heads]` `heads` is
the number of heads initialized while constructing the `Attention` class.
Args:
tensor (`torch.Tensor`): The tensor to reshape.
out_dim (`int`, *optional*, defaults to `3`): The output dimension of the tensor. If `3`, the tensor is
reshaped to `[batch_size * heads, seq_len, dim // heads]`.
Returns:
`torch.Tensor`: The reshaped tensor.
"""
head_size = self.heads
batch_size, seq_len, dim = tensor.shape
tensor = tensor.reshape(batch_size, seq_len, head_size, dim // head_size)
tensor = tensor.permute(0, 2, 1, 3)
if out_dim == 3:
tensor = tensor.reshape(batch_size * head_size, seq_len, dim // head_size)
return tensor
def get_attention_scores(
self,
query: torch.Tensor,
key: torch.Tensor,
attention_mask: torch.Tensor = None,
) -> torch.Tensor:
r"""
Compute the attention scores.
Args:
query (`torch.Tensor`): The query tensor.
key (`torch.Tensor`): The key tensor.
attention_mask (`torch.Tensor`, *optional*): The attention mask to use. If `None`, no mask is applied.
Returns:
`torch.Tensor`: The attention probabilities/scores.
"""
dtype = query.dtype
if self.upcast_attention:
query = query.float()
key = key.float()
if attention_mask is None:
baddbmm_input = torch.empty(
query.shape[0],
query.shape[1],
key.shape[1],
dtype=query.dtype,
device=query.device,
)
beta = 0
else:
baddbmm_input = attention_mask
beta = 1
attention_scores = torch.baddbmm(
baddbmm_input,
query,
key.transpose(-1, -2),
beta=beta,
alpha=self.scale,
)
del baddbmm_input
if self.upcast_softmax:
attention_scores = attention_scores.float()
attention_probs = attention_scores.softmax(dim=-1)
del attention_scores
attention_probs = attention_probs.to(dtype)
return attention_probs
def prepare_attention_mask(
self,
attention_mask: torch.Tensor,
target_length: int,
batch_size: int,
out_dim: int = 3,
) -> torch.Tensor:
r"""
Prepare the attention mask for the attention computation.
Args:
attention_mask (`torch.Tensor`):
The attention mask to prepare.
target_length (`int`):
The target length of the attention mask. This is the length of the attention mask after padding.
batch_size (`int`):
The batch size, which is used to repeat the attention mask.
out_dim (`int`, *optional*, defaults to `3`):
The output dimension of the attention mask. Can be either `3` or `4`.
Returns:
`torch.Tensor`: The prepared attention mask.
"""
head_size = self.heads
if attention_mask is None:
return attention_mask
current_length: int = attention_mask.shape[-1]
if current_length != target_length:
if attention_mask.device.type == "mps":
# HACK: MPS: Does not support padding by greater than dimension of input tensor.
# Instead, we can manually construct the padding tensor.
padding_shape = (
attention_mask.shape[0],
attention_mask.shape[1],
target_length,
)
padding = torch.zeros(
padding_shape,
dtype=attention_mask.dtype,
device=attention_mask.device,
)
attention_mask = torch.cat([attention_mask, padding], dim=2)
else:
# TODO: for pipelines such as stable-diffusion, padding cross-attn mask:
# we want to instead pad by (0, remaining_length), where remaining_length is:
# remaining_length: int = target_length - current_length
# TODO: re-enable tests/models/test_models_unet_2d_condition.py#test_model_xattn_padding
attention_mask = F.pad(attention_mask, (0, target_length), value=0.0)
if out_dim == 3:
if attention_mask.shape[0] < batch_size * head_size:
attention_mask = attention_mask.repeat_interleave(head_size, dim=0)
elif out_dim == 4:
attention_mask = attention_mask.unsqueeze(1)
attention_mask = attention_mask.repeat_interleave(head_size, dim=1)
return attention_mask
def norm_encoder_hidden_states(
self, encoder_hidden_states: torch.Tensor
) -> torch.Tensor:
r"""
Normalize the encoder hidden states. Requires `self.norm_cross` to be specified when constructing the
`Attention` class.
Args:
encoder_hidden_states (`torch.Tensor`): Hidden states of the encoder.
Returns:
`torch.Tensor`: The normalized encoder hidden states.
"""
assert (
self.norm_cross is not None
), "self.norm_cross must be defined to call self.norm_encoder_hidden_states"
if isinstance(self.norm_cross, nn.LayerNorm):
encoder_hidden_states = self.norm_cross(encoder_hidden_states)
elif isinstance(self.norm_cross, nn.GroupNorm):
# Group norm norms along the channels dimension and expects
# input to be in the shape of (N, C, *). In this case, we want
# to norm along the hidden dimension, so we need to move
# (batch_size, sequence_length, hidden_size) ->
# (batch_size, hidden_size, sequence_length)
encoder_hidden_states = encoder_hidden_states.transpose(1, 2)
encoder_hidden_states = self.norm_cross(encoder_hidden_states)
encoder_hidden_states = encoder_hidden_states.transpose(1, 2)
else:
assert False
return encoder_hidden_states
@torch.no_grad()
def fuse_projections(self, fuse=True):
is_cross_attention = self.cross_attention_dim != self.query_dim
device = self.to_q.weight.data.device
dtype = self.to_q.weight.data.dtype
if not is_cross_attention:
# fetch weight matrices.
concatenated_weights = torch.cat(
[self.to_q.weight.data, self.to_k.weight.data, self.to_v.weight.data]
)
in_features = concatenated_weights.shape[1]
out_features = concatenated_weights.shape[0]
# create a new single projection layer and copy over the weights.
self.to_qkv = self.linear_cls(
in_features, out_features, bias=False, device=device, dtype=dtype
)
self.to_qkv.weight.copy_(concatenated_weights)
else:
concatenated_weights = torch.cat(
[self.to_k.weight.data, self.to_v.weight.data]
)
in_features = concatenated_weights.shape[1]
out_features = concatenated_weights.shape[0]
self.to_kv = self.linear_cls(
in_features, out_features, bias=False, device=device, dtype=dtype
)
self.to_kv.weight.copy_(concatenated_weights)
self.fused_projections = fuse
class AttnProcessor:
r"""
Default processor for performing attention-related computations.
"""
def __call__(
self,
attn: Attention,
hidden_states: torch.FloatTensor,
encoder_hidden_states: Optional[torch.FloatTensor] = None,
attention_mask: Optional[torch.FloatTensor] = None,
) -> torch.Tensor:
residual = hidden_states
input_ndim = hidden_states.ndim
if input_ndim == 4:
batch_size, channel, height, width = hidden_states.shape
hidden_states = hidden_states.view(
batch_size, channel, height * width
).transpose(1, 2)
batch_size, sequence_length, _ = (
hidden_states.shape
if encoder_hidden_states is None
else encoder_hidden_states.shape
)
attention_mask = attn.prepare_attention_mask(
attention_mask, sequence_length, batch_size
)
if attn.group_norm is not None:
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(
1, 2
)
query = attn.to_q(hidden_states)
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
elif attn.norm_cross:
encoder_hidden_states = attn.norm_encoder_hidden_states(
encoder_hidden_states
)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
query = attn.head_to_batch_dim(query)
key = attn.head_to_batch_dim(key)
value = attn.head_to_batch_dim(value)
attention_probs = attn.get_attention_scores(query, key, attention_mask)
hidden_states = torch.bmm(attention_probs, value)
hidden_states = attn.batch_to_head_dim(hidden_states)
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
if input_ndim == 4:
hidden_states = hidden_states.transpose(-1, -2).reshape(
batch_size, channel, height, width
)
if attn.residual_connection:
hidden_states = hidden_states + residual
hidden_states = hidden_states / attn.rescale_output_factor
return hidden_states
class AttnProcessor2_0:
r"""
Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0).
"""
def __init__(self):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError(
"AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0."
)
def __call__(
self,
attn: Attention,
hidden_states: torch.FloatTensor,
encoder_hidden_states: Optional[torch.FloatTensor] = None,
attention_mask: Optional[torch.FloatTensor] = None,
) -> torch.FloatTensor:
residual = hidden_states
input_ndim = hidden_states.ndim
if input_ndim == 4:
batch_size, channel, height, width = hidden_states.shape
hidden_states = hidden_states.view(
batch_size, channel, height * width
).transpose(1, 2)
batch_size, sequence_length, _ = (
hidden_states.shape
if encoder_hidden_states is None
else encoder_hidden_states.shape
)
if attention_mask is not None:
attention_mask = attn.prepare_attention_mask(
attention_mask, sequence_length, batch_size
)
# scaled_dot_product_attention expects attention_mask shape to be
# (batch, heads, source_length, target_length)
attention_mask = attention_mask.view(
batch_size, attn.heads, -1, attention_mask.shape[-1]
)
if attn.group_norm is not None:
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(
1, 2
)
query = attn.to_q(hidden_states)
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
elif attn.norm_cross:
encoder_hidden_states = attn.norm_encoder_hidden_states(
encoder_hidden_states
)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
inner_dim = key.shape[-1]
head_dim = inner_dim // attn.heads
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
# the output of sdp = (batch, num_heads, seq_len, head_dim)
# TODO: add support for attn.scale when we move to Torch 2.1
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
hidden_states = hidden_states.transpose(1, 2).reshape(
batch_size, -1, attn.heads * head_dim
)
hidden_states = hidden_states.to(query.dtype)
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
if input_ndim == 4:
hidden_states = hidden_states.transpose(-1, -2).reshape(
batch_size, channel, height, width
)
if attn.residual_connection:
hidden_states = hidden_states + residual
hidden_states = hidden_states / attn.rescale_output_factor
return hidden_states
@@ -0,0 +1,334 @@
# Copyright 2023 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
# --------
#
# Modified 2024 by the Tripo AI and Stability AI Team.
#
# Copyright (c) 2024 Tripo AI & Stability AI
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
#
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
# SOFTWARE.
from typing import Optional
import torch
import torch.nn.functional as F
from torch import nn
from .attention import Attention
class BasicTransformerBlock(nn.Module):
r"""
A basic Transformer block.
Parameters:
dim (`int`): The number of channels in the input and output.
num_attention_heads (`int`): The number of heads to use for multi-head attention.
attention_head_dim (`int`): The number of channels in each head.
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention.
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
attention_bias (:
obj: `bool`, *optional*, defaults to `False`): Configure if the attentions should contain a bias parameter.
only_cross_attention (`bool`, *optional*):
Whether to use only cross-attention layers. In this case two cross attention layers are used.
double_self_attention (`bool`, *optional*):
Whether to use two self-attention layers. In this case no cross attention layers are used.
upcast_attention (`bool`, *optional*):
Whether to upcast the attention computation to float32. This is useful for mixed precision training.
norm_elementwise_affine (`bool`, *optional*, defaults to `True`):
Whether to use learnable elementwise affine parameters for normalization.
norm_type (`str`, *optional*, defaults to `"layer_norm"`):
The normalization layer to use. Can be `"layer_norm"`, `"ada_norm"` or `"ada_norm_zero"`.
final_dropout (`bool` *optional*, defaults to False):
Whether to apply a final dropout after the last feed-forward layer.
"""
def __init__(
self,
dim: int,
num_attention_heads: int,
attention_head_dim: int,
dropout=0.0,
cross_attention_dim: Optional[int] = None,
activation_fn: str = "geglu",
attention_bias: bool = False,
only_cross_attention: bool = False,
double_self_attention: bool = False,
upcast_attention: bool = False,
norm_elementwise_affine: bool = True,
norm_type: str = "layer_norm",
final_dropout: bool = False,
):
super().__init__()
self.only_cross_attention = only_cross_attention
assert norm_type == "layer_norm"
# Define 3 blocks. Each block has its own normalization layer.
# 1. Self-Attn
self.norm1 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine)
self.attn1 = Attention(
query_dim=dim,
heads=num_attention_heads,
dim_head=attention_head_dim,
dropout=dropout,
bias=attention_bias,
cross_attention_dim=cross_attention_dim if only_cross_attention else None,
upcast_attention=upcast_attention,
)
# 2. Cross-Attn
if cross_attention_dim is not None or double_self_attention:
# We currently only use AdaLayerNormZero for self attention where there will only be one attention block.
# I.e. the number of returned modulation chunks from AdaLayerZero would not make sense if returned during
# the second cross attention block.
self.norm2 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine)
self.attn2 = Attention(
query_dim=dim,
cross_attention_dim=(
cross_attention_dim if not double_self_attention else None
),
heads=num_attention_heads,
dim_head=attention_head_dim,
dropout=dropout,
bias=attention_bias,
upcast_attention=upcast_attention,
) # is self-attn if encoder_hidden_states is none
else:
self.norm2 = None
self.attn2 = None
# 3. Feed-forward
self.norm3 = nn.LayerNorm(dim, elementwise_affine=norm_elementwise_affine)
self.ff = FeedForward(
dim,
dropout=dropout,
activation_fn=activation_fn,
final_dropout=final_dropout,
)
# let chunk size default to None
self._chunk_size = None
self._chunk_dim = 0
def set_chunk_feed_forward(self, chunk_size: Optional[int], dim: int):
# Sets chunk feed-forward
self._chunk_size = chunk_size
self._chunk_dim = dim
def forward(
self,
hidden_states: torch.FloatTensor,
attention_mask: Optional[torch.FloatTensor] = None,
encoder_hidden_states: Optional[torch.FloatTensor] = None,
encoder_attention_mask: Optional[torch.FloatTensor] = None,
) -> torch.FloatTensor:
# Notice that normalization is always applied before the real computation in the following blocks.
# 0. Self-Attention
norm_hidden_states = self.norm1(hidden_states)
attn_output = self.attn1(
norm_hidden_states,
encoder_hidden_states=(
encoder_hidden_states if self.only_cross_attention else None
),
attention_mask=attention_mask,
)
hidden_states = attn_output + hidden_states
# 3. Cross-Attention
if self.attn2 is not None:
norm_hidden_states = self.norm2(hidden_states)
attn_output = self.attn2(
norm_hidden_states,
encoder_hidden_states=encoder_hidden_states,
attention_mask=encoder_attention_mask,
)
hidden_states = attn_output + hidden_states
# 4. Feed-forward
norm_hidden_states = self.norm3(hidden_states)
if self._chunk_size is not None:
# "feed_forward_chunk_size" can be used to save memory
if norm_hidden_states.shape[self._chunk_dim] % self._chunk_size != 0:
raise ValueError(
f"`hidden_states` dimension to be chunked: {norm_hidden_states.shape[self._chunk_dim]} has to be divisible by chunk size: {self._chunk_size}. Make sure to set an appropriate `chunk_size` when calling `unet.enable_forward_chunking`."
)
num_chunks = norm_hidden_states.shape[self._chunk_dim] // self._chunk_size
ff_output = torch.cat(
[
self.ff(hid_slice)
for hid_slice in norm_hidden_states.chunk(
num_chunks, dim=self._chunk_dim
)
],
dim=self._chunk_dim,
)
else:
ff_output = self.ff(norm_hidden_states)
hidden_states = ff_output + hidden_states
return hidden_states
class FeedForward(nn.Module):
r"""
A feed-forward layer.
Parameters:
dim (`int`): The number of channels in the input.
dim_out (`int`, *optional*): The number of channels in the output. If not given, defaults to `dim`.
mult (`int`, *optional*, defaults to 4): The multiplier to use for the hidden dimension.
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
final_dropout (`bool` *optional*, defaults to False): Apply a final dropout.
"""
def __init__(
self,
dim: int,
dim_out: Optional[int] = None,
mult: int = 4,
dropout: float = 0.0,
activation_fn: str = "geglu",
final_dropout: bool = False,
):
super().__init__()
inner_dim = int(dim * mult)
dim_out = dim_out if dim_out is not None else dim
linear_cls = nn.Linear
if activation_fn == "gelu":
act_fn = GELU(dim, inner_dim)
if activation_fn == "gelu-approximate":
act_fn = GELU(dim, inner_dim, approximate="tanh")
elif activation_fn == "geglu":
act_fn = GEGLU(dim, inner_dim)
elif activation_fn == "geglu-approximate":
act_fn = ApproximateGELU(dim, inner_dim)
self.net = nn.ModuleList([])
# project in
self.net.append(act_fn)
# project dropout
self.net.append(nn.Dropout(dropout))
# project out
self.net.append(linear_cls(inner_dim, dim_out))
# FF as used in Vision Transformer, MLP-Mixer, etc. have a final dropout
if final_dropout:
self.net.append(nn.Dropout(dropout))
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
for module in self.net:
hidden_states = module(hidden_states)
return hidden_states
class GELU(nn.Module):
r"""
GELU activation function with tanh approximation support with `approximate="tanh"`.
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
approximate (`str`, *optional*, defaults to `"none"`): If `"tanh"`, use tanh approximation.
"""
def __init__(self, dim_in: int, dim_out: int, approximate: str = "none"):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out)
self.approximate = approximate
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
if gate.device.type != "mps":
return F.gelu(gate, approximate=self.approximate)
# mps: gelu is not implemented for float16
return F.gelu(gate.to(dtype=torch.float32), approximate=self.approximate).to(
dtype=gate.dtype
)
def forward(self, hidden_states):
hidden_states = self.proj(hidden_states)
hidden_states = self.gelu(hidden_states)
return hidden_states
class GEGLU(nn.Module):
r"""
A variant of the gated linear unit activation function from https://arxiv.org/abs/2002.05202.
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
"""
def __init__(self, dim_in: int, dim_out: int):
super().__init__()
linear_cls = nn.Linear
self.proj = linear_cls(dim_in, dim_out * 2)
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
if gate.device.type != "mps":
return F.gelu(gate)
# mps: gelu is not implemented for float16
return F.gelu(gate.to(dtype=torch.float32)).to(dtype=gate.dtype)
def forward(self, hidden_states, scale: float = 1.0):
args = ()
hidden_states, gate = self.proj(hidden_states, *args).chunk(2, dim=-1)
return hidden_states * self.gelu(gate)
class ApproximateGELU(nn.Module):
r"""
The approximate form of Gaussian Error Linear Unit (GELU). For more details, see section 2:
https://arxiv.org/abs/1606.08415.
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
"""
def __init__(self, dim_in: int, dim_out: int):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.proj(x)
return x * torch.sigmoid(1.702 * x)
@@ -0,0 +1,219 @@
# Copyright 2023 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
# --------
#
# Modified 2024 by the Tripo AI and Stability AI Team.
#
# Copyright (c) 2024 Tripo AI & Stability AI
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
#
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
# SOFTWARE.
from dataclasses import dataclass
from typing import Optional
import torch
import torch.nn.functional as F
from torch import nn
from ...utils import BaseModule
from .basic_transformer_block import BasicTransformerBlock
class Transformer1D(BaseModule):
@dataclass
class Config(BaseModule.Config):
num_attention_heads: int = 16
attention_head_dim: int = 88
in_channels: Optional[int] = None
out_channels: Optional[int] = None
num_layers: int = 1
dropout: float = 0.0
norm_num_groups: int = 32
cross_attention_dim: Optional[int] = None
attention_bias: bool = False
activation_fn: str = "geglu"
only_cross_attention: bool = False
double_self_attention: bool = False
upcast_attention: bool = False
norm_type: str = "layer_norm"
norm_elementwise_affine: bool = True
gradient_checkpointing: bool = False
cfg: Config
def configure(self) -> None:
self.num_attention_heads = self.cfg.num_attention_heads
self.attention_head_dim = self.cfg.attention_head_dim
inner_dim = self.num_attention_heads * self.attention_head_dim
linear_cls = nn.Linear
# 2. Define input layers
self.in_channels = self.cfg.in_channels
self.norm = torch.nn.GroupNorm(
num_groups=self.cfg.norm_num_groups,
num_channels=self.cfg.in_channels,
eps=1e-6,
affine=True,
)
self.proj_in = linear_cls(self.cfg.in_channels, inner_dim)
# 3. Define transformers blocks
self.transformer_blocks = nn.ModuleList(
[
BasicTransformerBlock(
inner_dim,
self.num_attention_heads,
self.attention_head_dim,
dropout=self.cfg.dropout,
cross_attention_dim=self.cfg.cross_attention_dim,
activation_fn=self.cfg.activation_fn,
attention_bias=self.cfg.attention_bias,
only_cross_attention=self.cfg.only_cross_attention,
double_self_attention=self.cfg.double_self_attention,
upcast_attention=self.cfg.upcast_attention,
norm_type=self.cfg.norm_type,
norm_elementwise_affine=self.cfg.norm_elementwise_affine,
)
for d in range(self.cfg.num_layers)
]
)
# 4. Define output layers
self.out_channels = (
self.cfg.in_channels
if self.cfg.out_channels is None
else self.cfg.out_channels
)
self.proj_out = linear_cls(inner_dim, self.cfg.in_channels)
self.gradient_checkpointing = self.cfg.gradient_checkpointing
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
encoder_attention_mask: Optional[torch.Tensor] = None,
):
"""
The [`Transformer1DModel`] forward method.
Args:
hidden_states (`torch.LongTensor` of shape `(batch size, num latent pixels)` if discrete, `torch.FloatTensor` of shape `(batch size, channel, height, width)` if continuous):
Input `hidden_states`.
encoder_hidden_states ( `torch.FloatTensor` of shape `(batch size, sequence len, embed dims)`, *optional*):
Conditional embeddings for cross attention layer. If not given, cross-attention defaults to
self-attention.
attention_mask ( `torch.Tensor`, *optional*):
An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. If `1` the mask
is kept, otherwise if `0` it is discarded. Mask will be converted into a bias, which adds large
negative values to the attention scores corresponding to "discard" tokens.
encoder_attention_mask ( `torch.Tensor`, *optional*):
Cross-attention mask applied to `encoder_hidden_states`. Two formats supported:
* Mask `(batch, sequence_length)` True = keep, False = discard.
* Bias `(batch, 1, sequence_length)` 0 = keep, -10000 = discard.
If `ndim == 2`: will be interpreted as a mask, then converted into a bias consistent with the format
above. This bias will be added to the cross-attention scores.
Returns:
torch.FloatTensor
"""
# ensure attention_mask is a bias, and give it a singleton query_tokens dimension.
# we may have done this conversion already, e.g. if we came here via UNet2DConditionModel#forward.
# we can tell by counting dims; if ndim == 2: it's a mask rather than a bias.
# expects mask of shape:
# [batch, key_tokens]
# adds singleton query_tokens dimension:
# [batch, 1, key_tokens]
# this helps to broadcast it as a bias over attention scores, which will be in one of the following shapes:
# [batch, heads, query_tokens, key_tokens] (e.g. torch sdp attn)
# [batch * heads, query_tokens, key_tokens] (e.g. xformers or classic attn)
if attention_mask is not None and attention_mask.ndim == 2:
# assume that mask is expressed as:
# (1 = keep, 0 = discard)
# convert mask into a bias that can be added to attention scores:
# (keep = +0, discard = -10000.0)
attention_mask = (1 - attention_mask.to(hidden_states.dtype)) * -10000.0
attention_mask = attention_mask.unsqueeze(1)
# convert encoder_attention_mask to a bias the same way we do for attention_mask
if encoder_attention_mask is not None and encoder_attention_mask.ndim == 2:
encoder_attention_mask = (
1 - encoder_attention_mask.to(hidden_states.dtype)
) * -10000.0
encoder_attention_mask = encoder_attention_mask.unsqueeze(1)
# 1. Input
batch, _, seq_len = hidden_states.shape
residual = hidden_states
hidden_states = self.norm(hidden_states)
inner_dim = hidden_states.shape[1]
hidden_states = hidden_states.permute(0, 2, 1).reshape(
batch, seq_len, inner_dim
)
hidden_states = self.proj_in(hidden_states)
# 2. Blocks
for block in self.transformer_blocks:
if self.training and self.gradient_checkpointing:
hidden_states = torch.utils.checkpoint.checkpoint(
block,
hidden_states,
attention_mask,
encoder_hidden_states,
encoder_attention_mask,
use_reentrant=False,
)
else:
hidden_states = block(
hidden_states,
attention_mask=attention_mask,
encoder_hidden_states=encoder_hidden_states,
encoder_attention_mask=encoder_attention_mask,
)
# 3. Output
hidden_states = self.proj_out(hidden_states)
hidden_states = (
hidden_states.reshape(batch, seq_len, inner_dim)
.permute(0, 2, 1)
.contiguous()
)
output = hidden_states + residual
return output
+218
View File
@@ -0,0 +1,218 @@
import math
import os
from dataclasses import dataclass, field
from typing import List, Union
import numpy as np
import PIL.Image
import torch
import torch.nn.functional as F
import trimesh
from einops import rearrange
from huggingface_hub import hf_hub_download
from omegaconf import OmegaConf
from PIL import Image
from .models.isosurface import MarchingCubeHelper
from .utils import (
BaseModule,
ImagePreprocessor,
find_class,
get_spherical_cameras,
scale_tensor,
)
class TSR(BaseModule):
@dataclass
class Config(BaseModule.Config):
cond_image_size: int
image_tokenizer_cls: str
image_tokenizer: dict
tokenizer_cls: str
tokenizer: dict
backbone_cls: str
backbone: dict
post_processor_cls: str
post_processor: dict
decoder_cls: str
decoder: dict
renderer_cls: str
renderer: dict
cfg: Config
@classmethod
def from_pretrained(
cls, pretrained_model_name_or_path: str, config_name: str, weight_name: str
):
if os.path.isdir(pretrained_model_name_or_path):
config_path = os.path.join(pretrained_model_name_or_path, config_name)
weight_path = os.path.join(pretrained_model_name_or_path, weight_name)
else:
config_path = hf_hub_download(
repo_id=pretrained_model_name_or_path, filename=config_name
)
weight_path = hf_hub_download(
repo_id=pretrained_model_name_or_path, filename=weight_name
)
cfg = OmegaConf.load(config_path)
OmegaConf.resolve(cfg)
model = cls(cfg)
ckpt = torch.load(weight_path, map_location="cpu")
model.load_state_dict(ckpt)
return model
@classmethod
def from_pretrained_custom(
cls, weight_path: str, config_path: str
):
cfg = OmegaConf.load(config_path)
OmegaConf.resolve(cfg)
model = cls(cfg)
ckpt = torch.load(weight_path, map_location="cpu")
model.load_state_dict(ckpt)
return model
def configure(self):
self.image_tokenizer = find_class(self.cfg.image_tokenizer_cls)(
self.cfg.image_tokenizer
)
self.tokenizer = find_class(self.cfg.tokenizer_cls)(self.cfg.tokenizer)
self.backbone = find_class(self.cfg.backbone_cls)(self.cfg.backbone)
self.post_processor = find_class(self.cfg.post_processor_cls)(
self.cfg.post_processor
)
self.decoder = find_class(self.cfg.decoder_cls)(self.cfg.decoder)
self.renderer = find_class(self.cfg.renderer_cls)(self.cfg.renderer)
self.image_processor = ImagePreprocessor()
self.isosurface_helper = None
def forward(
self,
image: Union[
PIL.Image.Image,
np.ndarray,
torch.FloatTensor,
List[PIL.Image.Image],
List[np.ndarray],
List[torch.FloatTensor],
],
device: str,
) -> torch.FloatTensor:
rgb_cond = self.image_processor(image, self.cfg.cond_image_size)[:, None].to(
device
)
batch_size = rgb_cond.shape[0]
input_image_tokens: torch.Tensor = self.image_tokenizer(
rearrange(rgb_cond, "B Nv H W C -> B Nv C H W", Nv=1),
)
input_image_tokens = rearrange(
input_image_tokens, "B Nv C Nt -> B (Nv Nt) C", Nv=1
)
tokens: torch.Tensor = self.tokenizer(batch_size)
tokens = self.backbone(
tokens,
encoder_hidden_states=input_image_tokens,
)
scene_codes = self.post_processor(self.tokenizer.detokenize(tokens))
return scene_codes
def render(
self,
scene_codes,
n_views: int,
elevation_deg: float = 0.0,
camera_distance: float = 1.9,
fovy_deg: float = 40.0,
height: int = 256,
width: int = 256,
return_type: str = "pil",
):
rays_o, rays_d = get_spherical_cameras(
n_views, elevation_deg, camera_distance, fovy_deg, height, width
)
rays_o, rays_d = rays_o.to(scene_codes.device), rays_d.to(scene_codes.device)
def process_output(image: torch.FloatTensor):
if return_type == "pt":
return image
elif return_type == "np":
return image.detach().cpu().numpy()
elif return_type == "pil":
return Image.fromarray(
(image.detach().cpu().numpy() * 255.0).astype(np.uint8)
)
else:
raise NotImplementedError
images = []
for scene_code in scene_codes:
images_ = []
for i in range(n_views):
with torch.no_grad():
image = self.renderer(
self.decoder, scene_code, rays_o[i], rays_d[i]
)
images_.append(process_output(image))
images.append(images_)
return images
def set_marching_cubes_resolution(self, resolution: int):
if (
self.isosurface_helper is not None
and self.isosurface_helper.resolution == resolution
):
return
self.isosurface_helper = MarchingCubeHelper(resolution)
def extract_mesh(self, scene_codes, resolution: int = 256, threshold: float = 25.0,callback=None):
self.set_marching_cubes_resolution(resolution)
meshes = []
for scene_code in scene_codes:
with torch.no_grad():
density = self.renderer.query_triplane(
self.decoder,
scale_tensor(
self.isosurface_helper.grid_vertices.to(scene_codes.device),
self.isosurface_helper.points_range,
(-self.renderer.cfg.radius, self.renderer.cfg.radius),
),
scene_code,
)["density_act"]
v_pos, t_pos_idx = self.isosurface_helper(-(density - threshold))
v_pos = scale_tensor(
v_pos,
self.isosurface_helper.points_range,
(-self.renderer.cfg.radius, self.renderer.cfg.radius),
)
with torch.no_grad():
color = self.renderer.query_triplane(
self.decoder,
v_pos,
scene_code,
)["color"]
mesh = trimesh.Trimesh(
vertices=v_pos.cpu().numpy(),
faces=t_pos_idx.cpu().numpy(),
vertex_colors=color.cpu().numpy(),
)
meshes.append(mesh)
if callback:
callback(len(meshes))
return meshes
+475
View File
@@ -0,0 +1,475 @@
import importlib
import math
from collections import defaultdict
from dataclasses import dataclass
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
import imageio
import numpy as np
import PIL.Image
#import rembg
import torch
import torch.nn as nn
import torch.nn.functional as F
import trimesh
from omegaconf import DictConfig, OmegaConf
#from PIL import Image
def parse_structured(fields: Any, cfg: Optional[Union[dict, DictConfig]] = None) -> Any:
scfg = OmegaConf.merge(OmegaConf.structured(fields), cfg)
return scfg
def find_class(cls_string):
module_string = ".".join(cls_string.split(".")[:-1])
cls_name = cls_string.split(".")[-1]
module = importlib.import_module(module_string, package=None)
cls = getattr(module, cls_name)
return cls
def get_intrinsic_from_fov(fov, H, W, bs=-1):
focal_length = 0.5 * H / np.tan(0.5 * fov)
intrinsic = np.identity(3, dtype=np.float32)
intrinsic[0, 0] = focal_length
intrinsic[1, 1] = focal_length
intrinsic[0, 2] = W / 2.0
intrinsic[1, 2] = H / 2.0
if bs > 0:
intrinsic = intrinsic[None].repeat(bs, axis=0)
return torch.from_numpy(intrinsic)
class BaseModule(nn.Module):
@dataclass
class Config:
pass
cfg: Config # add this to every subclass of BaseModule to enable static type checking
def __init__(
self, cfg: Optional[Union[dict, DictConfig]] = None, *args, **kwargs
) -> None:
super().__init__()
self.cfg = parse_structured(self.Config, cfg)
self.configure(*args, **kwargs)
def configure(self, *args, **kwargs) -> None:
raise NotImplementedError
class ImagePreprocessor:
def convert_and_resize(
self,
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
size: int,
):
if isinstance(image, PIL.Image.Image):
image = torch.from_numpy(np.array(image).astype(np.float32) / 255.0)
elif isinstance(image, np.ndarray):
if image.dtype == np.uint8:
image = torch.from_numpy(image.astype(np.float32) / 255.0)
else:
image = torch.from_numpy(image)
elif isinstance(image, torch.Tensor):
pass
batched = image.ndim == 4
if not batched:
image = image[None, ...]
image = F.interpolate(
image.permute(0, 3, 1, 2),
(size, size),
mode="bilinear",
align_corners=False,
antialias=True,
).permute(0, 2, 3, 1)
if not batched:
image = image[0]
return image
def __call__(
self,
image: Union[
PIL.Image.Image,
np.ndarray,
torch.FloatTensor,
List[PIL.Image.Image],
List[np.ndarray],
List[torch.FloatTensor],
],
size: int,
) -> Any:
if isinstance(image, (np.ndarray, torch.FloatTensor)) and image.ndim == 4:
image = self.convert_and_resize(image, size)
else:
if not isinstance(image, list):
image = [image]
image = [self.convert_and_resize(im, size) for im in image]
image = torch.stack(image, dim=0)
return image
def rays_intersect_bbox(
rays_o: torch.Tensor,
rays_d: torch.Tensor,
radius: float,
near: float = 0.0,
valid_thresh: float = 0.01,
):
input_shape = rays_o.shape[:-1]
rays_o, rays_d = rays_o.view(-1, 3), rays_d.view(-1, 3)
rays_d_valid = torch.where(
rays_d.abs() < 1e-6, torch.full_like(rays_d, 1e-6), rays_d
)
if type(radius) in [int, float]:
radius = torch.FloatTensor(
[[-radius, radius], [-radius, radius], [-radius, radius]]
).to(rays_o.device)
radius = (
1.0 - 1.0e-3
) * radius # tighten the radius to make sure the intersection point lies in the bounding box
interx0 = (radius[..., 1] - rays_o) / rays_d_valid
interx1 = (radius[..., 0] - rays_o) / rays_d_valid
t_near = torch.minimum(interx0, interx1).amax(dim=-1).clamp_min(near)
t_far = torch.maximum(interx0, interx1).amin(dim=-1)
# check wheter a ray intersects the bbox or not
rays_valid = t_far - t_near > valid_thresh
t_near[torch.where(~rays_valid)] = 0.0
t_far[torch.where(~rays_valid)] = 0.0
t_near = t_near.view(*input_shape, 1)
t_far = t_far.view(*input_shape, 1)
rays_valid = rays_valid.view(*input_shape)
return t_near, t_far, rays_valid
def chunk_batch(func: Callable, chunk_size: int, *args, **kwargs) -> Any:
if chunk_size <= 0:
return func(*args, **kwargs)
B = None
for arg in list(args) + list(kwargs.values()):
if isinstance(arg, torch.Tensor):
B = arg.shape[0]
break
assert (
B is not None
), "No tensor found in args or kwargs, cannot determine batch size."
out = defaultdict(list)
out_type = None
# max(1, B) to support B == 0
for i in range(0, max(1, B), chunk_size):
out_chunk = func(
*[
arg[i : i + chunk_size] if isinstance(arg, torch.Tensor) else arg
for arg in args
],
**{
k: arg[i : i + chunk_size] if isinstance(arg, torch.Tensor) else arg
for k, arg in kwargs.items()
},
)
if out_chunk is None:
continue
out_type = type(out_chunk)
if isinstance(out_chunk, torch.Tensor):
out_chunk = {0: out_chunk}
elif isinstance(out_chunk, tuple) or isinstance(out_chunk, list):
chunk_length = len(out_chunk)
out_chunk = {i: chunk for i, chunk in enumerate(out_chunk)}
elif isinstance(out_chunk, dict):
pass
else:
print(
f"Return value of func must be in type [torch.Tensor, list, tuple, dict], get {type(out_chunk)}."
)
exit(1)
for k, v in out_chunk.items():
v = v if torch.is_grad_enabled() else v.detach()
out[k].append(v)
if out_type is None:
return None
out_merged: Dict[Any, Optional[torch.Tensor]] = {}
for k, v in out.items():
if all([vv is None for vv in v]):
# allow None in return value
out_merged[k] = None
elif all([isinstance(vv, torch.Tensor) for vv in v]):
out_merged[k] = torch.cat(v, dim=0)
else:
raise TypeError(
f"Unsupported types in return value of func: {[type(vv) for vv in v if not isinstance(vv, torch.Tensor)]}"
)
if out_type is torch.Tensor:
return out_merged[0]
elif out_type in [tuple, list]:
return out_type([out_merged[i] for i in range(chunk_length)])
elif out_type is dict:
return out_merged
ValidScale = Union[Tuple[float, float], torch.FloatTensor]
def scale_tensor(dat: torch.FloatTensor, inp_scale: ValidScale, tgt_scale: ValidScale):
if inp_scale is None:
inp_scale = (0, 1)
if tgt_scale is None:
tgt_scale = (0, 1)
if isinstance(tgt_scale, torch.FloatTensor):
assert dat.shape[-1] == tgt_scale.shape[-1]
dat = (dat - inp_scale[0]) / (inp_scale[1] - inp_scale[0])
dat = dat * (tgt_scale[1] - tgt_scale[0]) + tgt_scale[0]
return dat
def get_activation(name) -> Callable:
if name is None:
return lambda x: x
name = name.lower()
if name == "none":
return lambda x: x
elif name == "exp":
return lambda x: torch.exp(x)
elif name == "sigmoid":
return lambda x: torch.sigmoid(x)
elif name == "tanh":
return lambda x: torch.tanh(x)
elif name == "softplus":
return lambda x: F.softplus(x)
else:
try:
return getattr(F, name)
except AttributeError:
raise ValueError(f"Unknown activation function: {name}")
def get_ray_directions(
H: int,
W: int,
focal: Union[float, Tuple[float, float]],
principal: Optional[Tuple[float, float]] = None,
use_pixel_centers: bool = True,
normalize: bool = True,
) -> torch.FloatTensor:
"""
Get ray directions for all pixels in camera coordinate.
Reference: https://www.scratchapixel.com/lessons/3d-basic-rendering/
ray-tracing-generating-camera-rays/standard-coordinate-systems
Inputs:
H, W, focal, principal, use_pixel_centers: image height, width, focal length, principal point and whether use pixel centers
Outputs:
directions: (H, W, 3), the direction of the rays in camera coordinate
"""
pixel_center = 0.5 if use_pixel_centers else 0
if isinstance(focal, float):
fx, fy = focal, focal
cx, cy = W / 2, H / 2
else:
fx, fy = focal
assert principal is not None
cx, cy = principal
i, j = torch.meshgrid(
torch.arange(W, dtype=torch.float32) + pixel_center,
torch.arange(H, dtype=torch.float32) + pixel_center,
indexing="xy",
)
directions = torch.stack([(i - cx) / fx, -(j - cy) / fy, -torch.ones_like(i)], -1)
if normalize:
directions = F.normalize(directions, dim=-1)
return directions
def get_rays(
directions,
c2w,
keepdim=False,
normalize=False,
) -> Tuple[torch.FloatTensor, torch.FloatTensor]:
# Rotate ray directions from camera coordinate to the world coordinate
assert directions.shape[-1] == 3
if directions.ndim == 2: # (N_rays, 3)
if c2w.ndim == 2: # (4, 4)
c2w = c2w[None, :, :]
assert c2w.ndim == 3 # (N_rays, 4, 4) or (1, 4, 4)
rays_d = (directions[:, None, :] * c2w[:, :3, :3]).sum(-1) # (N_rays, 3)
rays_o = c2w[:, :3, 3].expand(rays_d.shape)
elif directions.ndim == 3: # (H, W, 3)
assert c2w.ndim in [2, 3]
if c2w.ndim == 2: # (4, 4)
rays_d = (directions[:, :, None, :] * c2w[None, None, :3, :3]).sum(
-1
) # (H, W, 3)
rays_o = c2w[None, None, :3, 3].expand(rays_d.shape)
elif c2w.ndim == 3: # (B, 4, 4)
rays_d = (directions[None, :, :, None, :] * c2w[:, None, None, :3, :3]).sum(
-1
) # (B, H, W, 3)
rays_o = c2w[:, None, None, :3, 3].expand(rays_d.shape)
elif directions.ndim == 4: # (B, H, W, 3)
assert c2w.ndim == 3 # (B, 4, 4)
rays_d = (directions[:, :, :, None, :] * c2w[:, None, None, :3, :3]).sum(
-1
) # (B, H, W, 3)
rays_o = c2w[:, None, None, :3, 3].expand(rays_d.shape)
if normalize:
rays_d = F.normalize(rays_d, dim=-1)
if not keepdim:
rays_o, rays_d = rays_o.reshape(-1, 3), rays_d.reshape(-1, 3)
return rays_o, rays_d
def get_spherical_cameras(
n_views: int,
elevation_deg: float,
camera_distance: float,
fovy_deg: float,
height: int,
width: int,
):
azimuth_deg = torch.linspace(0, 360.0, n_views + 1)[:n_views]
elevation_deg = torch.full_like(azimuth_deg, elevation_deg)
camera_distances = torch.full_like(elevation_deg, camera_distance)
elevation = elevation_deg * math.pi / 180
azimuth = azimuth_deg * math.pi / 180
# convert spherical coordinates to cartesian coordinates
# right hand coordinate system, x back, y right, z up
# elevation in (-90, 90), azimuth from +x to +y in (-180, 180)
camera_positions = torch.stack(
[
camera_distances * torch.cos(elevation) * torch.cos(azimuth),
camera_distances * torch.cos(elevation) * torch.sin(azimuth),
camera_distances * torch.sin(elevation),
],
dim=-1,
)
# default scene center at origin
center = torch.zeros_like(camera_positions)
# default camera up direction as +z
up = torch.as_tensor([0, 0, 1], dtype=torch.float32)[None, :].repeat(n_views, 1)
fovy = torch.full_like(elevation_deg, fovy_deg) * math.pi / 180
lookat = F.normalize(center - camera_positions, dim=-1)
right = F.normalize(torch.cross(lookat, up), dim=-1)
up = F.normalize(torch.cross(right, lookat), dim=-1)
c2w3x4 = torch.cat(
[torch.stack([right, up, -lookat], dim=-1), camera_positions[:, :, None]],
dim=-1,
)
c2w = torch.cat([c2w3x4, torch.zeros_like(c2w3x4[:, :1])], dim=1)
c2w[:, 3, 3] = 1.0
# get directions by dividing directions_unit_focal by focal length
focal_length = 0.5 * height / torch.tan(0.5 * fovy)
directions_unit_focal = get_ray_directions(
H=height,
W=width,
focal=1.0,
)
directions = directions_unit_focal[None, :, :, :].repeat(n_views, 1, 1, 1)
directions[:, :, :, :2] = (
directions[:, :, :, :2] / focal_length[:, None, None, None]
)
# must use normalize=True to normalize directions here
rays_o, rays_d = get_rays(directions, c2w, keepdim=True, normalize=True)
return rays_o, rays_d
# def remove_background(
# image: PIL.Image.Image,
# rembg_session: Any = None,
# force: bool = False,
# **rembg_kwargs,
# ) -> PIL.Image.Image:
# do_remove = True
# if image.mode == "RGBA" and image.getextrema()[3][0] < 255:
# do_remove = False
# do_remove = do_remove or force
# if do_remove:
# image = rembg.remove(image, session=rembg_session, **rembg_kwargs)
# return image
def resize_foreground(
image: PIL.Image.Image,
ratio: float,
) -> PIL.Image.Image:
image = np.array(image)
assert image.shape[-1] == 4
alpha = np.where(image[..., 3] > 0)
y1, y2, x1, x2 = (
alpha[0].min(),
alpha[0].max(),
alpha[1].min(),
alpha[1].max(),
)
# crop the foreground
fg = image[y1:y2, x1:x2]
# pad to square
size = max(fg.shape[0], fg.shape[1])
ph0, pw0 = (size - fg.shape[0]) // 2, (size - fg.shape[1]) // 2
ph1, pw1 = size - fg.shape[0] - ph0, size - fg.shape[1] - pw0
new_image = np.pad(
fg,
((ph0, ph1), (pw0, pw1), (0, 0)),
mode="constant",
constant_values=((0, 0), (0, 0), (0, 0)),
)
# compute padding according to the ratio
new_size = int(new_image.shape[0] / ratio)
# pad to size, double side
ph0, pw0 = (new_size - size) // 2, (new_size - size) // 2
ph1, pw1 = new_size - size - ph0, new_size - size - pw0
new_image = np.pad(
new_image,
((ph0, ph1), (pw0, pw1), (0, 0)),
mode="constant",
constant_values=((0, 0), (0, 0), (0, 0)),
)
new_image = PIL.Image.fromarray(new_image)
return new_image
def save_video(
frames: List[PIL.Image.Image],
output_path: str,
fps: int = 30,
):
# use imageio to save video
frames = [np.array(frame) for frame in frames]
writer = imageio.get_writer(output_path, fps=fps)
for frame in frames:
writer.append_data(frame)
writer.close()
def to_gradio_3d_orientation(mesh):
mesh.apply_transform(trimesh.transformations.rotation_matrix(-np.pi/2, [1, 0, 0]))
mesh.apply_scale([1, 1, -1])
mesh.apply_transform(trimesh.transformations.rotation_matrix(np.pi/2, [0, 1, 0]))
return mesh
+8 -2
View File
@@ -7,6 +7,12 @@ openai
simple-lama-inpainting
clip-interrogator==0.6.0
transformers>=4.36.0
zhipuai
lark-parser
imageio-ffmpeg
imageio-ffmpeg
rembg[gpu]
omegaconf==2.3.0
Pillow>=9.5.0
einops==0.7.0
trimesh>=4.0.5
huggingface-hub
scikit-image
+322 -53
View File
@@ -202,7 +202,7 @@
width: fit-content;
max-width: 100%;
margin-left: 12px;
min-height: 200px;
/*min-height: 200px;*/
}
.input_card {
@@ -278,8 +278,11 @@
}
button:hover {
border-color: yellow !important;
color: yellow !important;
/* border-color: yellow !important; */
color: #ffffff !important;
/* font-weight: 400; */
background: #232222;
cursor: pointer;
}
.disabled {
@@ -400,12 +403,58 @@
.hidden-caption-content {
display: none;
}
/* 定义滚动条轨道的背景颜色 */
::-webkit-scrollbar-track {
background-color: #f1f1f1;
/* 轨道背景颜色 */
}
/* 定义滚动条滑块的颜色 */
::-webkit-scrollbar-thumb {
background-color: #888;
/* 滑块颜色 */
border-radius: 4px;
/* 滑块圆角 */
}
/* 鼠标悬停时滚动条滑块的颜色 */
::-webkit-scrollbar-thumb:hover {
background-color: #555;
/* 悬停时滑块颜色 */
}
/* 定义滚动条角落的颜色 */
::-webkit-scrollbar-corner {
background-color: transparent;
/* 角落颜色 */
}
.dynamic_prompt::after {
content: attr(title);
position: absolute;
color: black;
padding: 4px;
padding-left: 25px;
border-radius: 4px;
}
summary{
user-select: none;
}
</style>
<!-- <script src="../../../scripts/api.js" type="module"></script> -->
<link href="/extensions/comfyui-mixlab-nodes/lib/photoswipe.min.css" rel="stylesheet">
<link href="/extensions/comfyui-mixlab-nodes/lib/classic.min.css" rel="stylesheet">
<script src="/extensions/comfyui-mixlab-nodes/lib/pickr.min.js"></script>
<script src="/extensions/comfyui-mixlab-nodes/lib/filerobot-image-editor.min.js"></script>
<script type="module" src="/extensions/comfyui-mixlab-nodes/lib/model-viewer.min.js"></script>
<link rel="stylesheet" href="/extensions/comfyui-mixlab-nodes/lib/login.css">
</head>
@@ -439,8 +488,39 @@
<script type="module">
import PhotoSwipeLightbox from '/extensions/comfyui-mixlab-nodes/lib/photoswipe-lightbox.esm.min.js'
// console.log(Lightbox)
import { api } from "../../../scripts/api.js";
// ComfyUI\web\extensions\core\dynamicPrompts.js
// 官方实现修改
// Allows for simple dynamic prompt replacement
// Inputs in the format {a|b} will have a random value of a or b chosen when the prompt is queued.
/*
* Strips C-style line and block comments from a string
*/
function stripComments(str) {
return str.replace(/\/\*[\s\S]*?\*\/|\/\/.*/g, '');
}
function dynamicPrompts(prompt) {
prompt = stripComments(prompt);
while (prompt.replace("\\{", "").includes("{") && prompt.replace("\\}", "").includes("}")) {
const startIndex = prompt.replace("\\{", "00").indexOf("{");
const endIndex = prompt.replace("\\}", "00").indexOf("}");
const optionsString = prompt.substring(startIndex + 1, endIndex);
const options = optionsString.split("|");
const randomIndex = Math.floor(Math.random() * options.length);
const randomOption = options[randomIndex];
prompt = prompt.substring(0, startIndex) + randomOption + prompt.substring(endIndex + 1);
}
return prompt
}
// console.log('api', api)
const base64Df =
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAwAAAAMCAYAAABWdVznAAAAAXNSR0IArs4c6QAAALZJREFUKFOFkLERwjAQBPdbgBkInECGaMLUQDsE0AkRVRAYWqAByxldPPOWHwnw4OBGye1p50UDSoA+W2ABLPN7i+C5dyC6R/uiAUXRQCs0bXoNIu4QPQzAxDKxHoALOrZcqtiyR/T6CXw7+3IGHhkYcy6BOR2izwT8LptG8rbMiCRAUb+CQ6WzQVb0SNOi5Z2/nX35DRyb/ENazhpWKoGwrpD6nICp5c2qogc4of+c7QcrhgF4Aa/aoAFHiL+RAAAAAElFTkSuQmCC'
@@ -498,6 +578,39 @@
})
}
//给load image to batch节点使用的输入
function createBase64ImageForLoadImageToBatch(imageElement, nodeId, bs) {
let im = new Image()
im.src = bs;
im.className = "base64"
imageElement.appendChild(im);
let base64s = imageElement.querySelectorAll('.base64')
//更新输入
window._appData.data[nodeId].inputs.images.base64 = Array.from(base64s, (b) => b.src)
// 删除
im.addEventListener('click', e => {
e.preventDefault();
im.remove();
let base64s = imageElement.querySelectorAll('.base64')
//更新输入
window._appData.data[nodeId].inputs.images.base64 = Array.from(base64s, (b) => b.src)
})
}
const blobToBase64 = blob => {
return new Promise((res, rej) => {
const reader = new FileReader()
reader.onloadend = () => {
const base64data = reader.result
res(base64data)
// 在这里可以将base64数据用于进一步处理或显示图片
}
reader.readAsDataURL(blob)
})
}
function base64ToBlob(base64) {
// 去除base64编码中的前缀
@@ -712,7 +825,7 @@
data[id].inputs.noise_seed = Math.round(Math.random() * max_seed)
}
// class_type:"Seed_"
if (data[id].class_type == "Seed_") {
if (data[id].class_type == "Seed_" && ['increment', 'decrement', 'randomize'].includes(seed[id])) {
data[id].inputs.seed = Math.round(Math.random() * max_seed)
}
console.log('new Seed', data[id])
@@ -737,6 +850,21 @@
// 随机seed
promptWorkflow = randomSeed(seed, promptWorkflow);
// //动态提示,改为输入的时候,手动触发
// for (const id in promptWorkflow) {
// let node = promptWorkflow[id]
// if (["TextInput_", "CLIPTextEncode", "PromptSimplification", "ChinesePrompt_Mix"].includes(
// node.class_type
// )) {
// if (node.class_type == "PromptSimplification") {
// promptWorkflow[id].inputs.prompt = dynamicPrompts(node.inputs.prompt);
// } else {
// promptWorkflow[id].inputs.text = dynamicPrompts(node.inputs.text);
// }
// console.log('#动态提示', promptWorkflow[id].inputs)
// }
// }
let url = get_url()
const data = JSON.stringify({ prompt: promptWorkflow, client_id });
fetch(`${url}/prompt`, {
@@ -912,6 +1040,9 @@
// copyImagesToClipboard(output_card.outerHTML)
})
//是否显示复制图片,复制html两个按钮
let isShowImageFn = false;
for (const node of outputData) {
// console.log('output', node)
if (node.class_type == "ShowTextForGPT") {
@@ -930,8 +1061,31 @@
output_card.appendChild(div);
};
if (["SaveImage", "PreviewImage", "PromptImage", "Image Save", "SaveImageAndMetadata_"].includes(node.class_type)) {
if (["SaveImage",
"PreviewImage",
"PromptImage",
"Image Save",
"SaveImageAndMetadata_",
"TransparentImage"].includes(node.class_type)) {
console.log('output#image', node)
let a = document.createElement('a');
a.id = `output_${node.id}`
a.setAttribute('data-pswp-width', "200");
a.setAttribute('data-pswp-height', "200");
a.setAttribute('target', "_blank");
a.setAttribute('href', base64Df);
a.setAttribute('title', node.title);
let img = new Image();
// img;
img.src = window._appData?.icon || base64Df;
a.appendChild(img)
output_card.appendChild(a);
isShowImageFn = true;
}
//3d
if (["SaveTripoSRMesh"].includes(node.class_type)) {
let a = document.createElement('a');
a.id = `output_${node.id}`
a.setAttribute('data-pswp-width', "200");
@@ -947,7 +1101,7 @@
}
// video ,gif
if (["VHS_VideoCombine"].includes(node.class_type)) {
if (["VHS_VideoCombine", "VideoCombine_Adv"].includes(node.class_type)) {
let a = document.createElement('a');
a.id = `output_${node.id}`
@@ -974,6 +1128,12 @@
}
}
if (isShowImageFn === false) {
copyImage.remove();
copyHTML.remove();
}
return container
}
@@ -1030,6 +1190,7 @@
}
async function handleClipboardImage(imageElement, data) {
//data.class_type === 'LoadImagesToBatch'
try {
const clipboardItems = await navigator.clipboard.read();
for (const clipboardItem of clipboardItems) {
@@ -1042,18 +1203,17 @@
if (hashId == window._appData.data[data.id].hashId) return
let { url, name } = await uploadImage(fileBlob);
// 在这里可以对 Blob 对象进行进一步处理
imageElement.src = url;
window._appData.data[data.id].inputs.image = name;
window._appData.data[data.id].hashId = hashId;
console.log("上传的文件:", url, data.id, name);
// const img = document.createElement('img');
// img.src = URL.createObjectURL(blob);
// document.body.appendChild(img);
// console.log( URL.createObjectURL(blob));
if (data.class_type === 'LoadImagesToBatch') {
let base64 = await blobToBase64(fileBlob)
createBase64ImageForLoadImageToBatch(imageElement, data.id, base64)
} else {
let { url, name } = await uploadImage(fileBlob);
// 在这里可以对 Blob 对象进行进一步处理
imageElement.src = url;
window._appData.data[data.id].inputs.image = name;
window._appData.data[data.id].hashId = hashId;
console.log("上传的文件:", url, data.id, name);
}
}
}
}
@@ -1217,10 +1377,16 @@
console.log('inputData', data);
// 图片 or 视频输入
if (data.class_type === "LoadImage" || data.class_type === "VHS_LoadVideo" || data.class_type === 'ImagesPrompt_') {
if (["LoadImage",
"VHS_LoadVideo",
"ImagesPrompt_",
"LoadImagesToBatch"].includes(data.class_type)) {
let isVideoUpload = data.class_type === "VHS_LoadVideo";
let isBase64Upload = data.class_type === "LoadImagesToBatch";
// Create a container for the upload control
const uploadContainer = document.createElement("div");
uploadContainer.className = 'card';
@@ -1252,7 +1418,7 @@
const btnForImageEdit = document.createElement("button");
btnForImageEdit.style = ` width: 32px; background: none;margin-left: 18px;`
btnForImageEdit.innerHTML = '<?xml version="1.0" ?><svg version="1.1" style="width: 24px;" viewBox="0 0 50 50" xml:space="preserve" xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink"><g id="Layer_1_1_"><path d="M18.293,31.707h6.414l24-24l-6.414-6.414l-24,24V31.707z M45.879,7.707l-3.586,3.586l-3.586-3.586l3.586-3.586 L45.879,7.707z M20.293,26.121l17-17l3.586,3.586l-17,17h-3.586V26.121z"/><polygon points="43.293,19.707 41.293,19.707 41.293,46.707 3.293,46.707 3.293,8.707 31.293,8.707 31.293,6.707 1.293,6.707 1.293,48.707 43.293,48.707 "/></g></svg>'
if (!isVideoUpload && data.class_type !== 'ImagesPrompt_') actionDiv.appendChild(btnForImageEdit);
if ((!isVideoUpload && !isBase64Upload) && data.class_type !== 'ImagesPrompt_') actionDiv.appendChild(btnForImageEdit);
uploadContainer.appendChild(actionDiv)
@@ -1293,11 +1459,22 @@
data.title,
data.options.images,
data.inputs.imageIndex,
(base64) => {
(base64, text) => {
window._appData.data[data.id].inputs.image_base64 = base64;
window._appData.data[data.id].inputs.text = text;
})
uploadContainer.appendChild(imgDiv);
window._appData.data[data.id].inputs.image_base64 = mainImage.querySelector('.images_prompt_main').src;
} else if (data.class_type === 'LoadImagesToBatch') {
// 多张base64 图片
let base64 = data.inputs.images.base64
imageElement = document.createElement('div');
imageElement.className = "images"
for (const bs of base64) {
createBase64ImageForLoadImageToBatch(imageElement, data.id, bs)
}
}
@@ -1305,7 +1482,7 @@
if (!isVideoUpload) btnFromClipboard.addEventListener('click', (event) => handleClipboardImage(imageElement, data));
if (!isVideoUpload) btnForImageEdit.addEventListener('click', e => editImage(imageElement, data))
if (!isVideoUpload && !isBase64Upload) btnForImageEdit.addEventListener('click', e => editImage(imageElement, data))
uploadImageInput.addEventListener('click', (event) => {
@@ -1328,33 +1505,43 @@
if (hashId == window._appData.data[data.id].hashId) return
let { url, name } = await uploadImage(fileBlob, '.' + file.type.split('/')[1])
if (data.class_type === 'ImagesPrompt_') {
//
let base64 = await parseImageToBase64(url);
uploadContainer.querySelector('.images_prompt_main').src = base64
window._appData.data[data.id].inputs.image_base64 = base64;
if (data.class_type === 'LoadImagesToBatch') {
// 上传 ,转为base64
let base64 = await blobToBase64(fileBlob)
createBase64ImageForLoadImageToBatch(imageElement, data.id, base64)
} else {
if (isVideoUpload) {
imageElement.srcObject = null;
}
// 在这里可以对 Blob 对象进行进一步处理
imageElement.src = url;
//上传,返回url
let { url, name } = await uploadImage(fileBlob, '.' + file.type.split('/')[1])
if (isVideoUpload) {
window._appData.data[data.id].inputs.video = name;
if (data.class_type === 'ImagesPrompt_') {
//
let base64 = await parseImageToBase64(url);
uploadContainer.querySelector('.images_prompt_main').src = base64
window._appData.data[data.id].inputs.image_base64 = base64;
} else {
window._appData.data[data.id].inputs.image = name;
if (isVideoUpload) {
imageElement.srcObject = null;
}
// 在这里可以对 Blob 对象进行进一步处理
imageElement.src = url;
if (isVideoUpload) {
window._appData.data[data.id].inputs.video = name;
} else {
window._appData.data[data.id].inputs.image = name;
}
}
window._appData.data[data.id].hashId = hashId;
console.log("上传的文件:", url, data.id, name);
}
window._appData.data[data.id].hashId = hashId;
console.log("上传的文件:", url, data.id, name);
};
// 开始读取文件
@@ -1441,7 +1628,7 @@
textInput.value = data.inputs.prompt;
} else {
textInput.value = data.inputs.text;
}
};
// uploadImageInput.type = "text";
let json = localStorage.getItem(`t_${data.id}`)
@@ -1457,7 +1644,6 @@
window._appData.data[data.id].inputs.text = textInput.value;
}
} catch (error) {
}
@@ -1465,6 +1651,27 @@
uploadContainer.appendChild(textInput);
// autoResize(textInput);
//动态提示功能
const dynamicPromptsBtn = document.createElement('button');
dynamicPromptsBtn.className = "dynamic_prompt"
dynamicPromptsBtn.innerText = 'dynamic';
dynamicPromptsBtn.style.width = '88px'
uploadContainer.appendChild(dynamicPromptsBtn);
dynamicPromptsBtn.addEventListener('click', e => {
e.preventDefault();
e.stopPropagation();
let prompt = dynamicPrompts(textInput.value)
textInput.setAttribute('title', prompt)
dynamicPromptsBtn.setAttribute('title', prompt)
if (data.class_type == "PromptSimplification") {
window._appData.data[data.id].inputs.prompt = prompt;
} else {
window._appData.data[data.id].inputs.text = prompt;
}
})
function autoResize(textarea) {
textarea.style.height = 'auto';
textarea.style.height = textarea.scrollHeight + 'px';
@@ -1836,7 +2043,7 @@
if (!opt.imgurl.match('data:image')) {
opt.imgurl = await parseImageToBase64(opt.imgurl)
}
callback(opt.imgurl);
callback(opt.imgurl, opt.keyword);
}
})
}
@@ -1907,6 +2114,7 @@
function createUI(data, share = true) {
// appData.input, appData.output, appData.seed, share, appData.link
if (!data) return
const { input: inputData, output: outputData, data: workflow, seed, seedTitle, link, name } = data;
let mainDiv = document.createElement('div');
@@ -2064,7 +2272,8 @@
// leftDiv.appendChild(des);
leftDiv.appendChild(statusDiv);
leftDiv.appendChild(input1);
mainDiv.appendChild(submitButton);
if (typeof (data.data) == 'object') mainDiv.appendChild(submitButton);
rightDiv.appendChild(output);
@@ -2199,6 +2408,51 @@
if (val && type == "text" && output.querySelector(`#output_${id}`)) output.querySelector(`#output_${id}`).innerText = val;
// 3d meshes
if (val && type == 'meshes' && output.querySelector(`#output_${id}`)) {
let threeD = output.querySelector('.threeD')
//判断默认的图片,需要去掉后创建model-viewer
let imgDf = output.querySelector(`#output_${id} img`);
if (!threeD) {
if (imgDf) imgDf.parentElement.remove();
threeD = document.createElement('div');
threeD.className = 'threeD'
threeD.id = `output_${id}`
output.querySelector('.output_card').appendChild(threeD)
// output.insertBefore(threeD, output.firstChild);
};
for (const meshUrl of val) {
const modelViewer = document.createElement('div');
modelViewer.style = `width:300px;margin:4px;height:300px;display:block`
modelViewer.innerHTML = `<model-viewer src="${meshUrl}"
min-field-of-view="0deg" max-field-of-view="180deg"
shadow-intensity="1"
camera-controls
touch-action="pan-y"
style="width:300px;height:300px;"
>
<div class="controls">
<button class="export">Save As</button>
</div>
</model-viewer>`
const btn = modelViewer.querySelector('.export');
btn.addEventListener('click', async e => {
e.preventDefault();
const glTF = await (modelViewer.querySelector('model-viewer')).exportScene()
const file = new File([glTF], 'mixlab.glb')
const link = document.createElement('a')
link.download = file.name
link.href = URL.createObjectURL(file)
link.click()
})
threeD.appendChild(modelViewer)
}
}
}
},
submitButton: {
@@ -2314,6 +2568,9 @@
const _images = detail?.output?._images;
const prompts = detail?.output?.prompts;
// 3d模型
const meshes = detail?.output?.mesh;
if (images) {
// if (!images) return;
@@ -2323,6 +2580,14 @@
return `${url}/view?filename=${encodeURIComponent(img.filename)}&type=${img.type}&subfolder=${encodeURIComponent(img.subfolder)}&t=${+new Date()}`;
}), detail.node, 'images');
} else if (meshes) {
//多个
let url = get_url();
show(Array.from(meshes, mesh => {
return `${url}/view?filename=${encodeURIComponent(mesh.filename)}&type=${mesh.type}&subfolder=${encodeURIComponent(mesh.subfolder)}&t=${+new Date()}`;
}), detail.node, 'meshes');
} else if (_images && prompts) {
let url = get_url();
@@ -2556,6 +2821,9 @@
// 创建app的选择菜单
function createAppList(apps = [], innerApp = false) {
window.prompt_ids = {};
if (document.body.querySelector('.apps')) {
document.body.querySelector('.apps').remove()
}
let details = document.createElement('details');
details.className = 'apps';
@@ -2756,20 +3024,21 @@
card.addEventListener('click', async e => {
e.preventDefault();
// console.log(c)
const { category, filename } = c.appInfo;
window._appData = (await get_my_app(category, filename))[0];
await createApp(window._appData);
executed(c.data, window._show);
try {
document.body.querySelector('#app_container').setAttribute('open', true)
document.body.querySelector('#app_input_pannel').removeAttribute('open')
document.body.querySelector('.apps').removeAttribute('open')
} catch (error) {
console.log(error)
}
const { category, filename } = c.appInfo;
window._appData = (await get_my_app(category, filename))[0];
createApp(window._appData);
executed(c.data, window._show);
})
}
+38 -4
View File
@@ -2,6 +2,24 @@ import { app } from '../../../scripts/app.js'
import { $el } from '../../../scripts/ui.js'
import { api } from '../../../scripts/api.js'
//本机安装的插件节点全集
window._nodesAll = null
//获取当前系统的插件,节点清单
function getObjectInfo () {
return new Promise(async (resolve, reject) => {
let url = getUrl()
try {
const response = await fetch(`${url}/object_info`)
const data = await response.json()
resolve(data)
} catch (error) {
reject(error)
}
})
}
const base64Df =
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAwAAAAMCAYAAABWdVznAAAAAXNSR0IArs4c6QAAALZJREFUKFOFkLERwjAQBPdbgBkInECGaMLUQDsE0AkRVRAYWqAByxldPPOWHwnw4OBGye1p50UDSoA+W2ABLPN7i+C5dyC6R/uiAUXRQCs0bXoNIu4QPQzAxDKxHoALOrZcqtiyR/T6CXw7+3IGHhkYcy6BOR2izwT8LptG8rbMiCRAUb+CQ6WzQVb0SNOi5Z2/nX35DRyb/ENazhpWKoGwrpD6nICp5c2qogc4of+c7QcrhgF4Aa/aoAFHiL+RAAAAAElFTkSuQmCC'
@@ -197,7 +215,8 @@ async function extractInputAndOutputData (
if (
node.type === 'KSampler' ||
node.type == 'SamplerCustom' ||
node.type === 'ChinesePrompt_Mix'
node.type === 'ChinesePrompt_Mix' ||
node.type === 'Seed_'
) {
// seed 的类型收集
try {
@@ -266,7 +285,10 @@ function downloadJsonFile (jsonData, fileName = 'mix_app.json') {
}
async function save (json, download = false, showInfo = true) {
console.log('####SAVE', json[0])
let nodesAll = window._nodesAll || (await getObjectInfo())
console.log('####SAVE', nodesAll, json[0])
const name = json[0],
version = json[5],
share_prefix = json[6], //用于分享的功能扩展
@@ -287,6 +309,13 @@ async function save (json, download = false, showInfo = true) {
try {
let data = await app.graphToPrompt()
//从output数据里把工作流的节点,插件数据统计出来
data.nodesMap = {}
for (const id in data.output) {
data.nodesMap[data.output[id].class_type] =
nodesAll[data.output[id].class_type]
}
let { input, output, seed, seedTitle } = await extractInputAndOutputData(
data,
inputIds,
@@ -349,11 +378,11 @@ async function save (json, download = false, showInfo = true) {
function getInputsAndOutputs () {
const inputs =
`LoadImage ImagesPrompt_ VHS_LoadVideo CLIPTextEncode PromptSlide TextInput_ Color FloatSlider IntNumber CheckpointLoaderSimple LoraLoader`.split(
`LoadImage LoadImagesToBatch ImagesPrompt_ VHS_LoadVideo CLIPTextEncode PromptSlide TextInput_ Color FloatSlider IntNumber CheckpointLoaderSimple LoraLoader`.split(
' '
),
outputs =
`PreviewImage,SaveImage,ShowTextForGPT,VHS_VideoCombine,Image Save,SaveImageAndMetadata_`.split(
`SaveTripoSRMesh,PreviewImage,SaveImage,TransparentImage,ShowTextForGPT,VHS_VideoCombine,VideoCombine_Adv,Image Save,SaveImageAndMetadata_,ClipInterrogator`.split(
','
)
@@ -378,6 +407,11 @@ function getInputsAndOutputs () {
app.registerExtension({
name: 'Mixlab.utils.AppInfo',
init () {
if (!window._nodesAll) {
getObjectInfo().then(r => (window._nodesAll = r))
}
},
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'AppInfo') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
+8 -2
View File
@@ -3,7 +3,7 @@ import { app } from '../../../scripts/app.js'
const repoOwner = 'shadowcz007' // 替换为仓库的所有者
const repoName = 'comfyui-mixlab-nodes' // 替换为仓库的名称
const version = 'v0.19.0'
const version = 'v0.24.0'
fetch(`https://api.github.com/repos/${repoOwner}/${repoName}/releases/latest`)
.then(response => response.json())
@@ -17,7 +17,13 @@ fetch(`https://api.github.com/repos/${repoOwner}/${repoName}/releases/latest`)
return
if (latestVersion && latestVersion != version) {
localStorage.setItem('_mixlab_nodes_vesion', latestVersion)
app.ui.dialog.show(`<h4 style="font-size: 18px;">${repoName} <br>
app.ui.dialog.show(`<a style="color: white;
font-size: 18px;
font-weight: 800;
letter-spacing: 2px;
}"
href="https://discord.gg/cXs9vZSqeK">Welcome to Mixlab nodes discord</a>
<h4 style="font-size: 18px;">${repoName} <br>
Latest release version: ${latestVersion}</h4>
<p>Please proceed to the official repository to download the latest version.</p>
<a style="color: #2196F3;
+1 -1
View File
@@ -61,7 +61,7 @@ app.registerExtension({
async getCustomWidgets (app) {
return {
KEY (node, inputName, inputData, app) {
console.log('##inputData', inputData)
// console.log('##inputData', inputData)
const widget = {
type: inputData[0], // the type, CHEESE
name: inputName, // the name, slice
+219 -2
View File
@@ -1,7 +1,39 @@
import { app } from '../../../scripts/app.js'
import { api } from '../../../scripts/api.js'
import { ComfyWidgets } from '../../../scripts/widgets.js'
// import { ComfyWidgets } from '../../../scripts/widgets.js'
import { $el } from '../../../scripts/ui.js'
import { applyTextReplacements } from '../../../scripts/utils.js'
function loadImageToCanvas (base64Image) {
var img = new Image()
var canvas = document.createElement('canvas')
var ctx = canvas.getContext('2d')
return new Promise((res, rej) => {
img.onload = function () {
// 等比例缩放图片
var width = img.width
var height = img.height
var max_width = 1024
if (width > max_width) {
height *= max_width / width
width = max_width
}
// 设置canvas尺寸
canvas.width = width
canvas.height = height
// 在canvas上绘制图片
ctx.drawImage(img, 0, 0, width, height)
// 将canvas转换为base64图片数据
var canvasData = canvas.toDataURL()
res(canvasData) // canvas转换后的base64图片数据
}
img.src = base64Image
})
}
async function uploadImage (blob, fileType = '.svg', filename) {
// const blob = await (await fetch(src)).blob();
@@ -618,9 +650,194 @@ app.registerExtension({
if (json && json[0]) {
uploadWidget.select.style.display = 'block'
createSelect(img, uploadWidget.select, json, prompt,text)
createSelect(img, uploadWidget.select, json, prompt, text)
}
} catch (error) {}
}
}
})
const createInputImageForBatch = (base64, widget) => {
let im = new Image()
im.src = base64
im.style = `width: 88px;`
im.addEventListener('click', e => {
let newValue = []
let items = widget.value?.base64 || []
for (const v of items) {
if (v != base64) newValue.push(v)
}
widget.value.base64 = newValue
im.remove()
})
return im
}
app.registerExtension({
name: 'Mixlab.Comfy.LoadImagesToBatch',
async getCustomWidgets (app) {
return {
IMAGEBASE64 (node, inputName, inputData, app) {
// console.log('##node', node)
const widget = {
value: {
base64: []
}, // 不能[x,x,x]
type: inputData[0], // the type
name: inputName, // the name, slice
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
}
// serializeValue (nodeId, widgetIndex) {
// return widget.value
// },
}
// widget.something = something; // maybe adds stuff to it
node.addCustomWidget(widget) // adds it to the node
return widget // and returns it.
}
}
},
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'LoadImagesToBatch') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
let imagesWidget = this.widgets.filter(w => w.name == 'images')[0]
const widget = {
type: 'div',
name: 'image_base64',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, 44, node.size[1])
)
},
serialize: false
}
widget.div = $el('div', {})
document.body.appendChild(widget.div)
let imagePreview = document.createElement('div')
let imagesDiv = document.createElement('div') //显示图片
imagesDiv.className = 'images_preview'
imagesDiv.style = `width: calc(100% - 14px);
display: flex;
flex-wrap: wrap;
padding: 7px; justify-content: space-between;
align-items: center;`
let inputImage = document.createElement('input')
inputImage.type = 'file'
inputImage.style.display = 'none'
inputImage.addEventListener('change', e => {
e.preventDefault()
const file = e.target.files[0]
const reader = new FileReader()
reader.onload = async event => {
let base64 = event.target.result
//压缩图片,控制1024以内
base64 = await loadImageToCanvas(base64)
// console.log(base64)
if (!imagesWidget.value) imagesWidget.value = { base64: [] }
imagesWidget.value.base64.push(base64)
let im = createInputImageForBatch(base64, imagesWidget)
imagesDiv.appendChild(im)
}
reader.readAsDataURL(file)
})
const btn = document.createElement('button')
btn.innerText = 'Upload Image'
btn.style = `cursor: pointer;
font-weight: 300;
margin: 2px;
color: var(--descrip-text);
background-color: var(--comfy-input-bg);
border-radius: 8px;
border-color: var(--border-color);
border-style: solid;height: 30px;min-width: 122px;
`
btn.addEventListener('click', e => {
e.preventDefault()
inputImage.click()
})
widget.div.appendChild(imagePreview)
imagePreview.appendChild(imagesDiv)
imagePreview.appendChild(btn)
imagePreview.appendChild(inputImage)
this.addCustomWidget(widget)
// document.addEventListener('wheel', handleMouseWheel)
const onRemoved = this.onRemoved
this.onRemoved = () => {
inputImage.remove()
widget.div.remove()
try {
// document.removeEventListener('wheel', handleMouseWheel)
} catch (error) {
console.log(error)
}
return onRemoved?.()
}
this.serialize_widgets = true //需要保存参数
}
}
if (nodeData.name === 'SaveImageAndMetadata_') {
const onNodeCreated = nodeType.prototype.onNodeCreated
// /web/extensions/core/saveImageExtraOutput.js
nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated
? onNodeCreated.apply(this, arguments)
: undefined
const widget = this.widgets.find(w => w.name === 'filename_prefix')
widget.serializeValue = () => {
return applyTextReplacements(app, widget.value)
}
return r
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments)
console.log('##onExecuted', this, message)
//TODO 是否 保存base64
if (message.base64) {
if (Array.isArray(message.base64)) {
}
}
}
}
},
async loadedGraphNode (node, app) {
if (node.type === 'LoadImagesToBatch') {
// await sleep(0)
let imagesWidget = node.widgets.filter(w => w.name === 'images')[0]
let imagePreview = node.widgets.filter(w => w.name == 'image_base64')[0]
let pre = imagePreview.div.querySelector('.images_preview')
for (const d of imagesWidget.value?.base64 || []) {
let im = createInputImageForBatch(d, imagesWidget)
pre.appendChild(im)
}
}
}
})
+203
View File
@@ -0,0 +1,203 @@
import { app } from '../../../scripts/app.js'
import { api } from '../../../scripts/api.js'
import { $el } from '../../../scripts/ui.js'
function get_position_style (ctx, widget_width, y, node_height) {
const MARGIN = 14 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
// height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
// outline: '1px solid red',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'space-around'
}
}
app.registerExtension({
name: 'Mixlab.3D.SaveTripoSRMesh',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeType.comfyClass == 'SaveTripoSRMesh') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = async function () {
orig_nodeCreated?.apply(this, arguments)
const widget = {
type: 'div',
name: 'preview',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, 88, node.size[1])
)
}
// value: [],
// async serializeValue (nodeId, widgetIndex) {
// return widget.value
// }
}
widget.div = $el('div', {})
widget.div.style.width = `120px`
document.body.appendChild(widget.div)
// preview.style = `margin-top: 12px;display: flex;
// justify-content: center;
// align-items: center;background-repeat: no-repeat;background-size: contain;`
this.addCustomWidget(widget)
const onResize = this.onResize
this.onResize = () => {
widget.div.style.width = `${this.size[0]}px`
widget.div.style.height = `${this.size[1] - 112}px`
let mvs = widget.div.querySelectorAll('model-viewer')
for (const m of mvs) {
m.style.height = `${Math.round(
(this.size[1] - 112) / mvs.length
)}px`
// console.log(m.style.height)
}
// console.log('resize', this.size)
return onResize?.apply(this, arguments)
}
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
return onRemoved?.()
}
if (this.onResize) {
this.onResize(this.size)
}
// this.isVirtualNode = true
this.serialize_widgets = false //需要保存参数
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
const r = onExecuted?.apply?.(this, arguments)
let widget = this.widgets.filter(d => d.name == 'preview')[0]
console.log('Test', widget, message)
let meshes = message.mesh
widget.div.innerHTML = ''
for (const mesh of meshes) {
if (mesh) {
const { filename, subfolder, type } = mesh
const fileURL = api.apiURL(
`/view?filename=${encodeURIComponent(
filename
)}&type=${type}&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
)
let modelViewer = document.createElement('div')
modelViewer.innerHTML = `<model-viewer src="${fileURL}"
min-field-of-view="0deg" max-field-of-view="180deg"
shadow-intensity="1"
camera-controls
touch-action="pan-y"
style="width:100%;margin:4px;min-height:88px"
>
<div class="controls">
<div><button class="export" style="
background-color: var(--comfy-input-bg);
border-radius: 8px;
border-color: var(--border-color);
border-style: solid;
color: var(--descrip-text);cursor: pointer;">Export GLB</button></div>
</div></model-viewer>`
widget.div.appendChild(modelViewer)
let modelViewerVariants= modelViewer
.querySelector('model-viewer');
modelViewer
.querySelector('.export')
.addEventListener('click', async e => {
e.preventDefault()
const glTF = await modelViewerVariants.exportScene()
const file = new File([glTF], filename)
const link = document.createElement('a')
link.download = file.name
link.href = URL.createObjectURL(file)
link.click()
})
}
}
// widget.value = [meshes]
this.onResize?.(this.size)
return r
}
}
},
async loadedGraphNode (node, app) {
const sleep = (t = 1000) => {
return new Promise((res, rej) => {
setTimeout(() => res(1), t)
})
}
// if (node.type === 'SaveTripoSRMesh') {
// await sleep(0)
// let widget = node.widgets.filter(w => w.name === 'preview')[0]
// widget.div.innerHTML = ''
// for (const mesh of widget.value) {
// if (mesh) {
// const { filename, subfolder, type } = mesh
// const fileURL = api.apiURL(
// `/view?filename=${encodeURIComponent(
// filename
// )}&type=${type}&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
// )
// let modelViewer = document.createElement('div')
// modelViewer.innerHTML = `<model-viewer src="${fileURL}"
// min-field-of-view="0deg" max-field-of-view="180deg"
// shadow-intensity="1"
// camera-controls
// touch-action="pan-y">
// <div class="controls">
// <div><button class="export">Export GLB</button></div>
// </div></model-viewer>`
// widget.div.appendChild(modelViewer)
// }
// }
// }
}
})
+28 -1
View File
@@ -46,6 +46,18 @@ const smart_connect_config_input = [
node_widget_name: 'image',
inputNodeName: 'LoadImage',
inputNode_output_name: 'IMAGE'
},
{
node_type: 'TripoSRSampler_',
node_widget_name: 'image',
inputNodeName: 'LoadImagesToBatch',
inputNode_output_name: 'IMAGE'
},
{
node_type: 'TripoSRSampler_',
node_widget_name: 'mask',
inputNodeName: 'RembgNode_Mix',
inputNode_output_name: 'masks'
}
]
@@ -74,6 +86,18 @@ const smart_connect_config_output = [
outputNodeName: 'SaveImage',
outputNode_input_name: 'images'
},
{
node_type: 'VAEDecode',
node_output_name: 'IMAGE',
outputNodeName: 'AppInfo',
outputNode_input_name: 'IMAGE'
},
{
node_type: 'VAEDecode',
node_output_name: 'IMAGE',
outputNodeName: 'SaveImageAndMetadata_',
outputNode_input_name: 'images'
},
{
node_type: 'Moondream',
node_output_name: 'STRING',
@@ -181,7 +205,10 @@ export function smart_init () {
]
let node_slotType = config[0]
// 如果input没有,则创建
if (!node.inputs?.filter(inp => inp.name === widget.name)[0]||!node.inputs)
if (
!node.inputs?.filter(inp => inp.name === widget.name)[0] ||
!node.inputs
)
convertToInput(node, widget, config)
input_node.connectByType(inputNode_slot, node, node_slotType)
}
+107 -13
View File
@@ -299,15 +299,16 @@ async function get_my_app (filename = null, category = '') {
data = []
for (const res of result.data) {
let { app, workflow } = res.data
if (app.filename)
data.push({
let { app, workflow } = res.data;
if (app?.filename) data.push({
...app,
data: workflow,
date: res.date
})
}
} catch (error) {}
} catch (error) {
console.log(error)
}
return data
}
@@ -1137,6 +1138,57 @@ app.registerExtension({
(this.canvas.height * 0.5) / (this.ds.scale * dpr) // 考虑设备像素比
this.setDirty(true, true)
}
// 支持app模式的json
const loadAppJson = async data => {
let workflow
try {
let w = JSON.parse(data)
if (w.app && w.output) workflow = w.workflow
} catch (err) {}
if (workflow && workflow.version && workflow.nodes && workflow.extra) {
await app.loadGraphData(workflow)
}
}
if (!window._mixlab_app_paste_listener) {
window._mixlab_app_paste_listener = true
//粘贴json的事件
document.addEventListener('paste', async e => {
// ctrl+shift+v is used to paste nodes with connections
// this is handled by litegraph
if (this.shiftDown) return
let data = e.clipboardData || window.clipboardData
// No image found. Look for node data
data = data.getData('text/plain')
loadAppJson(data)
})
// 把json往里 拖
document.addEventListener('drop', async event => {
event.preventDefault()
event.stopPropagation()
// Dragging from Chrome->Firefox there is a file but its a bmp, so ignore that
if (
event.dataTransfer.files.length &&
event.dataTransfer.files[0].type == 'application/json'
) {
const reader = new FileReader()
reader.onload = async () => {
loadAppJson(reader.result)
}
reader.readAsText(event.dataTransfer.files[0])
}
})
}
},
setup () {
setTimeout(async () => {
@@ -1146,6 +1198,8 @@ app.registerExtension({
const apps = await get_my_app()
if (!apps) return
console.log('apps',apps)
let apps_map = { 0: [] }
for (const app of apps) {
@@ -1159,7 +1213,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 => {
@@ -1511,14 +1565,16 @@ app.registerExtension({
document.body.appendChild(div)
}
},
{
content: 'Workflow App ♾️Mixlab',
has_submenu: true,
disabled: false,
submenu: {
options: apps_opts
}
}
apps_opts.length > 0
? {
content: 'Workflow App ♾️Mixlab',
has_submenu: true,
disabled: false,
submenu: {
options: apps_opts
}
}
: null
)
return options
@@ -1587,3 +1643,41 @@ app.registerExtension({
} catch (error) {}
}
})
//获取当前显存
function fetchSystemStats () {
return new Promise(async (resolve, reject) => {
try {
const response = await fetch('/system_stats')
const data = await response.json()
resolve(data)
} catch (error) {
reject(error)
}
})
}
//清理显存
function postFreeData () {
return new Promise(async (resolve, reject) => {
try {
const postData = {
unload_models: true,
free_memory: true
}
const response = await fetch('/free', {
method: 'POST',
headers: {
'Content-Type': 'application/json'
},
body: JSON.stringify(postData)
})
if (response.ok) {
resolve()
} else {
reject(new Error('Request failed'))
}
} catch (error) {
reject(error)
}
})
}
+515
View File
@@ -0,0 +1,515 @@
import { app } from '../../../scripts/app.js'
import { api } from '../../../scripts/api.js'
import { ComfyWidgets } from '../../../scripts/widgets.js'
import { $el } from '../../../scripts/ui.js'
// The code is based on ComfyUI-VideoHelperSuite modification.
function injectCSS (css) {
// 检查页面中是否已经存在具有相同内容的style标签
const existingStyle = document.querySelector('style')
if (existingStyle && existingStyle.textContent === css) {
return // 如果已经存在相同的样式,则不进行注入
}
// 创建一个新的style标签,并将CSS内容注入其中
const style = document.createElement('style')
style.textContent = css
// 将style标签插入到页面的head元素中
const head = document.querySelector('head')
head.appendChild(style)
}
injectCSS(`
.hidden{
display:none !important
}`)
function get_position_style (ctx, widget_width, y, node_height) {
const MARGIN = 4 // the margin around the html element
/* Create a transform that deals with all the scrolling and zooming */
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(MARGIN, MARGIN + y)
return {
transformOrigin: '0 0',
transform: transform,
left: `0`,
top: `0`,
cursor: 'pointer',
position: 'absolute',
maxWidth: `${widget_width - MARGIN * 2}px`,
// maxHeight: `${node_height - MARGIN * 2}px`, // we're assuming we have the whole height of the node
width: `${widget_width - MARGIN * 2}px`,
// height: `${node_height * 0.3 - MARGIN * 2}px`,
// background: '#EEEEEE',
display: 'flex',
flexDirection: 'column',
// alignItems: 'center',
justifyContent: 'space-around'
}
}
function videoUpload (node, inputName, inputData, app) {
const imageWidget = node.widgets.find(w => w.name === 'video')
let uploadWidget
const widget = {
type: 'div',
name: 'upload-preview',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, 220, node.size[1]),
{
outline: '1px solid'
}
)
}
}
widget.div = $el('div', {})
widget.div.style.width = `120px`
document.body.appendChild(widget.div)
node.addCustomWidget(widget)
// console.log('#imageWidget', imageWidget)
const displayDiv = document.createElement('video')
displayDiv.controls = true
// displayDiv.style=`width:200px;height:200px`
imageWidget.callback = () => {
displayDiv.src = `/view?filename=${
imageWidget.value
}&type=input&subfolder=${''}&rand=${Math.random()}`
// displayDiv.onloadedmetadata = function () {
// var frameCount = displayDiv.duration * displayDiv.webkitDecodedFrameCount
// console.log('视频帧数:' + frameCount)
// node.widgets.filter(w => w.name == 'video_segment_frames')[0].value =
// frameCount
// }
}
if (imageWidget.value) {
// console.log(imageWidget.value)
displayDiv.src = `/view?filename=${
imageWidget.value
}&type=input&subfolder=${''}&rand=${Math.random()}`
}
widget.div.appendChild(displayDiv)
const onRemoved = node.onRemoved
node.onRemoved = () => {
widget.div.remove()
return onRemoved?.()
}
var default_value = imageWidget.value
Object.defineProperty(imageWidget, 'value', {
set: function (value) {
this._real_value = value
},
get: function () {
let value = ''
if (this._real_value) {
value = this._real_value
} else {
return default_value
}
if (value.filename) {
let real_value = value
value = ''
if (real_value.subfolder) {
value = real_value.subfolder + '/'
}
value += real_value.filename
if (real_value.type && real_value.type !== 'input')
value += ` [${real_value.type}]`
}
return value
}
})
async function uploadFile (file, updateNode, pasted = false) {
try {
// Wrap file in formdata so it includes filename
const body = new FormData()
body.append('image', file)
if (pasted) body.append('subfolder', 'pasted')
const resp = await api.fetchApi('/upload/image', {
method: 'POST',
body
})
if (resp.status === 200) {
const data = await resp.json()
// Add the file to the dropdown list and update the widget value
let path = data.name
if (data.subfolder) path = data.subfolder + '/' + path
if (!imageWidget.options.values.includes(path)) {
imageWidget.options.values.push(path)
}
if (updateNode) {
imageWidget.value = path
}
return `/view?filename=${path}&type=input&subfolder=${
pasted ? 'pasted' : ''
}&rand=${Math.random()}`
} else {
alert(resp.status + ' - ' + resp.statusText)
}
} catch (error) {
alert(error)
}
}
const fileInput = document.createElement('input')
Object.assign(fileInput, {
type: 'file',
accept: 'video/*,.mkv,video/webm,video/mp4,video/x-matroska,image/gif',
style: 'display: none',
onchange: async () => {
if (fileInput.files.length) {
let file = fileInput.files[0]
const url = await uploadFile(file, true)
// console.log('fileInput', file)
var reader = new FileReader()
reader.onload = function () {
displayDiv.src = url
displayDiv.onloadedmetadata = function () {
// var frameCount =
// displayDiv.duration * displayDiv.webkitDecodedFrameCount
// console.log('视频帧数:' + frameCount)
// node.widgets.filter(
// w => w.name == 'video_segment_frames'
// )[0].value = frameCount
}
}
reader.readAsDataURL(file)
}
}
})
document.body.append(fileInput)
// Create the button widget for selecting the files
uploadWidget = node.addWidget('button', 'upload file', 'video', () => {
fileInput.click()
})
uploadWidget.serialize = false
return { widget: uploadWidget }
}
ComfyWidgets.VIDEOUPLOAD_ = videoUpload
app.registerExtension({
name: 'Mixlab.Video.LoadVideoAndSegment_',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeData?.name == 'LoadVideoAndSegment_') {
nodeData.input.required.upload = ['VIDEOUPLOAD_']
}
},
async loadedGraphNode (node, app) {
if (node.type === 'LoadVideoAndSegment_') {
const imageWidget = node.widgets.find(w => w.name === 'video')
const uploadPreview = node.widgets.find(w => w.name === 'upload-preview')
if (imageWidget.value) {
// console.log(imageWidget.value)
uploadPreview.div.querySelector('video').src = `/view?filename=${
imageWidget.value
}&type=input&subfolder=${''}&rand=${Math.random()}`
}
}
}
})
function offsetDOMWidget(
widget,
ctx,
node,
widgetWidth,
widgetY,
height
) {
const margin = 10
const elRect = ctx.canvas.getBoundingClientRect()
const transform = new DOMMatrix()
.scaleSelf(
elRect.width / ctx.canvas.width,
elRect.height / ctx.canvas.height
)
.multiplySelf(ctx.getTransform())
.translateSelf(0, widgetY + margin)
const scale = new DOMMatrix().scaleSelf(transform.a, transform.d)
Object.assign(widget.inputEl.style, {
transformOrigin: '0 0',
transform: scale,
left: `${transform.e}px`,
top: `${transform.d + transform.f}px`,
width: `${widgetWidth}px`,
height: `${(height || widget.parent?.inputHeight || 32) - margin}px`,
position: 'absolute',
background: !node.color ? '' : node.color,
color: !node.color ? '' : 'white',
zIndex: 5, //app.graph._nodes.indexOf(node),
})
}
export const hasWidgets = (node) => {
if (!node.widgets || !node.widgets?.[Symbol.iterator]) {
return false
}
return true
}
export const cleanupNode = (node) => {
if (!hasWidgets(node)) {
return
}
for (const w of node.widgets) {
if (w.canvas) {
w.canvas.remove()
}
if (w.inputEl) {
w.inputEl.remove()
}
// calls the widget remove callback
w.onRemoved?.()
}
}
const CreatePreviewElement = (name, val, format) => {
const [type] = format.split('/')
const w = {
name,
type,
value: val,
draw: function (ctx, node, widgetWidth, widgetY, height) {
const [cw, ch] = this.computeSize(widgetWidth)
offsetDOMWidget(this, ctx, node, widgetWidth, widgetY, ch)
},
computeSize: function (_) {
const ratio = this.inputRatio || 1
const width = Math.max(220, this.parent.size[0])
return [width, (width / ratio + 10)]
},
onRemoved: function () {
if (this.inputEl) {
this.inputEl.remove()
}
},
}
w.inputEl = document.createElement(type === 'video' ? 'video' : 'img')
w.inputEl.src = w.value
if (type === 'video') {
w.inputEl.setAttribute('type', 'video/webm');
w.inputEl.autoplay = true
w.inputEl.loop = true
w.inputEl.controls = false;
}
w.inputEl.onload = function () {
w.inputRatio = w.inputEl.naturalWidth / w.inputEl.naturalHeight
}
document.body.appendChild(w.inputEl)
return w
}
app.registerExtension({
name: 'Mixlab.Video.ImageListReplace',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeData?.name == 'ImageListReplace_') {
const orig_nodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
orig_nodeCreated?.apply(this, arguments)
const widget = {
type: 'div',
name: 'preview',
draw (ctx, node, widget_width, y, widget_height) {
Object.assign(
this.div.style,
get_position_style(ctx, widget_width, 188, node.size[1]),
{
outline: '1px solid',
display: 'flex',
flexWrap: 'wrap',
flexDirection: 'row',
justifyContent: 'flex-start'
}
)
}
}
widget.div = $el('div', {})
widget.div.style.width = `120px`
widget.div.className = 'hidden'
document.body.appendChild(widget.div)
this.addCustomWidget(widget)
// console.log('#ImageListReplace', widget)
const onRemoved = this.onRemoved
this.onRemoved = () => {
widget.div.remove()
return onRemoved?.()
}
}
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments)
// let _image_replace = message._image_replace[0]
// _image_replace = `/view?filename=${_image_replace.filename}&type=${
// _image_replace.type
// }&subfolder=${_image_replace.subfolder}&rand=${Math.random()}`
let preview = this.widgets.filter(w => w.name == 'preview')[0]
if (message._images.length > 0) {
preview.div.className = ''
// console.log('#ImageListReplace', preview.div)
}
preview.div.innerHTML = ''
for (const img_ of message._images) {
let img = new Image()
img.style = `width: 100px;
margin: 4px;`
img.src = `/view?filename=${img_.filename}&type=${
img_.type
}&subfolder=${img_.subfolder}&rand=${Math.random()}`
preview.div.appendChild(img)
}
let start_index = this.widgets.filter(w => w.name == 'start_index')[0]
let end_index = this.widgets.filter(w => w.name == 'end_index')[0]
let invert = this.widgets.filter(w => w.name == 'invert')[0]
let _sc = start_index.callback.bind(start_index)
let _ec = end_index.callback.bind(end_index)
const selectImages = () => {
// console.log(v)
let s = start_index.value,
e = end_index.value
let imgs = preview.div.querySelectorAll('img')
for (let index = 0; index < imgs.length; index++) {
if (invert.value) {
imgs[index].style.outline =
index >= s && index <= e ? 'none' : '4px solid #cbd3fe'
} else {
imgs[index].style.outline =
index >= s && index <= e ? '4px solid #cbd3fe' : 'none'
}
}
}
selectImages()
start_index.callback = v => {
let s = v,
e = end_index.value
let imgs = preview.div.querySelectorAll('img')
for (let index = 0; index < imgs.length; index++) {
if (invert.value) {
imgs[index].style.outline =
index >= s && index <= e ? 'none' : '4px solid #cbd3fe'
} else {
imgs[index].style.outline =
index >= s && index <= e ? '4px solid #cbd3fe' : 'none'
}
}
_sc(v)
}
end_index.callback = v => {
let s = start_index.value,
e = v
let imgs = preview.div.querySelectorAll('img')
for (let index = 0; index < imgs.length; index++) {
if (invert.value) {
imgs[index].style.outline =
index >= s && index <= e ? 'none' : '4px solid #cbd3fe'
} else {
imgs[index].style.outline =
index >= s && index <= e ? '4px solid #cbd3fe' : 'none'
}
}
_ec(v)
}
invert.callback = v => {
selectImages()
}
try {
} catch (error) {}
}
}
if (nodeData?.name == 'VideoCombine_Adv') {
const onExecuted = nodeType.prototype.onExecuted
nodeType.prototype.onExecuted = function (message) {
const prefix = 'vhs_gif_preview_'
const r = onExecuted ? onExecuted.apply(this, message) : undefined
if (this.widgets) {
const pos = this.widgets.findIndex(w => w.name === `${prefix}_0`)
if (pos !== -1) {
for (let i = pos; i < this.widgets.length; i++) {
this.widgets[i].onRemoved?.()
}
this.widgets.length = pos
}
if (message?.gifs) {
message.gifs.forEach((params, i) => {
const previewUrl = api.apiURL(
'/view?' + new URLSearchParams(params).toString()
)
const w = this.addCustomWidget(
CreatePreviewElement(
`${prefix}_${i}`,
previewUrl,
params.format || 'image/gif'
)
)
w.parent = this
})
}
const onRemoved = this.onRemoved
this.onRemoved = () => {
cleanupNode(this)
return onRemoved?.()
}
}
this.setSize([
this.size[0],
this.computeSize([this.size[0], this.size[1]])[1]
])
return r
}
}
}
})
-126
View File
@@ -1,126 +0,0 @@
import { app } from '../../../scripts/app.js'
import { api } from '../../../scripts/api.js'
import { ComfyWidgets } from '../../../scripts/widgets.js'
function videoUpload (node, inputName, inputData, app) {
const imageWidget = node.widgets.find(w => w.name === 'video')
let uploadWidget
const displayDiv = document.createElement('video')
console.log('imageWidget', node)
var default_value = imageWidget.value
Object.defineProperty(imageWidget, 'value', {
set: function (value) {
this._real_value = value
},
get: function () {
let value = ''
if (this._real_value) {
value = this._real_value
} else {
return default_value
}
if (value.filename) {
let real_value = value
value = ''
if (real_value.subfolder) {
value = real_value.subfolder + '/'
}
value += real_value.filename
if (real_value.type && real_value.type !== 'input')
value += ` [${real_value.type}]`
}
return value
}
})
async function uploadFile (file, updateNode, pasted = false) {
try {
// Wrap file in formdata so it includes filename
const body = new FormData()
body.append('image', file)
if (pasted) body.append('subfolder', 'pasted')
const resp = await api.fetchApi('/upload/image', {
method: 'POST',
body
})
if (resp.status === 200) {
const data = await resp.json()
// Add the file to the dropdown list and update the widget value
let path = data.name
if (data.subfolder) path = data.subfolder + '/' + path
if (!imageWidget.options.values.includes(path)) {
imageWidget.options.values.push(path)
}
if (updateNode) {
imageWidget.value = path
}
} else {
alert(resp.status + ' - ' + resp.statusText)
}
} catch (error) {
alert(error)
}
}
const fileInput = document.createElement('input')
Object.assign(fileInput, {
type: 'file',
accept: 'video/webm,video/mp4,video/mkv,image/gif',
style: 'display: none',
onchange: async () => {
if (fileInput.files.length) {
let file = fileInput.files[0]
console.log(file)
await uploadFile(file, true)
}
}
})
document.body.append(fileInput)
// Create the button widget for selecting the files
uploadWidget = node.addWidget('button', 'upload file', 'video', () => {
fileInput.click()
})
uploadWidget.serialize = false
return { widget: uploadWidget }
}
ComfyWidgets.VIDEOUPLOAD_ = videoUpload
app.registerExtension({
name: 'Mixlab.Video.LoadVideoAndSegment_',
async beforeRegisterNodeDef (nodeType, nodeData, app) {
if (nodeData?.name == 'LoadVideoAndSegment_') {
nodeData.input.required.upload = ['VIDEOUPLOAD_'];
// const onExecuted = nodeType.prototype.onExecuted
// nodeType.prototype.onExecuted = function (message) {
// onExecuted?.apply(this, arguments)
// console.log(message)
// // try {
// // let a = this.widgets.filter(w => w.name === 'AppInfoRun')[0]
// // if (a) {
// // if (!a.value) a.value = 0
// // a.value += 1
// // }
// // const div = this.widgets.filter(w => w.div)[0].div
// // Array.from(
// // div.querySelectorAll('button'),
// // b => (b.style.background = 'yellow')
// // )
// // } catch (error) {}
// }
}
}
})
-277
View File
@@ -1,277 +0,0 @@
* {
transition: all 0.6s cubic-bezier(0.77, 0, 0.175, 1);
}
#app-login {
width: 480px;
height: 90vh;
padding: 6vh;
background: white;
box-shadow: 0 0 2rem rgba(0, 0, 0, 0.1);
z-index: 999;
position: fixed;
top: 5vh;
left: calc(50vw - 240px);
}
.login-app-view {
position: absolute;
top: 0;
left: 0;
width: 100%;
height: 100%;
z-index: 999;
}
.login-background {
background-color: #202020e6;
position: fixed;
width: 100%;
height: 100vh;
left: 0;
top: 0;
z-index: 998;
}
.app-header {
padding: 6vh;
}
.app-header,
.app-header>* {
font-size: 1.2em;
margin: 0;
font-weight: 300;
}
.app-header>h1 {
font-size: 4.8vh;
font-weight: 400;
margin-bottom: 4.8vh;
}
.app-header>h2 {
font-size: 3vh;
}
.app-subheading {
color: rgba(0, 0, 0, 0.45);
}
.app-register {
position: absolute;
bottom: 0;
height: 10vh;
line-height: 10vh;
padding: 0 6vh;
color: rgba(0, 0, 0, 0.45);
}
.app-register>a {
font-weight: 400;
}
#app-login input {
font-size: 2.5vh;
width: calc(100% - 13vh);
height: 7.5vh;
margin-bottom: 2vh;
background: transparent;
position: absolute;
top: 0;
left: 6.5vh;
z-index: 2;
border: none;
box-shadow: inset 0 -0.5vh rgba(0, 0, 0, 0.1);
}
#app-login input:focus {
outline: none;
box-shadow: inset 0 -0.5vh transparent;
}
#app-login input[type=email] {
top: 58%;
}
#app-login input[type=password] {
top: calc(58% + 7.5vh);
}
#app-login input[type=email]:valid~* .st1 {
transition-timing-function: ease-in-out;
stroke-dasharray: 50, 153;
stroke-dashoffset: 25;
}
#app-login input[type=password]:focus~* .st0,
#app-login input[type=password]:valid~* .st0,
#login_run:focus~* .st0 {
stroke-dasharray: 210, 900;
stroke-dashoffset: -305;
}
#app-login input[type=email]:focus~* .st0 {
stroke-dasharray: 210, 900;
stroke-dashoffset: 0;
}
#app-login input:not(:valid)~#login_run {
/* pointer-events: none; */
opacity: 0.6;
}
#login_run {
text-decoration: none;
color: #0f9ede;
font-size: 1.5em;
padding: 0 6vh;
position: absolute;
bottom: 10vh;
font-weight: 400;
z-index: 998;
cursor: pointer;
}
#login_run:focus {
outline: none;
}
.login-app-view:nth-child(2) {
display: flex;
flex-direction: column;
pointer-events: none;
}
.login-app-view:nth-child(2)>.app-header {
font-size: 1rem;
flex-basis: 25%;
display: flex;
flex-direction: column;
justify-content: space-between;
padding: 4vh;
padding-bottom: 1rem;
}
.login-app-view:nth-child(2)>.app-header>h2 {
transform: translateY(1rem);
}
.login-app-view:nth-child(2)>.app-header>h2>em {
color: #0f9ede;
font-style: normal;
}
.login-app-view:nth-child(2)>.app-header>h2,
.login-app-view:nth-child(2) .app-item>*:not(.app-graphic) {
transition-duration: 0.9s;
opacity: 0;
}
.st0,
.st1,
.svg-loader-segment {
fill: none;
stroke: #0f9ede;
stroke-width: 0.5vh;
stroke-alignment: inside;
opacity: 1;
transition: all 0.6s cubic-bezier(0.77, 0, 0.175, 1);
}
.svg-loader {
opacity: 0;
}
.st0 {
stroke-dasharray: 0, 900;
stroke-dashoffset: 0;
}
.st1 {
transition-delay: 0.3s;
stroke-dasharray: 50, 153;
stroke-dashoffset: -153;
}
.svg-loader-segment {
transition: transform 1.2s cubic-bezier(0.77, 0, 0.175, 1), opacity 0.85s cubic-bezier(0.77, 0, 0.175, 1), stroke 0.85s cubic-bezier(0.77, 0, 0.175, 1);
}
#svg-lines {
position: absolute;
top: 45%;
left: 0;
width: 100%;
z-index: 0;
overflow: visible;
transform-origin: center 4vh;
}
.svg-data {
fill: none;
stroke-width: 0.5vh;
}
.svg-data.-temp {
stroke: #f4814b;
stroke-dasharray: 20, 118;
}
.svg-data.-cal {
stroke: #08b5cf;
stroke-dasharray: 20, 113;
}
.svg-data.-steps-bg {
stroke: #e0e1e0;
stroke-dasharray: 40, 100;
stroke-dashoffset: -60;
}
.svg-data.-steps {
stroke: #0f9ede;
stroke-dasharray: 20, 73;
stroke-dashoffset: -53;
}
.svg-data.-heart {
stroke: #9965aa;
stroke-dasharray: 50, 200;
stroke-dashoffset: -150;
}
.svg-activity-fill {
fill: #c4e4f8;
}
.svg-activity-line {
fill: none;
stroke: #65bcea;
stroke-miterlimit: 10;
stroke-width: 0.25vh;
}
.svg-activity-avg,
.svg-activity-indicator {
fill: none;
stroke: #d0dff0;
stroke-width: 0.25vh;
mix-blend-mode: multiply;
}
.svg-activity-fill,
.svg-activity-line {
transform: translateY(10vh);
opacity: 0;
}
*,
*:before,
*:after {
box-sizing: border-box;
position: relative;
}
-67
View File
@@ -1,67 +0,0 @@
;(() => {
let div = document.createElement('div')
div.innerHTML = `
<div id="app-login">
<div class="login-app-view">
<header class="app-header">
<h1>Hi</h1>
Welcome back,<br />
<span class="app-subheading">
sign in to continue<br />
</span>
</header>
<input class="email" type="email" required pattern=".*\.\w{2,}" placeholder="Email Address" />
<input class="password" type="password" required placeholder="Password" />
<a class="app-button" id="login_run">登录</a>
<!-- <div class="app-register">
Don't have an account? <a>Sign Up</a>
</div> -->
<svg id="svg-lines" version="1.1" xmlns="http://www.w3.org/2000/svg"
xmlns:xlink="http://www.w3.org/1999/xlink" x="0px" y="0px" viewBox="0 0 284.2 152.7"
xml:space="preserve">
<path class="st0"
d="M37.7,107.3h222.6c12,0,21.8,9.7,21.8,21.7s-9.7,21.8-21.8,21.8c0,0-203.6,0-222.6,0S2.2,138.6,2.2,103.3 c0-52,113.5-101.5,141-101.5c13.5,0,21.8,9.7,21.8,21.8s-9.7,21.7-21.8,21.7s-21.8-9.7-21.8-21.7s9.7-21.8,21.8-21.8" />
<path class="st1"
d="M260.2,76.3L250,87.8l-9-9c-6.2-6.2,2-24.7,17.2-24.7c15.2,0,23.9,17.7,23.9,29.7s-11.7,23.5-23.9,23.5h-10.2">
</path>
<g class="svg-loader" xmlns="http://www.w3.org/2000/svg">
<path class="svg-loader-segment -cal" d="M164.7,23.5c0-12-9.7-21.8-21.8-21.8" />
<path class="svg-loader-segment -heart" d="M143,45.2c12,0,21.8-9.7,21.8-21.7" />
<path class="svg-loader-segment -steps" d="M121.2,23.5c0,12,9.7,21.7,21.8,21.7" />
<path class="svg-loader-segment -temp" d="M143,1.7c-12,0-21.8,9.7-21.8,21.8" />
</g>
</svg>
</div>
</div>
<div class="login-background"></div>
`
document.body.appendChild(div)
let bg = div.querySelector('.login-background')
bg.addEventListener('click', e => {
div.style.display = 'none'
})
let login_btn = document.body.querySelector('#login_btn')
// login_btn.href="";
if (login_btn) {
login_btn.innerHTML =
'<svg stroke="currentColor" fill="none" stroke-width="0" viewBox="0 0 24 24" height="40px" width="40px" xmlns="http://www.w3.org/2000/svg"><path d="M12 17C14.2091 17 16 15.2091 16 13H8C8 15.2091 9.79086 17 12 17Z" fill="currentColor"></path><path d="M10 10C10 10.5523 9.55228 11 9 11C8.44772 11 8 10.5523 8 10C8 9.44772 8.44772 9 9 9C9.55228 9 10 9.44772 10 10Z" fill="currentColor"></path><path d="M15 11C15.5523 11 16 10.5523 16 10C16 9.44772 15.5523 9 15 9C14.4477 9 14 9.44772 14 10C14 10.5523 14.4477 11 15 11Z" fill="currentColor"></path><path fill-rule="evenodd" clip-rule="evenodd" d="M22 12C22 17.5228 17.5228 22 12 22C6.47715 22 2 17.5228 2 12C2 6.47715 6.47715 2 12 2C17.5228 2 22 6.47715 22 12ZM20 12C20 16.4183 16.4183 20 12 20C7.58172 20 4 16.4183 4 12C4 7.58172 7.58172 4 12 4C16.4183 4 20 7.58172 20 12Z" fill="currentColor"></path></svg>LOGIN'
login_btn.addEventListener('click', e => {
e.preventDefault()
div.style.display = 'block'
})
}
let login_run = div.querySelector('#login_run')
if (login_run) {
login_run.addEventListener('click', e => {
e.preventDefault()
let ps = div.querySelector('.password')
let email = div.querySelector('.email')
div.style.display = 'none'
console.log(ps.value, email.value)
})
}
})()
+81 -64
View File
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long