Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c92e43b920 | ||
|
|
be6c32b0e0 | ||
|
|
3d2062e810 | ||
|
|
dd816e95cd | ||
|
|
6d1b51890d | ||
|
|
11f03ec99a | ||
|
|
bd192f43e7 | ||
|
|
36e4b11983 | ||
|
|
97397ba8c2 | ||
|
|
052eee4111 | ||
|
|
e319496044 | ||
|
|
b83b63c362 | ||
|
|
4d6b1675bb | ||
|
|
13110fab39 | ||
|
|
74ea509848 | ||
|
|
192bff9d2c | ||
|
|
44ed8812dc | ||
|
|
200696ba21 | ||
|
|
51cf3b0c04 | ||
|
|
6a4831c83b | ||
|
|
c5e7ed95a3 | ||
|
|
42a97fa4d9 | ||
|
|
45240d0012 | ||
|
|
8ed085febd | ||
|
|
37803ea61b | ||
|
|
acd416952c | ||
|
|
6ec46cbc44 | ||
|
|
a9d971e476 | ||
|
|
9fe064675d | ||
|
|
c84fa467d0 | ||
|
|
3d7a55f6d3 | ||
|
|
d49baa1540 | ||
|
|
5c686af842 | ||
|
|
22425b5bc6 | ||
|
|
b6d9b338d2 | ||
|
|
a191a13751 | ||
|
|
3eccdbcc9b | ||
|
|
50063903f9 | ||
|
|
95a1b70533 | ||
|
|
c3679ac90b | ||
|
|
29e48eb6a2 | ||
|
|
71d02e9651 | ||
|
|
f5193f3eec | ||
|
|
480c4d6919 | ||
|
|
1c5e030540 | ||
|
|
27e83a5908 | ||
|
|
d3cbf8fa8d | ||
|
|
74b1f8129b | ||
|
|
56ed513cfd | ||
|
|
5f412371c4 | ||
|
|
c6b0b67585 | ||
|
|
16d18e681a | ||
|
|
1fe99f33b2 | ||
|
|
a2ece25ac0 | ||
|
|
4865f4d148 | ||
|
|
41e88824cf | ||
|
|
0961ab138e | ||
|
|
fe8271a12f | ||
|
|
f3866ede89 | ||
|
|
d938adf3cc | ||
|
|
9908cff64b | ||
|
|
1b9b0bb4e6 | ||
|
|
03acd9bea5 | ||
|
|
4c42949023 | ||
|
|
2e4d9836e5 | ||
|
|
43c6b58354 | ||
|
|
df637e8196 | ||
|
|
8e78f9786c | ||
|
|
f4130f06ed | ||
|
|
b766e714a4 | ||
|
|
1100a90be3 | ||
|
|
33aaf80c82 | ||
|
|
b5861dbc24 | ||
|
|
93731416fc | ||
|
|
a305e736ca | ||
|
|
a046ebbafb | ||
|
|
af65e96723 | ||
|
|
c9eb0ab5f0 | ||
|
|
9802e841a8 | ||
|
|
b896df8d54 | ||
|
|
720b8c237b | ||
|
|
371f9f813f | ||
|
|
33e229c41c | ||
|
|
2c33c0d801 | ||
|
|
d64fee5954 | ||
|
|
0217678c8c | ||
|
|
2959a9c31f | ||
|
|
1928a18992 | ||
|
|
a168171009 | ||
|
|
6f767f9700 | ||
|
|
e5459f63fd | ||
|
|
d0ab85a8c4 |
@@ -5,3 +5,4 @@ workflow/my_workflow.json
|
||||
workflow/my_workflow_app.json
|
||||
workflow/prompt_result.json
|
||||
app/*
|
||||
workflow/prompt_result.json
|
||||
|
||||
@@ -3,11 +3,15 @@
|
||||
> [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) -->
|
||||
|
||||
|
||||
## 🚀🚗🚚🏃 Workflow-to-APP
|
||||
@@ -15,7 +19,9 @@
|
||||
- 支持多个web app 切换
|
||||
- 发布为app的workflow,可以在右键里再次编辑了
|
||||
- web app可以设置分类,在comfyui右键菜单可以编辑更新web app
|
||||
- 支持动态提示
|
||||
|
||||

|
||||
|
||||
- Support multiple web app switching.
|
||||
- Add the AppInfo node, which allows you to transform the workflow into a web app by simple configuration.
|
||||
@@ -110,28 +116,39 @@ https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43
|
||||
|
||||
|
||||
### 3D
|
||||

|
||||

|
||||
[workflow](./assets/Image-to-3D_1.json)
|
||||
|
||||

|
||||
[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.
|
||||
|
||||

|
||||
|
||||
[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)
|
||||
|
||||

|
||||
|
||||
> 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)
|
||||
@@ -144,7 +161,7 @@ https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43
|
||||
|
||||
|
||||
|
||||
## Other Nodes
|
||||
### Other Nodes
|
||||
|
||||

|
||||

|
||||
@@ -191,15 +208,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
|
||||
|
||||
|
||||
+77
-23
@@ -263,7 +263,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 +273,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:
|
||||
@@ -576,9 +577,6 @@ async def post_prompt_result(request):
|
||||
|
||||
return web.json_response({"result":res})
|
||||
|
||||
|
||||
|
||||
|
||||
# 扩展api接口
|
||||
# from server import PromptServer
|
||||
# from aiohttp import web
|
||||
@@ -591,17 +589,22 @@ async def post_prompt_result(request):
|
||||
|
||||
|
||||
# 导入节点
|
||||
from .nodes.PromptNode import EmbeddingPrompt,RandomPrompt,PromptSlide,PromptSimplification,PromptImage,JoinWithDelimiter
|
||||
from .nodes.ImageNode import SaveImageToLocal,SplitImage,GridOutput,GetImageSize_,MirroredImage,ImageColorTransfer,NoiseImage,TransparentImage,GradientImage,LoadImagesFromPath,LoadImagesFromURL,ResizeImage,TextImage,SvgImage,Image3D,ShowLayer,NewLayer,MergeLayers,CenterImage,AreaToMask,SmoothMask,SplitLongMask,ImageCropByAlpha,EnhanceImage,FaceToMask
|
||||
from .nodes.PromptNode import GLIGENTextBoxApply_Advanced,EmbeddingPrompt,RandomPrompt,PromptSlide,PromptSimplification,PromptImage,JoinWithDelimiter
|
||||
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 CreateLoraNames,CreateSampler_names,CreateCkptNames,CreateSeedNode,TESTNODE_,TESTNODE_TOKEN,AppInfo,IntNumber,FloatSlider,TextInput,ColorInput,FontInput,TextToNumber,DynamicDelayProcessor,LimitNumber,SwitchByIndex,MultiplicationNode
|
||||
from .nodes.Mask import 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 VideoCombine_Adv,LoadVideoAndSegment,ImageListReplace,VAEEncodeForInpaint_Frames
|
||||
|
||||
from .nodes.TripoSR import LoadTripoSRModel,TripoSRSampler,SaveTripoSRMesh
|
||||
|
||||
from .nodes.Style import ApplyVisualStylePrompting
|
||||
|
||||
# 要导出的所有节点及其名称的字典
|
||||
# 注意:名称应全局唯一
|
||||
@@ -613,6 +616,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
# "LoraPrompt":LoraPrompt,
|
||||
"EmbeddingPrompt":EmbeddingPrompt,
|
||||
"PromptSlide":PromptSlide,
|
||||
"GLIGENTextBoxApply_Advanced":GLIGENTextBoxApply_Advanced,
|
||||
"PromptSimplification":PromptSimplification,
|
||||
"PromptImage":PromptImage,
|
||||
"MirroredImage":MirroredImage,
|
||||
@@ -622,6 +626,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"ResizeImageMixlab":ResizeImage,
|
||||
"LoadImagesFromPath":LoadImagesFromPath,
|
||||
"LoadImagesFromURL":LoadImagesFromURL,
|
||||
"LoadImagesToBatch":LoadImages_,
|
||||
"TextImage":TextImage,
|
||||
"EnhanceImage":EnhanceImage,
|
||||
"SvgImage":SvgImage,
|
||||
@@ -629,9 +634,12 @@ NODE_CLASS_MAPPINGS = {
|
||||
"ImageColorTransfer":ImageColorTransfer,
|
||||
"ShowLayer":ShowLayer,
|
||||
"NewLayer":NewLayer,
|
||||
"CompositeImages_":CompositeImages,
|
||||
"SplitImage":SplitImage,
|
||||
"CenterImage":CenterImage,
|
||||
"GridOutput":GridOutput,
|
||||
"GridDisplayAndSave":GridDisplayAndSave,
|
||||
"GridInput":GridInput,
|
||||
"MergeLayers":MergeLayers,
|
||||
"SplitLongMask":SplitLongMask,
|
||||
"FeatheredMask":FeatheredMask,
|
||||
@@ -639,8 +647,10 @@ NODE_CLASS_MAPPINGS = {
|
||||
"FaceToMask":FaceToMask,
|
||||
"AreaToMask":AreaToMask,
|
||||
"ImageCropByAlpha":ImageCropByAlpha,
|
||||
"ImagesPrompt_":ImagesPrompt,
|
||||
# "VAELoaderConsistencyDecoder":VAELoader,
|
||||
"SaveImageToLocal":SaveImageToLocal,
|
||||
"SaveImageAndMetadata_":SaveImageAndMetadata,
|
||||
# "VAEDecodeConsistencyDecoder":VAEDecode,
|
||||
"ScreenShare":ScreenShareNode,
|
||||
"FloatingVideo":FloatingVideo,
|
||||
@@ -662,41 +672,85 @@ NODE_CLASS_MAPPINGS = {
|
||||
"SwitchByIndex":SwitchByIndex,
|
||||
"LimitNumber":LimitNumber,
|
||||
"OutlineMask":OutlineMask,
|
||||
"MaskListMerge_":MaskListMerge,
|
||||
"JoinWithDelimiter":JoinWithDelimiter,
|
||||
"Seed_":CreateSeedNode,
|
||||
"CkptNames_":CreateCkptNames,
|
||||
"SamplerNames_":CreateSampler_names,
|
||||
"LoraNames_":CreateLoraNames,
|
||||
"ApplyVisualStylePrompting_":ApplyVisualStylePrompting
|
||||
# "LaMaInpainting":LaMaInpainting
|
||||
"ApplyVisualStylePrompting_":ApplyVisualStylePrompting,
|
||||
"StyleAlignedReferenceSampler_": StyleAlignedReferenceSampler,
|
||||
"StyleAlignedSampleReferenceLatents_": StyleAlignedSampleReferenceLatents,
|
||||
"StyleAlignedBatchAlign_": StyleAlignedBatchAlign,
|
||||
"LoadVideoAndSegment_":LoadVideoAndSegment,
|
||||
"VideoCombine_Adv":VideoCombine_Adv,
|
||||
"ListSplit_":ListSplit,
|
||||
"MaskListReplace_":MaskListReplace,
|
||||
"ImageListReplace_":ImageListReplace,
|
||||
"VAEEncodeForInpaint_Frames":VAEEncodeForInpaint_Frames,
|
||||
"IncrementingListNode_":IncrementingListNode,
|
||||
"PreviewMask_":PreviewMask_,
|
||||
"LoadTripoSRModel_": LoadTripoSRModel,
|
||||
"TripoSRSampler_": TripoSRSampler,
|
||||
"SaveTripoSRMesh": SaveTripoSRMesh
|
||||
# "GamePal":GamePal
|
||||
}
|
||||
|
||||
# 一个包含节点友好/可读的标题的字典
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"AppInfo":"AppInfo ♾️Mixlab",
|
||||
"ResizeImageMixlab":"ResizeImage ♾️Mixlab",
|
||||
"AppInfo":"App Info ♾️MixlabApp",
|
||||
"Color":"Color Input ♾️MixlabApp",
|
||||
"TextInput_":"Text Input ♾️MixlabApp",
|
||||
"FloatSlider":"Float Slider Input ♾️MixlabApp",
|
||||
"IntNumber":"Int Input ♾️MixlabApp",
|
||||
"ImagesPrompt_":"Images Input ♾️MixlabApp",
|
||||
"SaveImageAndMetadata_":"Save Image Output ♾️MixlabApp",
|
||||
"ResizeImageMixlab":"Resize Image ♾️Mixlab",
|
||||
"RandomPrompt": "Random Prompt ♾️Mixlab",
|
||||
"PromptImage":"Output Prompt and Image ♾️Mixlab",
|
||||
"SplitLongMask":"Splitting a long image into sections",
|
||||
"VAELoaderConsistencyDecoder":"Consistency Decoder Loader",
|
||||
"VAEDecodeConsistencyDecoder":"Consistency Decoder Decode",
|
||||
"ScreenShare":"ScreenShare ♾️Mixlab",
|
||||
"ScreenShare":"Screen Share ♾️Mixlab",
|
||||
"FloatingVideo":"FloatingVideo ♾️Mixlab",
|
||||
"ChatGPTOpenAI":"ChatGPT ♾️Mixlab",
|
||||
"ShowTextForGPT":"ShowTextForGPT ♾️Mixlab",
|
||||
"MergeLayers":"MergeLayers ♾️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":"PromptSlide ♾️Mixlab",
|
||||
"PromptGenerate_Mix":"PromptGenerate ♾️Mixlab",
|
||||
"ChinesePrompt_Mix":"ChinesePrompt ♾️Mixlab",
|
||||
"PromptSlide":"Prompt Slide ♾️Mixlab",
|
||||
"PromptGenerate_Mix":"Prompt Generate ♾️Mixlab",
|
||||
"ChinesePrompt_Mix":"Chinese Prompt ♾️Mixlab",
|
||||
"GamePal":"GamePal ♾️Mixlab",
|
||||
"RembgNode_Mix":"Removebg",
|
||||
"LoraNames_":"LoraName",
|
||||
"ApplyVisualStylePrompting_":"Apply VisualStyle Prompting"
|
||||
"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 ♾️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的节点功能
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 135 KiB |
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 |
@@ -0,0 +1,10 @@
|
||||
[
|
||||
{
|
||||
"keyword":"Dog",
|
||||
"imgurl":"http://127.0.0.1:8188/view?filename=1709966910233.png&type=input&subfolder=&rand=0.2734446552394221"
|
||||
},
|
||||
{
|
||||
"keyword":"x",
|
||||
"imgurl":"http://127.0.0.1:8188/view?filename=image%20(33).png&type=input&subfolder=pasted&rand=0.6984318219852814"
|
||||
}
|
||||
]
|
||||
+57
-15
@@ -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,14 +57,36 @@ 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
|
||||
|
||||
|
||||
|
||||
def chat(client, model_name,messages ):
|
||||
|
||||
try_count = 0
|
||||
@@ -105,7 +138,15 @@ 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-35-turbo","gpt-3.5-turbo-16k", "gpt-3.5-turbo-16k-0613", "gpt-4-0613","gpt-4-1106-preview","glm-4"],
|
||||
"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"}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "step": 1}),
|
||||
"context_size":("INT", {"default": 1, "min": 0, "max":30, "step": 1}),
|
||||
@@ -119,7 +160,7 @@ class ChatGPTNode:
|
||||
RETURN_TYPES = ("STRING","STRING","STRING",)
|
||||
RETURN_NAMES = ("text","messages","session_history",)
|
||||
FUNCTION = "generate_contextual_text"
|
||||
CATEGORY = "♾️Mixlab/Prompt/GPT"
|
||||
CATEGORY = "♾️Mixlab/GPT"
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,False,False,)
|
||||
|
||||
@@ -209,7 +250,7 @@ class ShowTextForGPT:
|
||||
OUTPUT_NODE = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt/GPT"
|
||||
CATEGORY = "♾️Mixlab/Text"
|
||||
|
||||
def run(self, text,output_dir=[""]):
|
||||
|
||||
@@ -293,7 +334,7 @@ class CharacterInText:
|
||||
# OUTPUT_NODE = True
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt/GPT"
|
||||
CATEGORY = "♾️Mixlab/Text"
|
||||
|
||||
def run(self, text,character,start_index):
|
||||
# print(text,character,start_index)
|
||||
@@ -306,8 +347,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
|
||||
@@ -338,15 +379,16 @@ class TextSplitByDelimiter:
|
||||
# OUTPUT_NODE = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt/GPT"
|
||||
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,)
|
||||
|
||||
+561
-40
@@ -8,16 +8,90 @@ 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):
|
||||
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']
|
||||
)
|
||||
|
||||
bg_image=bg_image.convert('RGB')
|
||||
|
||||
return bg_image
|
||||
|
||||
|
||||
def count_files_in_directory(directory):
|
||||
file_count = 0
|
||||
for _, _, files in os.walk(directory):
|
||||
file_count += len(files)
|
||||
return file_count
|
||||
|
||||
def save_json_to_file(data, file_path):
|
||||
with open(file_path, 'w') as file:
|
||||
json.dump(data, file)
|
||||
|
||||
def draw_rectangle(image, grid, color,width):
|
||||
x, y, w, h = grid
|
||||
draw = ImageDraw.Draw(image)
|
||||
draw.rectangle([(x, y), (x+w, y+h)], outline=color,width=width)
|
||||
|
||||
def generate_random_string(length):
|
||||
letters = string.ascii_letters + string.digits
|
||||
return ''.join(random.choice(letters) for _ in range(length))
|
||||
|
||||
def padding_rectangle(grid, padding):
|
||||
x, y, w, h = grid
|
||||
x -= padding
|
||||
y -= padding
|
||||
w += 2 * padding
|
||||
h += 2 * padding
|
||||
return (x, y, w, h)
|
||||
|
||||
class AnyType(str):
|
||||
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
|
||||
|
||||
@@ -559,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)
|
||||
@@ -592,7 +666,46 @@ 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")
|
||||
|
||||
@@ -621,13 +734,43 @@ 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:
|
||||
# 在底图上粘贴图层
|
||||
bg_image.paste(layer_image, (x, y), mask=mask)
|
||||
|
||||
# 输出合成后的图片
|
||||
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")
|
||||
@@ -657,8 +800,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
|
||||
|
||||
|
||||
@@ -982,6 +1127,35 @@ class TransparentImage:
|
||||
|
||||
|
||||
|
||||
class ImagesPrompt:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
# input_dir = folder_paths.get_input_directory()
|
||||
# files = [f for f in os.listdir(input_dir) if os.path.isfile(os.path.join(input_dir, f))]
|
||||
return {
|
||||
"required": {
|
||||
"image_base64": ("STRING",{"multiline": False,"default": "","dynamicPrompts": False}),
|
||||
"text": ("STRING",{"multiline": True,"default": "","dynamicPrompts": True}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE","STRING",)
|
||||
RETURN_NAMES = ("image","text",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,False,)
|
||||
OUTPUT_NODE = False
|
||||
|
||||
# 运行的函数
|
||||
def run(self,image_base64,text):
|
||||
image = base64_to_image(image_base64)
|
||||
image=image.convert('RGB')
|
||||
image=pil2tensor(image)
|
||||
return (image,text,)
|
||||
|
||||
|
||||
class EnhanceImage:
|
||||
@@ -1029,6 +1203,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 里:
|
||||
@@ -1378,7 +1588,7 @@ class Image3D:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Image"
|
||||
CATEGORY = "♾️Mixlab/3D"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,False,False,False,)
|
||||
@@ -1488,6 +1698,37 @@ class FaceToMask:
|
||||
return (mask,)
|
||||
|
||||
|
||||
class CompositeImages:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"foreground": ("IMAGE",),
|
||||
"mask":("MASK",),
|
||||
"background": ("IMAGE",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("IMAGE",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Layer"
|
||||
|
||||
# OUTPUT_IS_LIST = (True,)
|
||||
|
||||
def run(self, foreground,mask,background):
|
||||
foreground= tensor2pil(foreground)
|
||||
mask= tensor2pil(mask)
|
||||
background= tensor2pil(background)
|
||||
res=composite_images(foreground,background,mask)
|
||||
|
||||
return (pil2tensor(res),)
|
||||
|
||||
|
||||
|
||||
|
||||
class EmptyLayer:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -1583,7 +1824,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}),
|
||||
@@ -1640,21 +1881,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
|
||||
@@ -1804,20 +2044,198 @@ class CenterImage:
|
||||
|
||||
return (grid,pil2tensor(mask),)
|
||||
|
||||
class GridDisplayAndSave:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"labels": ("STRING",
|
||||
{
|
||||
"multiline": True,
|
||||
"default": "",
|
||||
"forceInput": True,
|
||||
"dynamicPrompts": False
|
||||
}),
|
||||
"grids": ("_GRID",),
|
||||
|
||||
"image": ("IMAGE",),
|
||||
"filename_prefix": ("STRING", {"default": "mixlab/grids"})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ( )
|
||||
RETURN_NAMES = ( )
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Layer"
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_NODE = True
|
||||
# OUTPUT_IS_LIST = (True,)
|
||||
|
||||
def run(self,labels,grids,image,filename_prefix):
|
||||
|
||||
# print(image.shape)
|
||||
|
||||
img= tensor2pil(image[0])
|
||||
|
||||
for grid in grids:
|
||||
draw_rectangle(img, grid, 'red',8)
|
||||
|
||||
#获取临时目录:temp
|
||||
output_dir = folder_paths.get_temp_directory()
|
||||
|
||||
(
|
||||
full_output_folder,
|
||||
filename,
|
||||
counter,
|
||||
subfolder,
|
||||
_,
|
||||
) = folder_paths.get_save_image_path('tmp_', output_dir)
|
||||
|
||||
image_file = f"{filename}_{counter:05}.png"
|
||||
|
||||
image_path=os.path.join(full_output_folder, image_file)
|
||||
# 保存图片
|
||||
img.save(image_path,compress_level=6)
|
||||
width, height = img.size
|
||||
|
||||
(
|
||||
full_output_folder,
|
||||
filename,
|
||||
counter,
|
||||
_,
|
||||
_,
|
||||
) = folder_paths.get_save_image_path(filename_prefix[0], output_dir)
|
||||
|
||||
|
||||
data_converted = [{
|
||||
"label":labels[i],
|
||||
"grid":[float(grids[i][0]),
|
||||
float(grids[i][1]),
|
||||
float(grids[i][2]),
|
||||
float(grids[i][3])
|
||||
]
|
||||
} for i in range(len(grids))]
|
||||
|
||||
data={
|
||||
"width":int(width),
|
||||
"height":int(height),
|
||||
"grids":data_converted
|
||||
}
|
||||
|
||||
save_json_to_file(data,os.path.join(full_output_folder,f"${filename}_{counter:05}.json"))
|
||||
|
||||
return {"ui":{"image": [{
|
||||
"filename": image_file,
|
||||
"subfolder": subfolder,
|
||||
"type":"temp"
|
||||
}],
|
||||
"json":[data["width"],data['height'],data["grids"]]
|
||||
},"result": ()}
|
||||
# return {"ui":{"image": [ ],
|
||||
|
||||
# },"result": ()}
|
||||
|
||||
class GridInput:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"grids": ("STRING",
|
||||
{
|
||||
"multiline": True,
|
||||
"default": "",
|
||||
"dynamicPrompts": False
|
||||
}),
|
||||
"padding":("INT",{
|
||||
"default": 24,
|
||||
"min": -500, #Minimum value
|
||||
"max": 5000, #Maximum value
|
||||
"step": 1, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
|
||||
},
|
||||
"optional":{
|
||||
"width":("INT",{
|
||||
"forceInput": True,
|
||||
}),
|
||||
"height":("INT",{
|
||||
"forceInput": True,
|
||||
}),
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("_GRID","STRING","IMAGE",)
|
||||
RETURN_NAMES = ("grids","labels","image",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,True,False,)
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def run(self,grids,padding,width=[-1],height=[-1]):
|
||||
# print(padding[0],grids[0])
|
||||
width=width[0]
|
||||
height=height[0]
|
||||
|
||||
grids=grids[0]
|
||||
data=json.loads(grids)
|
||||
grids=data['grids']
|
||||
|
||||
if width>-1:
|
||||
data['width']=width
|
||||
if height>-1:
|
||||
data['height']=height
|
||||
|
||||
new_grids=[]
|
||||
labels=[]
|
||||
|
||||
for g in grids:
|
||||
labels.append(g['label'])
|
||||
new_grids.append(padding_rectangle(g['grid'],padding[0]))
|
||||
|
||||
image = Image.new("RGB", (int(data['width']),int(data["height"])), "white")
|
||||
im=pil2tensor(image)
|
||||
# image=create_temp_file(im)
|
||||
|
||||
data_converted = [{
|
||||
"label":labels[i],
|
||||
"grid":[float(new_grids[i][0]),
|
||||
float(new_grids[i][1]),
|
||||
float(new_grids[i][2]),
|
||||
float(new_grids[i][3])
|
||||
]
|
||||
} for i in range(len(new_grids))]
|
||||
|
||||
# 传递到前端节点的数据 报错,需要处理成 key:[x,x,x,x]
|
||||
return {"ui":{
|
||||
"json":[data["width"],data["height"],data_converted]
|
||||
},"result": (new_grids,labels,im,)}
|
||||
|
||||
# return (new_grids,labels,pil2tensor(image),)
|
||||
|
||||
class GridOutput:
|
||||
@classmethod
|
||||
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"
|
||||
|
||||
@@ -1826,9 +2244,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
|
||||
@@ -1920,10 +2358,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",)
|
||||
@@ -1936,11 +2378,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])
|
||||
@@ -1971,6 +2414,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,
|
||||
@@ -1978,7 +2423,8 @@ class MergeLayers:
|
||||
layer['y'],
|
||||
layer['width'],
|
||||
layer['height'],
|
||||
layer['scale_option']
|
||||
layer['scale_option'],
|
||||
is_multiply_blend
|
||||
)
|
||||
|
||||
final_mask=merge_images(final_mask,
|
||||
@@ -2225,11 +2671,15 @@ class ResizeImage:
|
||||
for ims in image:
|
||||
for im in ims:
|
||||
im=tensor2pil(im)
|
||||
im=resize_image(im,scale_option,w,h,fill_color)
|
||||
im=im.convert('RGB')
|
||||
|
||||
im=im.convert('RGB')
|
||||
a_im,hex=get_average_color_image(im)
|
||||
|
||||
if average_color=='on':
|
||||
fill_color=hex
|
||||
|
||||
im=resize_image(im,scale_option,w,h,fill_color)
|
||||
|
||||
im=pil2tensor(im)
|
||||
imgs.append(im)
|
||||
|
||||
@@ -2329,7 +2779,57 @@ class GetImageSize_:
|
||||
|
||||
return (width, height,min_width,min_height,)
|
||||
|
||||
class SaveImageAndMetadata:
|
||||
def __init__(self):
|
||||
self.output_dir = folder_paths.get_output_directory()
|
||||
self.type = "output"
|
||||
self.prefix_append = ""
|
||||
self.compress_level = 4
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"images": ("IMAGE", ),
|
||||
"filename_prefix": ("STRING", {"default": "Mixlab"}),
|
||||
"metadata": (["disable","enable"],),
|
||||
},
|
||||
"hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
FUNCTION = "save_images"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
|
||||
CATEGORY = "♾️Mixlab/Output"
|
||||
|
||||
def save_images(self, images, filename_prefix="Mixlab",metadata="disable", prompt=None, extra_pnginfo=None):
|
||||
filename_prefix += self.prefix_append
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, self.output_dir, images[0].shape[1], images[0].shape[0])
|
||||
results = list()
|
||||
for image in images:
|
||||
i = 255. * image.cpu().numpy()
|
||||
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
|
||||
metadata = None
|
||||
if (not args.disable_metadata) and (metadata=="enable"):
|
||||
print('##enable_metadata')
|
||||
metadata = PngInfo()
|
||||
if prompt is not None:
|
||||
metadata.add_text("prompt", json.dumps(prompt))
|
||||
if extra_pnginfo is not None:
|
||||
for x in extra_pnginfo:
|
||||
metadata.add_text(x, json.dumps(extra_pnginfo[x]))
|
||||
|
||||
file = f"{filename}_{counter:05}_.png"
|
||||
img.save(os.path.join(full_output_folder, file), pnginfo=metadata, compress_level=self.compress_level)
|
||||
results.append({
|
||||
"filename": file,
|
||||
"subfolder": subfolder,
|
||||
"type": self.type
|
||||
})
|
||||
counter += 1
|
||||
|
||||
return { "ui": { "images": results } }
|
||||
|
||||
class ImageColorTransfer:
|
||||
@classmethod
|
||||
@@ -2337,6 +2837,7 @@ class ImageColorTransfer:
|
||||
return {"required": {
|
||||
"source": ("IMAGE",),
|
||||
"target": ("IMAGE",),
|
||||
"weight": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -2347,28 +2848,48 @@ class ImageColorTransfer:
|
||||
FUNCTION = "run"
|
||||
|
||||
# 右键菜单目录
|
||||
CATEGORY = "♾️Mixlab/Image"
|
||||
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,)
|
||||
|
||||
@@ -2395,7 +2916,7 @@ class SaveImageToLocal:
|
||||
|
||||
OUTPUT_NODE = True
|
||||
|
||||
CATEGORY = "♾️Mixlab/Image"
|
||||
CATEGORY = "♾️Mixlab/Output"
|
||||
|
||||
def save_images(self, images,file_path , prompt=None, extra_pnginfo=None):
|
||||
filename_prefix = os.path.basename(file_path)
|
||||
|
||||
+108
-5
@@ -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):
|
||||
@@ -22,6 +20,19 @@ def pil2tensor(image):
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
|
||||
def add_masks(mask1, mask2):
|
||||
mask1 = mask1.cpu()
|
||||
mask2 = mask2.cpu()
|
||||
cv2_mask1 = np.array(mask1) * 255
|
||||
cv2_mask2 = np.array(mask2) * 255
|
||||
|
||||
if cv2_mask1.shape == cv2_mask2.shape:
|
||||
cv2_mask = cv2.add(cv2_mask1, cv2_mask2)
|
||||
return torch.clamp(torch.from_numpy(cv2_mask) / 255.0, min=0, max=1)
|
||||
else:
|
||||
return mask1
|
||||
|
||||
|
||||
def grow(mask, expand, tapered_corners):
|
||||
c = 0 if tapered_corners else 1
|
||||
kernel = np.array([[c, 1, c],
|
||||
@@ -58,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
|
||||
@@ -87,6 +127,69 @@ class OutlineMask:
|
||||
return (m3,)
|
||||
|
||||
|
||||
class MaskListReplace:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"masks": ("MASK",),
|
||||
"mask_replace": ("MASK",),
|
||||
"start_index":("INT", {"default": 0, "min": 0, "step": 1}),
|
||||
"end_index":("INT", {"default": 0, "min": 0, "step": 1}),
|
||||
"invert": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MASK",)
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
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]
|
||||
invert=invert[0]
|
||||
|
||||
new_masks=[]
|
||||
for i in range(len(masks)):
|
||||
if i>=start_index and i<=end_index:
|
||||
if invert:
|
||||
new_masks.append(masks[i])
|
||||
else:
|
||||
new_masks.append(mask_replace)
|
||||
else:
|
||||
if invert:
|
||||
new_masks.append(mask_replace)
|
||||
else:
|
||||
new_masks.append(masks[i])
|
||||
|
||||
return (new_masks,)
|
||||
|
||||
|
||||
class MaskListMerge:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"masks": ("MASK",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MASK",)
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/Mask"
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
def run(self, masks):
|
||||
mask=masks[0]
|
||||
if isinstance(masks, list):
|
||||
for m in masks:
|
||||
# print(m.shape)
|
||||
mask = add_masks(mask, m)
|
||||
return (mask,)
|
||||
|
||||
|
||||
class FeatheredMask:
|
||||
|
||||
+109
-4
@@ -189,7 +189,7 @@ class PromptImage:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
CATEGORY = "♾️Mixlab/Output"
|
||||
|
||||
# 运行的函数
|
||||
def run(self,prompts,images,save_to_image):
|
||||
@@ -514,7 +514,7 @@ class RandomPrompt:
|
||||
# return (lora_name,prompt,output_tags.split(','),)
|
||||
|
||||
|
||||
import folder_paths
|
||||
|
||||
class EmbeddingPrompt:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -546,7 +546,112 @@ class EmbeddingPrompt:
|
||||
# return (new_prompt)
|
||||
return (prompt,)
|
||||
|
||||
RETURN_TYPES = (any_type,)
|
||||
# RETURN_TYPES = (any_type,)
|
||||
|
||||
# conditioning :提示,正向or负向
|
||||
# clip:clip模型
|
||||
# gligen_textbox_model:gligen模型
|
||||
# grids:矩形框的集合
|
||||
# labels:每个矩形框对应的标签的集合
|
||||
# index:选取第几个矩形框作为gligen的box
|
||||
|
||||
class GLIGENTextBoxApply_Advanced:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {"conditioning": ("CONDITIONING", ),
|
||||
"clip": ("CLIP", ),
|
||||
"gligen_textbox_model": ("GLIGEN", ),
|
||||
"grids": ("_GRID",),
|
||||
"labels": ("STRING",
|
||||
{
|
||||
"multiline": True,
|
||||
"default": "",
|
||||
"forceInput": True
|
||||
}),
|
||||
"index": ("INT", {"default": -1, "min": -1, "max": 300, "step": 1}),
|
||||
"max_size": ("INT", {"default": 8, "min": 1, "max": 300, "step": 1}),
|
||||
"random_shuffle":(["on","off"],),
|
||||
},
|
||||
"optional":{
|
||||
"seed": (any_type, {"default": 0, "min": 0, "max": 0xffffffffffffffff,"step": 1}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("CONDITIONING","STRING",)
|
||||
RETURN_NAMES = ("CONDITIONING","label",)
|
||||
|
||||
FUNCTION = "run"
|
||||
# INPUT_IS_LIST = True
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
|
||||
def run(self, conditioning, clip, gligen_textbox_model, grids, labels, index,max_size,random_shuffle,seed=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
|
||||
|
||||
if index>-1:
|
||||
texts=[labels[index]]
|
||||
grids=[grids[index]]
|
||||
|
||||
if random_shuffle=='on':
|
||||
sss=[[texts[i],grids[i]] for i in range(len(texts))]
|
||||
random.shuffle(sss)
|
||||
texts=[s[0] for s in sss]
|
||||
grids=[s[1] for s in sss]
|
||||
|
||||
if len(texts) > max_size:
|
||||
texts = texts[:max_size]
|
||||
|
||||
c = []
|
||||
|
||||
for t in conditioning:
|
||||
n = [t[0], t[1].copy()]
|
||||
|
||||
|
||||
# 多个
|
||||
position_params=[]
|
||||
for i in range(len(texts)):
|
||||
text=texts[i]
|
||||
grid=grids[i]
|
||||
x,y,width,height=grid
|
||||
# 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)]
|
||||
|
||||
# 前一个
|
||||
prev = []
|
||||
if "gligen" in n[1]:
|
||||
prev = n[1]['gligen'][2]
|
||||
|
||||
n[1]['gligen'] = ("position", gligen_textbox_model, prev + position_params)
|
||||
# print('gligen',n)
|
||||
c.append(n)
|
||||
|
||||
# 下面这个写法有bug
|
||||
# for i in range(len(texts)):
|
||||
# text=texts[i]
|
||||
# grid=grids[i]
|
||||
# x,y,width,height=grid
|
||||
|
||||
# cond, cond_pooled = clip.encode_from_tokens(clip.tokenize(text), return_pooled=True)
|
||||
# for t in conditioning:
|
||||
# n = [t[0], t[1].copy()]
|
||||
# position_params = [(cond_pooled, height // 8, width // 8, y // 8, x // 8)]
|
||||
# prev = []
|
||||
# if "gligen" in n[1]:
|
||||
# prev = n[1]['gligen'][2]
|
||||
|
||||
# n[1]['gligen'] = ("position", gligen_textbox_model, prev + position_params)
|
||||
# c.append(n)
|
||||
|
||||
|
||||
return (c,texts, )
|
||||
|
||||
|
||||
class JoinWithDelimiter:
|
||||
@classmethod
|
||||
@@ -561,7 +666,7 @@ class JoinWithDelimiter:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
CATEGORY = "♾️Mixlab/Text"
|
||||
|
||||
INPUT_IS_LIST = True # 当true的时候,输入时list,当false的时候,如果输入是list,则会自动包一层for循环调用
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
@@ -93,7 +93,7 @@ class ScreenShareNode:
|
||||
RETURN_NAMES = ("IMAGE","PROMPT","FLOAT","INT")
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Image"
|
||||
CATEGORY = "♾️Mixlab/Screen"
|
||||
|
||||
# INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (False,False,False,False)
|
||||
@@ -118,7 +118,7 @@ class FloatingVideo:
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Image"
|
||||
CATEGORY = "♾️Mixlab/Screen"
|
||||
|
||||
# INPUT_IS_LIST = True
|
||||
# OUTPUT_IS_LIST = (False,False,)
|
||||
|
||||
+427
@@ -1,6 +1,18 @@
|
||||
import comfy
|
||||
import torch
|
||||
|
||||
from dataclasses import dataclass
|
||||
import torch.nn as nn
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
import comfy.ops
|
||||
from typing import Union
|
||||
import comfy.sample
|
||||
import latent_preview
|
||||
import comfy.utils
|
||||
|
||||
T = torch.Tensor
|
||||
|
||||
|
||||
from .VisualStylePrompting.attention_functions import VisualStyleProcessor
|
||||
|
||||
class ApplyVisualStylePrompting:
|
||||
@@ -74,3 +86,418 @@ class ApplyVisualStylePrompting:
|
||||
|
||||
return (model, conditioning_prompt, negative_prompt, {"samples": latents, "noise_mask": denoise_mask})
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def exists(val):
|
||||
return val is not None
|
||||
|
||||
def default(val, d):
|
||||
if exists(val):
|
||||
return val
|
||||
return d
|
||||
|
||||
|
||||
class StyleAlignedArgs:
|
||||
def __init__(self, share_attn: str) -> None:
|
||||
self.adain_keys = "k" in share_attn
|
||||
self.adain_values = "v" in share_attn
|
||||
self.adain_queries = "q" in share_attn
|
||||
|
||||
share_attention: bool = True
|
||||
adain_queries: bool = True
|
||||
adain_keys: bool = True
|
||||
adain_values: bool = True
|
||||
|
||||
|
||||
def expand_first(
|
||||
feat: T,
|
||||
scale=1.0,
|
||||
) -> T:
|
||||
"""
|
||||
Expand the first element so it has the same shape as the rest of the batch.
|
||||
"""
|
||||
b = feat.shape[0]
|
||||
feat_style = torch.stack((feat[0], feat[b // 2])).unsqueeze(1)
|
||||
if scale == 1:
|
||||
feat_style = feat_style.expand(2, b // 2, *feat.shape[1:])
|
||||
else:
|
||||
feat_style = feat_style.repeat(1, b // 2, 1, 1, 1)
|
||||
feat_style = torch.cat([feat_style[:, :1], scale * feat_style[:, 1:]], dim=1)
|
||||
return feat_style.reshape(*feat.shape)
|
||||
|
||||
|
||||
def concat_first(feat: T, dim=2, scale=1.0) -> T:
|
||||
"""
|
||||
concat the the feature and the style feature expanded above
|
||||
"""
|
||||
feat_style = expand_first(feat, scale=scale)
|
||||
return torch.cat((feat, feat_style), dim=dim)
|
||||
|
||||
|
||||
def calc_mean_std(feat, eps: float = 1e-5) -> "tuple[T, T]":
|
||||
feat_std = (feat.var(dim=-2, keepdims=True) + eps).sqrt()
|
||||
feat_mean = feat.mean(dim=-2, keepdims=True)
|
||||
return feat_mean, feat_std
|
||||
|
||||
def adain(feat: T) -> T:
|
||||
feat_mean, feat_std = calc_mean_std(feat)
|
||||
feat_style_mean = expand_first(feat_mean)
|
||||
feat_style_std = expand_first(feat_std)
|
||||
feat = (feat - feat_mean) / feat_std
|
||||
feat = feat * feat_style_std + feat_style_mean
|
||||
return feat
|
||||
|
||||
class SharedAttentionProcessor:
|
||||
def __init__(self, args: StyleAlignedArgs, scale: float):
|
||||
self.args = args
|
||||
self.scale = scale
|
||||
|
||||
def __call__(self, q, k, v, extra_options):
|
||||
if self.args.adain_queries:
|
||||
q = adain(q)
|
||||
if self.args.adain_keys:
|
||||
k = adain(k)
|
||||
if self.args.adain_values:
|
||||
v = adain(v)
|
||||
if self.args.share_attention:
|
||||
k = concat_first(k, -2, scale=self.scale)
|
||||
v = concat_first(v, -2)
|
||||
|
||||
return q, k, v
|
||||
|
||||
|
||||
def get_norm_layers(
|
||||
layer: nn.Module,
|
||||
norm_layers_: "dict[str, list[Union[nn.GroupNorm, nn.LayerNorm]]]",
|
||||
share_layer_norm: bool,
|
||||
share_group_norm: bool,
|
||||
):
|
||||
if isinstance(layer, nn.LayerNorm) and share_layer_norm:
|
||||
norm_layers_["layer"].append(layer)
|
||||
if isinstance(layer, nn.GroupNorm) and share_group_norm:
|
||||
norm_layers_["group"].append(layer)
|
||||
else:
|
||||
for child_layer in layer.children():
|
||||
get_norm_layers(
|
||||
child_layer, norm_layers_, share_layer_norm, share_group_norm
|
||||
)
|
||||
|
||||
|
||||
def register_norm_forward(
|
||||
norm_layer: Union[nn.GroupNorm, nn.LayerNorm],
|
||||
) -> Union[nn.GroupNorm, nn.LayerNorm]:
|
||||
if not hasattr(norm_layer, "orig_forward"):
|
||||
setattr(norm_layer, "orig_forward", norm_layer.forward)
|
||||
orig_forward = norm_layer.orig_forward
|
||||
|
||||
def forward_(hidden_states: T) -> T:
|
||||
n = hidden_states.shape[-2]
|
||||
hidden_states = concat_first(hidden_states, dim=-2)
|
||||
hidden_states = orig_forward(hidden_states) # type: ignore
|
||||
return hidden_states[..., :n, :]
|
||||
|
||||
norm_layer.forward = forward_ # type: ignore
|
||||
return norm_layer
|
||||
|
||||
|
||||
def register_shared_norm(
|
||||
model: ModelPatcher,
|
||||
share_group_norm: bool = True,
|
||||
share_layer_norm: bool = True,
|
||||
):
|
||||
norm_layers = {"group": [], "layer": []}
|
||||
get_norm_layers(model.model, norm_layers, share_layer_norm, share_group_norm)
|
||||
print(
|
||||
f"Patching {len(norm_layers['group'])} group norms, {len(norm_layers['layer'])} layer norms."
|
||||
)
|
||||
return [register_norm_forward(layer) for layer in norm_layers["group"]] + [
|
||||
register_norm_forward(layer) for layer in norm_layers["layer"]
|
||||
]
|
||||
|
||||
|
||||
SHARE_NORM_OPTIONS = ["both", "group", "layer", "disabled"]
|
||||
SHARE_ATTN_OPTIONS = ["q+k", "q+k+v", "disabled"]
|
||||
|
||||
class StyleAlignedSampleReferenceLatents:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{
|
||||
"reference_image": ("IMAGE",),
|
||||
"positive": ("CONDITIONING",),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"model": ("MODEL",),
|
||||
"vae": ("VAE", ),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
|
||||
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
|
||||
|
||||
"scheduler": (comfy.samplers.KSampler.SCHEDULERS.reverse(), ),
|
||||
"denoise": ("FLOAT", {"default": 1, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STEP_LATENTS","LATENT")
|
||||
RETURN_NAMES = ("ref_latents", "noised_output")
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
# CATEGORY = "style_aligned"
|
||||
CATEGORY = "♾️Mixlab/Style"
|
||||
|
||||
def run(self, reference_image, positive, negative, model, vae, seed, steps, cfg,scheduler,denoise):
|
||||
|
||||
# TODO noise_mask?
|
||||
def vae_encode_crop_pixels(pixels):
|
||||
x = (pixels.shape[1] // 8) * 8
|
||||
y = (pixels.shape[2] // 8) * 8
|
||||
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, :]
|
||||
return pixels
|
||||
|
||||
pixels=vae_encode_crop_pixels(reference_image)
|
||||
t = vae.encode(pixels[:,:,:,:3])
|
||||
latent_image = {"samples":t}
|
||||
|
||||
noise_seed=seed
|
||||
|
||||
sampler_name="ddim"
|
||||
|
||||
sampler = comfy.samplers.sampler_object(sampler_name)
|
||||
|
||||
total_steps = steps
|
||||
if denoise < 1.0:
|
||||
total_steps = int(steps/denoise)
|
||||
|
||||
comfy.model_management.load_models_gpu([model])
|
||||
sigmas = comfy.samplers.calculate_sigmas_scheduler(model.model, scheduler, total_steps).cpu()
|
||||
sigmas = sigmas[-(steps + 1):]
|
||||
|
||||
sigmas = sigmas.flip(0)
|
||||
if sigmas[0] == 0:
|
||||
sigmas[0] = 0.0001
|
||||
|
||||
latent = latent_image
|
||||
latent_image = latent["samples"]
|
||||
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
|
||||
|
||||
|
||||
noise_mask = None
|
||||
if "noise_mask" in latent:
|
||||
noise_mask = latent["noise_mask"]
|
||||
|
||||
ref_latents = []
|
||||
def callback(step: int, x0: T, x: T, steps: int):
|
||||
ref_latents.insert(0, x[0])
|
||||
|
||||
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED
|
||||
samples = comfy.sample.sample_custom(model, noise, cfg, sampler, sigmas, positive, negative, latent_image, noise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=noise_seed)
|
||||
|
||||
out = latent.copy()
|
||||
out["samples"] = samples
|
||||
out_noised = out
|
||||
|
||||
ref_latents = torch.stack(ref_latents)
|
||||
|
||||
return (ref_latents, out_noised)
|
||||
|
||||
class StyleAlignedReferenceSampler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
|
||||
"ref_latents": ("STEP_LATENTS",),
|
||||
"reference_image_text": ("STRING", {"multiline": True}),
|
||||
"model": ("MODEL",),
|
||||
"clip": ("CLIP", ),
|
||||
|
||||
"positive": ("CONDITIONING",),
|
||||
"negative": ("CONDITIONING",),
|
||||
|
||||
"share_norm": (SHARE_NORM_OPTIONS,),
|
||||
"share_attn": (SHARE_ATTN_OPTIONS,),
|
||||
"scale": ("FLOAT", {"default": 1, "min": 0, "max": 2.0, "step": 0.01}),
|
||||
"batch_size": ("INT", {"default": 2, "min": 1, "max": 8, "step": 1}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
|
||||
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
|
||||
"scheduler": (comfy.samplers.KSampler.SCHEDULERS, ),
|
||||
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT", "LATENT")
|
||||
RETURN_NAMES = ("output", "denoised_output")
|
||||
FUNCTION = "patch"
|
||||
# CATEGORY = "style_aligned"
|
||||
CATEGORY = "♾️Mixlab/Style"
|
||||
def patch(
|
||||
self,
|
||||
ref_latents,
|
||||
reference_image_text,
|
||||
model,
|
||||
clip,
|
||||
positive,
|
||||
negative,
|
||||
share_norm,
|
||||
share_attn,
|
||||
scale,
|
||||
batch_size,
|
||||
seed,steps,cfg,scheduler,denoise
|
||||
|
||||
) -> "tuple[dict, dict]":
|
||||
|
||||
m = model.clone()
|
||||
|
||||
# ref_latents = vae.encode(reference_image[:,:,:,:3])
|
||||
|
||||
tokens = clip.tokenize(reference_image_text)
|
||||
cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True)
|
||||
ref_positive=[[cond, {"pooled_output": pooled}]]
|
||||
|
||||
noise_seed=seed
|
||||
|
||||
|
||||
total_steps = steps
|
||||
if denoise < 1.0:
|
||||
total_steps = int(steps/denoise)
|
||||
|
||||
# comfy.model_management.load_models_gpu([model])
|
||||
sigmas = comfy.samplers.calculate_sigmas_scheduler(model.model, scheduler, total_steps).cpu()
|
||||
sigmas = sigmas[-(steps + 1):]
|
||||
|
||||
sampler_name="ddim"
|
||||
|
||||
sampler = comfy.samplers.sampler_object(sampler_name)
|
||||
|
||||
args = StyleAlignedArgs(share_attn)
|
||||
|
||||
# Concat batch with style latent
|
||||
style_latent_tensor = ref_latents[0].unsqueeze(0)
|
||||
height, width = style_latent_tensor.shape[-2:]
|
||||
latent_t = torch.zeros(
|
||||
[batch_size, 4, height, width], device=ref_latents.device
|
||||
)
|
||||
latent = {"samples": latent_t}
|
||||
noise = comfy.sample.prepare_noise(latent_t, noise_seed)
|
||||
|
||||
latent_t = torch.cat((style_latent_tensor, latent_t), dim=0)
|
||||
ref_noise = torch.zeros_like(noise[0]).unsqueeze(0)
|
||||
noise = torch.cat((ref_noise, noise), dim=0)
|
||||
|
||||
x0_output = {}
|
||||
preview_callback = latent_preview.prepare_callback(m, sigmas.shape[-1] - 1, x0_output)
|
||||
|
||||
# Replace first latent with the corresponding reference latent after each step
|
||||
def callback(step: int, x0: T, x: T, steps: int):
|
||||
preview_callback(step, x0, x, steps)
|
||||
if (step + 1 < steps):
|
||||
# 当ref_latents的step不够时
|
||||
if step+1>len(ref_latents)-1:
|
||||
step=len(ref_latents)-2
|
||||
|
||||
x[0] = ref_latents[step+1]
|
||||
x0[0] = ref_latents[step+1]
|
||||
|
||||
# Register shared norms
|
||||
share_group_norm = share_norm in ["group", "both"]
|
||||
share_layer_norm = share_norm in ["layer", "both"]
|
||||
register_shared_norm(m, share_group_norm, share_layer_norm)
|
||||
|
||||
# Patch cross attn
|
||||
m.set_model_attn1_patch(SharedAttentionProcessor(args, scale))
|
||||
|
||||
# Add reference conditioning to batch
|
||||
batched_condition = []
|
||||
for i,condition in enumerate(positive):
|
||||
additional = condition[1].copy()
|
||||
batch_with_reference = torch.cat([ref_positive[i][0], condition[0].repeat([batch_size] + [1] * len(condition[0].shape[1:]))], dim=0)
|
||||
if 'pooled_output' in additional and 'pooled_output' in ref_positive[i][1]:
|
||||
# combine pooled output
|
||||
pooled_output = torch.cat([ref_positive[i][1]['pooled_output'], additional['pooled_output'].repeat([batch_size]
|
||||
+ [1] * len(additional['pooled_output'].shape[1:]))], dim=0)
|
||||
additional['pooled_output'] = pooled_output
|
||||
if 'control' in additional:
|
||||
if 'control' in ref_positive[i][1]:
|
||||
# combine control conditioning
|
||||
control_hint = torch.cat([ref_positive[i][1]['control'].cond_hint_original, additional['control'].cond_hint_original.repeat([batch_size]
|
||||
+ [1] * len(additional['control'].cond_hint_original.shape[1:]))], dim=0)
|
||||
cloned_controlnet = additional['control'].copy()
|
||||
cloned_controlnet.set_cond_hint(control_hint, strength=additional['control'].strength, timestep_percent_range=additional['control'].timestep_percent_range)
|
||||
additional['control'] = cloned_controlnet
|
||||
else:
|
||||
# add zeros for first in batch
|
||||
control_hint = torch.cat([torch.zeros_like(additional['control'].cond_hint_original), additional['control'].cond_hint_original.repeat([batch_size]
|
||||
+ [1] * len(additional['control'].cond_hint_original.shape[1:]))], dim=0)
|
||||
cloned_controlnet = additional['control'].copy()
|
||||
cloned_controlnet.set_cond_hint(control_hint, strength=additional['control'].strength, timestep_percent_range=additional['control'].timestep_percent_range)
|
||||
additional['control'] = cloned_controlnet
|
||||
batched_condition.append([batch_with_reference, additional])
|
||||
|
||||
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED
|
||||
samples = comfy.sample.sample_custom(
|
||||
m,
|
||||
noise,
|
||||
cfg,
|
||||
sampler,
|
||||
sigmas,
|
||||
batched_condition,
|
||||
negative,
|
||||
latent_t,
|
||||
callback=callback,
|
||||
disable_pbar=disable_pbar,
|
||||
seed=noise_seed,
|
||||
)
|
||||
|
||||
# remove reference image
|
||||
samples = samples[1:]
|
||||
|
||||
out = latent.copy()
|
||||
out["samples"] = samples
|
||||
if "x0" in x0_output:
|
||||
out_denoised = latent.copy()
|
||||
x0 = x0_output["x0"][1:]
|
||||
out_denoised["samples"] = m.model.process_latent_out(x0.cpu())
|
||||
else:
|
||||
out_denoised = out
|
||||
return (out, out_denoised)
|
||||
|
||||
|
||||
class StyleAlignedBatchAlign:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"share_norm": (SHARE_NORM_OPTIONS,),
|
||||
"share_attn": (SHARE_ATTN_OPTIONS,),
|
||||
"scale": ("FLOAT", {"default": 1, "min": 0, "max": 1.0, "step": 0.1}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "patch"
|
||||
# CATEGORY = "style_aligned"
|
||||
CATEGORY = "♾️Mixlab/Style"
|
||||
def patch(
|
||||
self,
|
||||
model: ModelPatcher,
|
||||
share_norm: str,
|
||||
share_attn: str,
|
||||
scale: float,
|
||||
):
|
||||
m = model.clone()
|
||||
share_group_norm = share_norm in ["group", "both"]
|
||||
share_layer_norm = share_norm in ["layer", "both"]
|
||||
register_shared_norm(model, share_group_norm, share_layer_norm)
|
||||
args = StyleAlignedArgs(share_attn)
|
||||
m.set_model_attn1_patch(SharedAttentionProcessor(args, scale))
|
||||
return (m,)
|
||||
|
||||
|
||||
|
||||
+14
-11
@@ -12,6 +12,8 @@ import comfy.utils
|
||||
# import numpy as np
|
||||
import torch
|
||||
import random
|
||||
from lark import Lark, Transformer, v_args
|
||||
|
||||
|
||||
global _available
|
||||
_available=True
|
||||
@@ -109,7 +111,7 @@ def text_generate(text_pipe,input,seed=None):
|
||||
|
||||
import re
|
||||
|
||||
def correct_prompt_syntax(prompt):
|
||||
def correct_prompt_syntax(prompt=""):
|
||||
|
||||
# print("input prompt",prompt)
|
||||
corrected_elements = []
|
||||
@@ -191,7 +193,7 @@ NUMBER: /\s*-?\d+(\.\d+)?\s*/
|
||||
WORD: /[^,:\(\)\[\]<>]+/
|
||||
"""
|
||||
|
||||
from lark import Lark, Transformer, v_args
|
||||
|
||||
|
||||
@v_args(inline=True) # Decorator to flatten the tree directly into the function arguments
|
||||
class ChinesePromptTranslate(Transformer):
|
||||
@@ -304,7 +306,6 @@ class ChinesePrompt:
|
||||
pbar = comfy.utils.ProgressBar(len(text)+1)
|
||||
texts = [correct_prompt_syntax(t) for t in text]
|
||||
|
||||
|
||||
global text_pipe,zh_en_model,zh_en_tokenizer
|
||||
if zh_en_model==None:
|
||||
zh_en_model = AutoModelForSeq2SeqLM.from_pretrained(zh_en_model_path).eval()
|
||||
@@ -323,13 +324,14 @@ class ChinesePrompt:
|
||||
en_texts=[]
|
||||
|
||||
for t in texts:
|
||||
# translated_text = translated_word = translate(zh_en_tokenizer,zh_en_model,str(t))
|
||||
parser = Lark(grammar, start="start", parser="lalr", transformer=ChinesePromptTranslate())
|
||||
# print('t',t)
|
||||
result = parser.parse(t).children
|
||||
# print('en_result',result)
|
||||
# en_text=translate(zh_en_tokenizer,zh_en_model,text_without_syntax)
|
||||
en_texts.append(result[0])
|
||||
if t:
|
||||
# translated_text = translated_word = translate(zh_en_tokenizer,zh_en_model,str(t))
|
||||
parser = Lark(grammar, start="start", parser="lalr", transformer=ChinesePromptTranslate())
|
||||
# print('t',t)
|
||||
result = parser.parse(t).children
|
||||
# print('en_result',result)
|
||||
# en_text=translate(zh_en_tokenizer,zh_en_model,text_without_syntax)
|
||||
en_texts.append(result[0])
|
||||
|
||||
zh_en_model.to('cpu')
|
||||
print("test en_text",en_texts)
|
||||
@@ -352,7 +354,8 @@ class ChinesePrompt:
|
||||
|
||||
print('prompt_result',prompt_result,)
|
||||
# prompt_result = [','.join(correct_prompt_syntax(p)) for p in prompt_result]
|
||||
|
||||
if len(prompt_result)==0:
|
||||
prompt_result=[""]
|
||||
return {
|
||||
"ui":{
|
||||
"prompt": prompt_result
|
||||
|
||||
@@ -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}}
|
||||
|
||||
|
||||
|
||||
+126
-20
@@ -8,6 +8,18 @@ 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 = []
|
||||
for i in range(0, len(lst), chunk_size):
|
||||
start = i - transition_size
|
||||
end = i + chunk_size + transition_size
|
||||
result.append(lst[max(start, 0):end])
|
||||
return result
|
||||
|
||||
def recursive_search(directory, excluded_dir_names=None):
|
||||
if not os.path.isdir(directory):
|
||||
@@ -146,7 +158,6 @@ class ColorInput:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
|
||||
"color":("TCOLOR",),
|
||||
},
|
||||
}
|
||||
@@ -156,7 +167,7 @@ class ColorInput:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Color"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,False,False,False,False,)
|
||||
@@ -185,7 +196,7 @@ class FontInput:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -218,7 +229,7 @@ class TextToNumber:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Text"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -273,10 +284,10 @@ class FloatSlider:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("FLOAT",)
|
||||
|
||||
RETURN_NAMES = ('weight(0-1)',)
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -329,7 +340,7 @@ class IntNumber:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -389,7 +400,7 @@ class TextInput:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -398,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
|
||||
@@ -577,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'):
|
||||
|
||||
@@ -605,10 +671,43 @@ 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
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"optional":{
|
||||
"A":(any_type,),
|
||||
},
|
||||
"required": {
|
||||
"chunk_size": ("INT", {"default": 10, "min": 1, "step": 1}),
|
||||
"transition_size": ("INT", {"default": 0, "min": 0, "step": 1}),
|
||||
"index": ("INT", {"default": -1, "min": -1, "step": 1}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = (any_type,)
|
||||
RETURN_NAMES = ("B",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
def run(self, A=[],chunk_size=[10],transition_size=[0],index=[-1]):
|
||||
# print(len(A))
|
||||
B=split_list(A,chunk_size[0],transition_size[0])
|
||||
|
||||
if index[0]>-1:
|
||||
B=B[index[0]]
|
||||
|
||||
return (B,)
|
||||
|
||||
|
||||
|
||||
class LimitNumber:
|
||||
@@ -639,7 +738,7 @@ class LimitNumber:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
@@ -706,14 +805,21 @@ class TESTNODE_:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/__TEST"
|
||||
CATEGORY = "♾️Mixlab/Test"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
def run(self,ANY):
|
||||
# print(ANY)
|
||||
print(type(ANY))
|
||||
try:
|
||||
print(ANY[0].shape)
|
||||
img= tensor2pil(ANY[0])
|
||||
print(img.size)
|
||||
except:
|
||||
print('')
|
||||
|
||||
# data=ANY
|
||||
list_stats = ListStatistics()
|
||||
|
||||
@@ -751,7 +857,7 @@ class TESTNODE_TOKEN:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/__TEST"
|
||||
CATEGORY = "♾️Mixlab/Test"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = False
|
||||
@@ -788,7 +894,7 @@ class CreateSeedNode:
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Experiment"
|
||||
|
||||
def run(self, seed):
|
||||
return (seed,)
|
||||
@@ -815,7 +921,7 @@ class CreateCkptNames:
|
||||
# OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Experiment"
|
||||
|
||||
def run(self, ckpt_names):
|
||||
ckpt_names=ckpt_names.split('\n')
|
||||
@@ -844,7 +950,7 @@ class CreateLoraNames:
|
||||
# OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Experiment"
|
||||
|
||||
def run(self, lora_names):
|
||||
lora_names=lora_names.split('\n')
|
||||
@@ -875,7 +981,7 @@ class CreateSampler_names:
|
||||
# OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Utils"
|
||||
CATEGORY = "♾️Mixlab/Experiment"
|
||||
|
||||
def run(self, sampler_names):
|
||||
sampler_names=sampler_names.split('\n')
|
||||
|
||||
+692
@@ -0,0 +1,692 @@
|
||||
import os
|
||||
import hashlib
|
||||
import json
|
||||
import subprocess
|
||||
import shutil
|
||||
import re
|
||||
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,random,string
|
||||
from pathlib import Path
|
||||
|
||||
import folder_paths
|
||||
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"),
|
||||
],
|
||||
[".json"]
|
||||
)
|
||||
|
||||
ffmpeg_path = shutil.which("ffmpeg")
|
||||
if ffmpeg_path is None:
|
||||
print("ffmpeg could not be found. Using ffmpeg from imageio-ffmpeg.")
|
||||
from imageio_ffmpeg import get_ffmpeg_exe
|
||||
try:
|
||||
ffmpeg_path = get_ffmpeg_exe()
|
||||
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 = []
|
||||
for i in range(0, len(lst), chunk_size):
|
||||
start = i - transition_size
|
||||
end = i + chunk_size + transition_size
|
||||
result.append(lst[max(start, 0):end])
|
||||
return result
|
||||
|
||||
# images = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10]
|
||||
# chunk_size = 3
|
||||
# transition_size = 1
|
||||
|
||||
# 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):
|
||||
video_extensions = ['webm', 'mp4', 'mkv', 'gif']
|
||||
input_dir = folder_paths.get_input_directory()
|
||||
files = []
|
||||
for f in os.listdir(input_dir):
|
||||
if os.path.isfile(os.path.join(input_dir, f)):
|
||||
file_parts = f.split('.')
|
||||
if len(file_parts) > 1 and (file_parts[-1] in video_extensions):
|
||||
files.append(f)
|
||||
return {"required": {
|
||||
"video": (sorted(files), {"video_upload": True}),
|
||||
"video_segment_frames": ("INT", {"default": 10, "min": 1, "step": 1}),
|
||||
"transition_frames": ("INT", {"default": 0, "min": 0, "step": 1}),
|
||||
},}
|
||||
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
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,False,False,)
|
||||
|
||||
|
||||
def is_gif(self, filename):
|
||||
file_parts = filename.split('.')
|
||||
return len(file_parts) > 1 and file_parts[-1] == "gif"
|
||||
|
||||
def load_video_cv_fallback(self, video, frame_load_cap, skip_first_frames):
|
||||
try:
|
||||
video_cap = cv2.VideoCapture(folder_paths.get_annotated_filepath(video))
|
||||
if not video_cap.isOpened():
|
||||
raise ValueError(f"{video} could not be loaded with cv fallback.")
|
||||
# set video_cap to look at start_index frame
|
||||
images = []
|
||||
total_frame_count = 0
|
||||
frames_added = 0
|
||||
base_frame_time = 1/video_cap.get(cv2.CAP_PROP_FPS)
|
||||
|
||||
target_frame_time = base_frame_time
|
||||
|
||||
time_offset=0.0
|
||||
while video_cap.isOpened():
|
||||
if time_offset < target_frame_time:
|
||||
is_returned, frame = video_cap.read()
|
||||
# if didn't return frame, video has ended
|
||||
if not is_returned:
|
||||
break
|
||||
time_offset += base_frame_time
|
||||
if time_offset < target_frame_time:
|
||||
continue
|
||||
time_offset -= target_frame_time
|
||||
# if not at start_index, skip doing anything with frame
|
||||
total_frame_count += 1
|
||||
if total_frame_count <= skip_first_frames:
|
||||
continue
|
||||
# TODO: do whatever operations need to happen, like force_size, etc
|
||||
|
||||
# opencv loads images in BGR format (yuck), so need to convert to RGB for ComfyUI use
|
||||
# follow up: can videos ever have an alpha channel?
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
# convert frame to comfyui's expected format (taken from comfy's load image code)
|
||||
image = Image.fromarray(frame)
|
||||
image = ImageOps.exif_transpose(image)
|
||||
image = np.array(image, dtype=np.float32) / 255.0
|
||||
image = torch.from_numpy(image)[None,]
|
||||
images.append(image)
|
||||
frames_added += 1
|
||||
# if cap exists and we've reached it, stop processing frames
|
||||
if frame_load_cap > 0 and frames_added >= frame_load_cap:
|
||||
break
|
||||
finally:
|
||||
video_cap.release()
|
||||
images = torch.cat(images, dim=0)
|
||||
|
||||
return (images, frames_added)
|
||||
|
||||
def load_video(self, video,video_segment_frames,transition_frames ):
|
||||
|
||||
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)
|
||||
# otherwise, continue with ffmpeg
|
||||
|
||||
# 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_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=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 (scenes_video,len(scenes_video), total_frames,fps,)
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(s, video, **kwargs):
|
||||
image_path = folder_paths.get_annotated_filepath(video)
|
||||
m = hashlib.sha256()
|
||||
with open(image_path, 'rb') as f:
|
||||
m.update(f.read())
|
||||
return m.digest().hex()
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(s, video, **kwargs):
|
||||
if not folder_paths.exists_annotated_filepath(video):
|
||||
return "Invalid image file: {}".format(video)
|
||||
|
||||
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, )
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
+9
-2
@@ -7,5 +7,12 @@ openai
|
||||
simple-lama-inpainting
|
||||
clip-interrogator==0.6.0
|
||||
transformers>=4.36.0
|
||||
zhipuai
|
||||
lark-parser
|
||||
lark-parser
|
||||
imageio-ffmpeg
|
||||
rembg[gpu]
|
||||
omegaconf==2.3.0
|
||||
Pillow>=9.5.0
|
||||
einops==0.7.0
|
||||
trimesh>=4.0.5
|
||||
huggingface-hub
|
||||
scikit-image
|
||||
+445
-66
@@ -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,53 @@
|
||||
.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;
|
||||
}
|
||||
</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>
|
||||
|
||||
@@ -413,7 +457,7 @@
|
||||
|
||||
<div id="editor_container"></div>
|
||||
<div class="header">
|
||||
<div style="margin: 0 24px;
|
||||
<div id="logo" style="margin: 0 24px;
|
||||
margin-bottom: 24px;
|
||||
padding: 8px;
|
||||
color: #4a4a4a;
|
||||
@@ -439,8 +483,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'
|
||||
@@ -479,6 +554,59 @@
|
||||
};
|
||||
|
||||
|
||||
const parseImageToBase64 = url => {
|
||||
return new Promise((res, rej) => {
|
||||
fetch(url)
|
||||
.then(response => response.blob())
|
||||
.then(blob => {
|
||||
const reader = new FileReader()
|
||||
reader.onloadend = () => {
|
||||
const base64data = reader.result
|
||||
res(base64data)
|
||||
// 在这里可以将base64数据用于进一步处理或显示图片
|
||||
}
|
||||
reader.readAsDataURL(blob)
|
||||
})
|
||||
.catch(error => {
|
||||
console.log('发生错误:', error)
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
//给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编码中的前缀
|
||||
const base64WithoutPrefix = base64.replace(/^data:image\/\w+;base64,/, '');
|
||||
@@ -677,18 +805,23 @@
|
||||
|
||||
// 种子的处理
|
||||
function randomSeed(seed, data) {
|
||||
let max_seed = 4294967295
|
||||
//1849378600828930
|
||||
for (const id in data) {
|
||||
if (data[id].inputs.seed != undefined
|
||||
&& !Array.isArray(data[id].inputs.seed) //如果是数组,则由其他节点控制
|
||||
&& ['increment', 'decrement', 'randomize'].includes(seed[id])) {
|
||||
data[id].inputs.seed = Math.round(Math.random() * 1849378600828930)
|
||||
data[id].inputs.seed = Math.round(Math.random() * max_seed)
|
||||
// console.log('new Seed', data[id])
|
||||
}
|
||||
if (data[id].inputs.noise_seed != undefined
|
||||
&& !Array.isArray(data[id].inputs.noise_seed) //如果是数组,则由其他节点控制
|
||||
&& ['increment', 'decrement', 'randomize'].includes(seed[id])) {
|
||||
data[id].inputs.noise_seed = Math.round(Math.random() * 1849378600828930)
|
||||
|
||||
data[id].inputs.noise_seed = Math.round(Math.random() * max_seed)
|
||||
}
|
||||
// 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])
|
||||
}
|
||||
@@ -712,6 +845,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`, {
|
||||
@@ -747,6 +895,10 @@
|
||||
let url = get_url()
|
||||
const res = await fetch(`${url}/mixlab/workflow`, {
|
||||
method: 'POST',
|
||||
mode: 'cors', // 允许跨域请求
|
||||
headers: {
|
||||
'Content-Type': 'application/json'
|
||||
},
|
||||
body: JSON.stringify({
|
||||
task: 'my_app',
|
||||
filename,
|
||||
@@ -811,11 +963,6 @@
|
||||
action.appendChild(copyImage)
|
||||
copyImage.style.marginLeft = '18px';
|
||||
|
||||
const copyText = document.createElement('button');
|
||||
copyText.innerText = 'copy text for share'
|
||||
action.appendChild(copyText)
|
||||
copyText.style.marginLeft = '18px';
|
||||
|
||||
let isURL = false;
|
||||
try {
|
||||
new URL(link)
|
||||
@@ -865,12 +1012,16 @@
|
||||
output_card.className = 'output_card'
|
||||
container.appendChild(output_card)
|
||||
|
||||
|
||||
copyText.addEventListener('click', e => {
|
||||
e.preventDefault();
|
||||
copyTextToClipboard((window._appData.share_prefix || '') + " " + output_card.outerHTML, (r) => success(r, copyText, 'copy text for share'))
|
||||
|
||||
})
|
||||
if (window._appData.share_prefix) {
|
||||
const copyText = document.createElement('button');
|
||||
copyText.innerText = 'copy text for share'
|
||||
action.appendChild(copyText)
|
||||
copyText.style.marginLeft = '18px';
|
||||
copyText.addEventListener('click', e => {
|
||||
e.preventDefault();
|
||||
copyTextToClipboard((window._appData.share_prefix || '') + " " + output_card.outerHTML, (r) => success(r, copyText, 'copy text for share'))
|
||||
})
|
||||
}
|
||||
|
||||
copyImage.addEventListener('click', e => {
|
||||
e.preventDefault();
|
||||
@@ -884,6 +1035,9 @@
|
||||
// copyImagesToClipboard(output_card.outerHTML)
|
||||
})
|
||||
|
||||
//是否显示复制图片,复制html两个按钮
|
||||
let isShowImageFn = false;
|
||||
|
||||
for (const node of outputData) {
|
||||
// console.log('output', node)
|
||||
if (node.class_type == "ShowTextForGPT") {
|
||||
@@ -902,8 +1056,31 @@
|
||||
output_card.appendChild(div);
|
||||
};
|
||||
|
||||
if (["SaveImage", "PreviewImage", "PromptImage"].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");
|
||||
@@ -919,7 +1096,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}`
|
||||
@@ -946,6 +1123,12 @@
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
if (isShowImageFn === false) {
|
||||
copyImage.remove();
|
||||
copyHTML.remove();
|
||||
}
|
||||
|
||||
return container
|
||||
}
|
||||
|
||||
@@ -1002,6 +1185,7 @@
|
||||
}
|
||||
|
||||
async function handleClipboardImage(imageElement, data) {
|
||||
//data.class_type === 'LoadImagesToBatch'
|
||||
try {
|
||||
const clipboardItems = await navigator.clipboard.read();
|
||||
for (const clipboardItem of clipboardItems) {
|
||||
@@ -1014,18 +1198,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);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1186,13 +1369,19 @@
|
||||
inputData = inputData.filter(inp => inp);
|
||||
// console.log('inputData',inputData)
|
||||
inputData.forEach(async data => {
|
||||
console.log(data)
|
||||
console.log('inputData', data);
|
||||
|
||||
// 图片 or 视频输入
|
||||
if (data.class_type === "LoadImage" || data.class_type === "VHS_LoadVideo") {
|
||||
|
||||
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';
|
||||
@@ -1224,7 +1413,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) actionDiv.appendChild(btnForImageEdit);
|
||||
if ((!isVideoUpload && !isBase64Upload) && data.class_type !== 'ImagesPrompt_') actionDiv.appendChild(btnForImageEdit);
|
||||
|
||||
|
||||
uploadContainer.appendChild(actionDiv)
|
||||
@@ -1246,7 +1435,7 @@
|
||||
|
||||
// imageElement.innerHTML=`<img src="${base64Df}"/>`
|
||||
|
||||
} else {
|
||||
} else if (data.class_type === 'LoadImage') {
|
||||
// 图片
|
||||
let [subfolder, name] = data.inputs.image.split('/');
|
||||
if (!name) {
|
||||
@@ -1259,15 +1448,36 @@
|
||||
imageElement.src = data.options?.defaultImage || url;
|
||||
imageElement.setAttribute('onerror', `this.src='${base64Df}'`)
|
||||
|
||||
} else if (data.class_type === 'ImagesPrompt_') {
|
||||
// 图片库模式
|
||||
const [imgDiv, mainImage] = createSelectForImages(
|
||||
data.title,
|
||||
data.options.images,
|
||||
data.inputs.imageIndex,
|
||||
(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)
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
|
||||
imageElement.style.maxWidth = '200px';
|
||||
|
||||
|
||||
|
||||
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) => {
|
||||
@@ -1290,22 +1500,43 @@
|
||||
|
||||
if (hashId == window._appData.data[data.id].hashId) return
|
||||
|
||||
let { url, name } = await uploadImage(fileBlob, '.' + file.type.split('/')[1])
|
||||
if (isVideoUpload) {
|
||||
imageElement.srcObject = null;
|
||||
}
|
||||
// 在这里可以对 Blob 对象进行进一步处理
|
||||
imageElement.src = url;
|
||||
|
||||
if (isVideoUpload) {
|
||||
window._appData.data[data.id].inputs.video = name;
|
||||
if (data.class_type === 'LoadImagesToBatch') {
|
||||
// 上传 ,转为base64
|
||||
let base64 = await blobToBase64(fileBlob)
|
||||
createBase64ImageForLoadImageToBatch(imageElement, data.id, base64)
|
||||
} else {
|
||||
window._appData.data[data.id].inputs.image = name;
|
||||
//上传,返回url
|
||||
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;
|
||||
} else {
|
||||
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);
|
||||
|
||||
};
|
||||
|
||||
// 开始读取文件
|
||||
@@ -1315,7 +1546,7 @@
|
||||
})
|
||||
|
||||
// imageElement.src = `${get_url()}/view?filename=${encodeURIComponent(data.inputs.image)}&type=${type}&subfolder=${subfolder}`;
|
||||
uploadContainer.appendChild(imageElement);
|
||||
if (data.class_type !== 'ImagesPrompt_') uploadContainer.appendChild(imageElement);
|
||||
|
||||
// Append the upload container to the main container
|
||||
container.appendChild(uploadContainer);
|
||||
@@ -1392,7 +1623,7 @@
|
||||
textInput.value = data.inputs.prompt;
|
||||
} else {
|
||||
textInput.value = data.inputs.text;
|
||||
}
|
||||
};
|
||||
|
||||
// uploadImageInput.type = "text";
|
||||
let json = localStorage.getItem(`t_${data.id}`)
|
||||
@@ -1408,7 +1639,6 @@
|
||||
window._appData.data[data.id].inputs.text = textInput.value;
|
||||
}
|
||||
|
||||
|
||||
} catch (error) {
|
||||
|
||||
}
|
||||
@@ -1416,6 +1646,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';
|
||||
@@ -1729,6 +1980,73 @@
|
||||
|
||||
}
|
||||
|
||||
function createImageForSelect(isMain, imgurl, keyword) {
|
||||
let im = new Image();
|
||||
im.className = isMain ? 'images_prompt_main' : ''
|
||||
im.src = imgurl;
|
||||
im.title = keyword
|
||||
im.style = `
|
||||
width:${isMain ? 120 : 56}px;
|
||||
height:auto;
|
||||
min-height:${isMain ? 120 : 56}px;
|
||||
filter: brightness(${isMain ? 1 : 0.8});
|
||||
${isMain ? 'filter: drop-shadow(1px 1px 4px black);' : ''}
|
||||
`
|
||||
let p = document.createElement('p');
|
||||
p.innerText = keyword;
|
||||
p.style = `position: absolute;
|
||||
top: 14px;
|
||||
left: 20px;
|
||||
background-color: #00000075;
|
||||
padding: 2px 4px;
|
||||
color: white;
|
||||
font-size: 12px;`
|
||||
const div = document.createElement("div");
|
||||
// div.className = 'card';
|
||||
div.appendChild(im)
|
||||
if (isMain) div.appendChild(p)
|
||||
im.setAttribute('onerror', `this.src='${base64Df}'`)
|
||||
return div
|
||||
}
|
||||
|
||||
// 创建图库选择
|
||||
function createSelectForImages(title, options, index = 0, callback) {
|
||||
|
||||
const div = document.createElement("div");
|
||||
// div.className = 'card';
|
||||
div.style = `margin-top: 24px;`
|
||||
|
||||
// Create a label for the upload control
|
||||
// const nameLabel = document.createElement("label");
|
||||
// nameLabel.textContent = title;
|
||||
// div.appendChild(nameLabel);
|
||||
|
||||
var mainImg = createImageForSelect(true, options[index].imgurl, options[index].keyword);
|
||||
div.appendChild(mainImg);
|
||||
|
||||
let imgs = document.createElement('div');
|
||||
div.appendChild(imgs);
|
||||
imgs.style = `display:flex; flex-wrap: wrap;`
|
||||
for (const opt of options) {
|
||||
var selectElement = createImageForSelect(false, opt.imgurl, opt.keyword);
|
||||
imgs.appendChild(selectElement);
|
||||
selectElement.addEventListener('click', async e => {
|
||||
e.preventDefault();
|
||||
mainImg.querySelector('p').innerText = opt.keyword;
|
||||
mainImg.querySelector('img').src = opt.imgurl;
|
||||
if (callback) {
|
||||
if (!opt.imgurl.match('data:image')) {
|
||||
opt.imgurl = await parseImageToBase64(opt.imgurl)
|
||||
}
|
||||
callback(opt.imgurl, opt.keyword);
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
return [div, mainImg];
|
||||
}
|
||||
|
||||
|
||||
// 创建下拉选择
|
||||
function createSelect(options, defaultValue) {
|
||||
var selectElement = document.createElement("select");
|
||||
@@ -1791,7 +2109,8 @@
|
||||
|
||||
function createUI(data, share = true) {
|
||||
// appData.input, appData.output, appData.seed, share, appData.link
|
||||
const { input: inputData, output: outputData, data: workflow, seed, link, name } = data;
|
||||
if (!data) return
|
||||
const { input: inputData, output: outputData, data: workflow, seed, seedTitle, link, name } = data;
|
||||
|
||||
let mainDiv = document.createElement('div');
|
||||
|
||||
@@ -1889,7 +2208,7 @@
|
||||
|
||||
let em = document.createElement('em');
|
||||
let emText = document.createElement('span');
|
||||
emText.innerText = `#${id} ${s.toUpperCase()}`;
|
||||
emText.innerText = `#${seedTitle && seedTitle[id] ? seedTitle[id] : id} ${s.toUpperCase()}`;
|
||||
em.appendChild(emText);
|
||||
|
||||
seedInput.appendChild(em)
|
||||
@@ -1910,7 +2229,7 @@
|
||||
data.seed[id] = 'randomize';
|
||||
inSeed.style.display = 'none'
|
||||
}
|
||||
emText.innerText = `#${id} ${data.seed[id].toUpperCase()}`;
|
||||
emText.innerText = `#${seedTitle && seedTitle[id] ? seedTitle[id] : id} ${data.seed[id].toUpperCase()}`;
|
||||
});
|
||||
|
||||
seedInput.appendChild(inSeed);
|
||||
@@ -2083,6 +2402,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: {
|
||||
@@ -2198,6 +2562,9 @@
|
||||
const _images = detail?.output?._images;
|
||||
const prompts = detail?.output?.prompts;
|
||||
|
||||
// 3d模型
|
||||
const meshes = detail?.output?.mesh;
|
||||
|
||||
if (images) {
|
||||
// if (!images) return;
|
||||
|
||||
@@ -2207,6 +2574,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();
|
||||
|
||||
@@ -2440,6 +2815,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';
|
||||
|
||||
@@ -2640,20 +3018,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);
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -2,9 +2,46 @@ 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'
|
||||
|
||||
const parseImageToBase64 = url => {
|
||||
return new Promise((res, rej) => {
|
||||
fetch(url)
|
||||
.then(response => response.blob())
|
||||
.then(blob => {
|
||||
const reader = new FileReader()
|
||||
reader.onloadend = () => {
|
||||
const base64data = reader.result
|
||||
res(base64data)
|
||||
// 在这里可以将base64数据用于进一步处理或显示图片
|
||||
}
|
||||
reader.readAsDataURL(blob)
|
||||
})
|
||||
.catch(error => {
|
||||
console.log('发生错误:', error)
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
function get_position_style (ctx, widget_width, y, node_height) {
|
||||
const MARGIN = 12 // the margin around the html element
|
||||
|
||||
@@ -82,6 +119,7 @@ async function extractInputAndOutputData (
|
||||
let input = []
|
||||
let output = []
|
||||
const seed = {}
|
||||
const seedTitle = {}
|
||||
|
||||
for (const id in data) {
|
||||
if (data.hasOwnProperty(id)) {
|
||||
@@ -121,6 +159,28 @@ async function extractInputAndOutputData (
|
||||
}
|
||||
}
|
||||
|
||||
if (node.type == 'ImagesPrompt_') {
|
||||
//图库
|
||||
// console.log('ImagesPrompt_', data[id])
|
||||
let image_base64 = data[id].inputs.image_base64
|
||||
let img_index = 0
|
||||
let imgsData = JSON.parse(data[id].inputs.upload)
|
||||
for (let index = 0; index < imgsData.length; index++) {
|
||||
const imgd = imgsData[index].imgurl
|
||||
imgsData[index].index = index
|
||||
//TODO缩放大小
|
||||
imgsData[index].imgurl = await parseImageToBase64(imgd)
|
||||
if (image_base64 == imgsData[index].imgurl) {
|
||||
img_index = index
|
||||
}
|
||||
}
|
||||
options.images = imgsData
|
||||
delete data[id].inputs.upload
|
||||
delete data[id].inputs.image_base64
|
||||
|
||||
data[id].inputs.imageIndex = img_index
|
||||
}
|
||||
|
||||
if (node.type == 'Color') {
|
||||
}
|
||||
|
||||
@@ -133,9 +193,9 @@ async function extractInputAndOutputData (
|
||||
}
|
||||
// loadImage的默认图,转为base64
|
||||
let imgurl = app.graph.getNodeById(id).imgs[0].src
|
||||
|
||||
|
||||
options.defaultImage = await drawImageToCanvas(imgurl, 512)
|
||||
console.log('#loadImage的默认图',options)
|
||||
console.log('#loadImage的默认图', options)
|
||||
}
|
||||
|
||||
input[inputIds.indexOf(id)] = {
|
||||
@@ -152,12 +212,18 @@ async function extractInputAndOutputData (
|
||||
output[outputIds.indexOf(id)] = { ...data[id], title: node.title, id }
|
||||
}
|
||||
|
||||
if (node.type === 'KSampler' || node.type == 'SamplerCustom') {
|
||||
if (
|
||||
node.type === 'KSampler' ||
|
||||
node.type == 'SamplerCustom' ||
|
||||
node.type === 'ChinesePrompt_Mix' ||
|
||||
node.type === 'Seed_'
|
||||
) {
|
||||
// seed 的类型收集
|
||||
try {
|
||||
seed[id] = node.widgets.filter(
|
||||
w => w.name === 'seed' || w.name == 'noise_seed'
|
||||
)[0].linkedWidgets[0].value
|
||||
seedTitle[id] = node.title
|
||||
} catch (error) {}
|
||||
}
|
||||
}
|
||||
@@ -167,7 +233,7 @@ async function extractInputAndOutputData (
|
||||
input = input.filter(i => i)
|
||||
output = output.filter(i => i)
|
||||
|
||||
return { input, output, seed }
|
||||
return { input, output, seed, seedTitle }
|
||||
}
|
||||
|
||||
function getUrl () {
|
||||
@@ -219,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], //用于分享的功能扩展
|
||||
@@ -240,7 +309,14 @@ async function save (json, download = false, showInfo = true) {
|
||||
try {
|
||||
let data = await app.graphToPrompt()
|
||||
|
||||
let { input, output, seed } = await extractInputAndOutputData(
|
||||
//从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,
|
||||
outputIds
|
||||
@@ -260,6 +336,7 @@ async function save (json, download = false, showInfo = true) {
|
||||
input,
|
||||
output,
|
||||
seed, //控制是fixed 还是random
|
||||
seedTitle,
|
||||
share_prefix,
|
||||
link,
|
||||
category,
|
||||
@@ -301,12 +378,13 @@ async function save (json, download = false, showInfo = true) {
|
||||
|
||||
function getInputsAndOutputs () {
|
||||
const inputs =
|
||||
`LoadImage 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`.split(
|
||||
' '
|
||||
)
|
||||
outputs =
|
||||
`SaveTripoSRMesh,PreviewImage,SaveImage,TransparentImage,ShowTextForGPT,VHS_VideoCombine,VideoCombine_Adv,Image Save,SaveImageAndMetadata_,ClipInterrogator`.split(
|
||||
','
|
||||
)
|
||||
|
||||
let inputsId = [],
|
||||
outputsId = []
|
||||
@@ -329,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
|
||||
|
||||
@@ -3,7 +3,7 @@ import { app } from '../../../scripts/app.js'
|
||||
const repoOwner = 'shadowcz007' // 替换为仓库的所有者
|
||||
const repoName = 'comfyui-mixlab-nodes' // 替换为仓库的名称
|
||||
|
||||
const version = 'v0.18.0'
|
||||
const version = 'v0.23.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;
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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();
|
||||
@@ -28,6 +60,9 @@ async function uploadImage (blob, fileType = '.svg', filename) {
|
||||
return src
|
||||
}
|
||||
|
||||
const base64Df =
|
||||
'data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAwAAAAMCAYAAABWdVznAAAAAXNSR0IArs4c6QAAALZJREFUKFOFkLERwjAQBPdbgBkInECGaMLUQDsE0AkRVRAYWqAByxldPPOWHwnw4OBGye1p50UDSoA+W2ABLPN7i+C5dyC6R/uiAUXRQCs0bXoNIu4QPQzAxDKxHoALOrZcqtiyR/T6CXw7+3IGHhkYcy6BOR2izwT8LptG8rbMiCRAUb+CQ6WzQVb0SNOi5Z2/nX35DRyb/ENazhpWKoGwrpD6nICp5c2qogc4of+c7QcrhgF4Aa/aoAFHiL+RAAAAAElFTkSuQmCC'
|
||||
|
||||
function base64ToBlobFromURL (base64URL, contentType) {
|
||||
return fetch(base64URL).then(response => response.blob())
|
||||
}
|
||||
@@ -108,7 +143,7 @@ function createImage (url) {
|
||||
})
|
||||
}
|
||||
|
||||
const parseImage = url => {
|
||||
const parseImageToBase64 = url => {
|
||||
return new Promise((res, rej) => {
|
||||
fetch(url)
|
||||
.then(response => response.blob())
|
||||
@@ -406,9 +441,7 @@ app.registerExtension({
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
// Fires every time a node is constructed
|
||||
@@ -442,3 +475,369 @@ app.registerExtension({
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
const createSelect = (imgDiv, select, opts, targetWidget, textWidget) => {
|
||||
select.style.display = 'block'
|
||||
let html = ''
|
||||
let isMatch = false
|
||||
for (const opt of opts) {
|
||||
html += `<option value='${opt.keyword}' ${opt.selected ? 'selected' : ''}>${
|
||||
opt.keyword
|
||||
}</option>`
|
||||
if (opt.selected) {
|
||||
isMatch = true
|
||||
imgDiv.src = opt.imgurl
|
||||
// targetWidget.value = opt.keyword
|
||||
}
|
||||
}
|
||||
select.innerHTML = html
|
||||
if (!isMatch) {
|
||||
// targetWidget.value = opts[0].keyword
|
||||
imgDiv.src = opts[0].imgurl
|
||||
}
|
||||
|
||||
// 添加change事件监听器
|
||||
select.addEventListener('change', async function () {
|
||||
// 获取选中的选项的值
|
||||
var selectedOption = select.options[select.selectedIndex].value
|
||||
let t = opts.filter(opt => opt.keyword === selectedOption)[0]
|
||||
|
||||
targetWidget.value = await parseImageToBase64(t.imgurl)
|
||||
imgDiv.src = targetWidget.value
|
||||
textWidget.value = t.keyword
|
||||
})
|
||||
// console.log(select)
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.prompt.ImagesPrompt_',
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'ImagesPrompt_') {
|
||||
const orig_nodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = async function () {
|
||||
orig_nodeCreated?.apply(this, arguments)
|
||||
|
||||
const image_prompt = this.widgets.filter(
|
||||
w => w.name == 'image_base64'
|
||||
)[0]
|
||||
const image_text = this.widgets.filter(w => w.name == 'text')[0]
|
||||
|
||||
const node = this
|
||||
|
||||
const widget = {
|
||||
type: 'div',
|
||||
name: 'upload',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, widget_width, y, node.size[1])
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
widget.div = $el('div', {})
|
||||
|
||||
// console.log('image_prompt',image_prompt)
|
||||
const img = new Image()
|
||||
img.src = image_prompt?.value || base64Df
|
||||
widget.div.appendChild(img)
|
||||
|
||||
const btn = document.createElement('button')
|
||||
btn.innerText = 'Upload Images JSON'
|
||||
|
||||
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;
|
||||
`
|
||||
|
||||
const select = document.createElement('select')
|
||||
select.style = `display:none;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: 100px;
|
||||
`
|
||||
widget.select = select
|
||||
|
||||
// const btn=document.createElement('button');
|
||||
// btn.innerText='Upload'
|
||||
btn.addEventListener('click', () => {
|
||||
let inp = document.createElement('input')
|
||||
inp.type = 'file'
|
||||
inp.accept = '.json'
|
||||
inp.click()
|
||||
inp.addEventListener('change', event => {
|
||||
// 获取选择的文件
|
||||
// [{title,imageUrl}]
|
||||
const file = event.target.files[0]
|
||||
this.title = file.name.split('.')[0]
|
||||
|
||||
// console.log(file.name.split('.')[0])
|
||||
// 创建文件读取器
|
||||
const reader = new FileReader()
|
||||
|
||||
// 定义读取完成事件的回调函数
|
||||
reader.onload = async event => {
|
||||
// 读取完成后的文本内容
|
||||
const json = JSON.parse(event.target.result)
|
||||
console.log(node, json)
|
||||
|
||||
widget.value = JSON.stringify(json)
|
||||
|
||||
let img = widget.div.querySelector('img')
|
||||
|
||||
createSelect(img, select, json, image_prompt, image_text)
|
||||
|
||||
image_prompt.value = await parseImageToBase64(json[0].imgurl)
|
||||
image_text.value = json[0].keyword
|
||||
|
||||
if (img) {
|
||||
img.src = image_prompt.value
|
||||
}
|
||||
|
||||
inp.remove()
|
||||
}
|
||||
|
||||
// 以文本方式读取文件
|
||||
reader.readAsText(file)
|
||||
})
|
||||
})
|
||||
|
||||
widget.div.appendChild(btn)
|
||||
widget.div.appendChild(select)
|
||||
document.body.appendChild(widget.div)
|
||||
this.addCustomWidget(widget)
|
||||
|
||||
const onRemoved = this.onRemoved
|
||||
this.onRemoved = () => {
|
||||
widget.div.remove()
|
||||
return onRemoved?.()
|
||||
}
|
||||
|
||||
if (this.onResize) {
|
||||
this.onResize(this.size)
|
||||
}
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
}
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
if (node.type === 'ImagesPrompt_') {
|
||||
try {
|
||||
let prompt = node.widgets.filter(w => w.name === 'image_base64')[0]
|
||||
let text = node.widgets.filter(w => w.name === 'text')[0]
|
||||
let uploadWidget = node.widgets.filter(w => w.name == 'upload')[0]
|
||||
// console.log('##prompt',prompt.value)
|
||||
let img = uploadWidget.div.querySelector('img')
|
||||
let json = JSON.parse(uploadWidget.value)
|
||||
|
||||
for (let index = 0; index < json.length; index++) {
|
||||
const j = json[index]
|
||||
let base64 = await parseImageToBase64(j.imgurl)
|
||||
if (base64 === prompt.value) {
|
||||
json[index].selected = true
|
||||
}
|
||||
}
|
||||
|
||||
if (json && json[0]) {
|
||||
uploadWidget.select.style.display = 'block'
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
+566
-14
@@ -1,8 +1,70 @@
|
||||
import { app } from '../../../scripts/app.js'
|
||||
// import { api } from '../../../scripts/api.js'
|
||||
import { api } from '../../../scripts/api.js'
|
||||
import { ComfyWidgets } from '../../../scripts/widgets.js'
|
||||
import { $el } from '../../../scripts/ui.js'
|
||||
|
||||
function downloadJsonFile (jsonData, fileName = 'grid.json') {
|
||||
const dataString = JSON.stringify(jsonData)
|
||||
const blob = new Blob([dataString], { type: 'application/json' })
|
||||
const url = URL.createObjectURL(blob)
|
||||
|
||||
const link = document.createElement('a')
|
||||
link.href = url
|
||||
link.download = fileName
|
||||
link.click()
|
||||
|
||||
// 释放URL对象
|
||||
setTimeout(() => {
|
||||
URL.revokeObjectURL(url)
|
||||
}, 0)
|
||||
}
|
||||
|
||||
function createSelectWithOptions (options) {
|
||||
const select = document.createElement('select')
|
||||
|
||||
options.forEach(option => {
|
||||
const optionElement = document.createElement('option')
|
||||
optionElement.text = option
|
||||
optionElement.value = option
|
||||
select.appendChild(optionElement)
|
||||
})
|
||||
|
||||
select.style = `cursor: pointer;
|
||||
font-weight: 300;
|
||||
height: 30px;
|
||||
min-width: 122px;
|
||||
position: absolute;
|
||||
top: 24px;
|
||||
left: 88px;
|
||||
z-index: 999999999999999;
|
||||
`
|
||||
|
||||
return select
|
||||
}
|
||||
|
||||
function drawCanvasWithText (w, h, tag, color = 'rgba(255,255,255,0.4)') {
|
||||
const canvas = document.createElement('canvas')
|
||||
const ctx = canvas.getContext('2d')
|
||||
|
||||
// 设置画布大小
|
||||
canvas.width = w
|
||||
canvas.height = h
|
||||
|
||||
// 绘制白色背景
|
||||
ctx.fillStyle = color
|
||||
ctx.fillRect(0, 0, canvas.width, canvas.height)
|
||||
|
||||
// 绘制文字
|
||||
ctx.fillStyle = '#000000'
|
||||
ctx.font = '20px Arial'
|
||||
ctx.fillText(tag, 50, 50)
|
||||
|
||||
// 导出为Base64
|
||||
const base64 = canvas.toDataURL()
|
||||
|
||||
return base64
|
||||
}
|
||||
|
||||
function get_position_style (ctx, widget_width, y, node_height) {
|
||||
const MARGIN = 4 // the margin around the html element
|
||||
|
||||
@@ -156,32 +218,29 @@ const parseSvg = async svgContent => {
|
||||
return { data, image: base64, svgElement }
|
||||
}
|
||||
|
||||
|
||||
function findImages(nodeId) {
|
||||
function findImages (nodeId) {
|
||||
// 检查当前节点是否有 imgs 字段
|
||||
const n = app.graph.getNodeById(nodeId)
|
||||
if (n.imgs) {
|
||||
return n.imgs;
|
||||
return n.imgs
|
||||
}
|
||||
|
||||
// 检查当前节点的 inputs 是否有 image 字段
|
||||
if (n.inputs) {
|
||||
for (let i = 0; i < n.inputs.length; i++) {
|
||||
if (n.inputs[i].name==='image'||n.inputs[i].name==='images') {
|
||||
if (n.inputs[i].name === 'image' || n.inputs[i].name === 'images') {
|
||||
// 获取新的 nodeId,并递归调用 findImages 函数
|
||||
var linkId = n.inputs[i]?.link;
|
||||
var linkId = n.inputs[i]?.link
|
||||
var origin_id = app.graph.links[linkId].origin_id
|
||||
return findImages(origin_id);
|
||||
return findImages(origin_id)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 如果没有找到 imgs 字段或者 image 字段,则返回 null
|
||||
return null;
|
||||
return null
|
||||
}
|
||||
|
||||
|
||||
|
||||
async function setArea (cw, ch, topBase64, base64, data, fn) {
|
||||
let displayHeight = Math.round(window.screen.availHeight * 0.8)
|
||||
let div = document.createElement('div')
|
||||
@@ -353,6 +412,196 @@ async function setArea (cw, ch, topBase64, base64, data, fn) {
|
||||
}
|
||||
}
|
||||
|
||||
async function setAreaTags (cw, ch, grids, fn) {
|
||||
let base64 = drawCanvasWithText(cw, ch, '', 'white')
|
||||
let displayHeight = Math.round(window.screen.availHeight * 0.8)
|
||||
let div = document.createElement('div')
|
||||
div.innerHTML = `
|
||||
<div id='ml_overlay' style='position: absolute;top:0;background: #251f1fc4;
|
||||
height: 100vh;
|
||||
z-index:999999;
|
||||
width: 100%;'>
|
||||
<img id='ml_video' style='position: absolute;
|
||||
height: ${displayHeight}px;user-select: none;
|
||||
-webkit-user-drag: none;
|
||||
outline: 2px solid #eaeaea;
|
||||
box-shadow: 8px 9px 17px #575757;' />
|
||||
${Array.from(grids, g => {
|
||||
const { label: tag, grid } = g
|
||||
const [dx, dy, dw, dh] = grid
|
||||
const base64Data = drawCanvasWithText(dw, dh, tag)
|
||||
|
||||
let x = 0,
|
||||
y = 0,
|
||||
width = (cw * displayHeight) / ch,
|
||||
height = displayHeight
|
||||
|
||||
let imgWidth = cw
|
||||
let imgHeight = ch
|
||||
|
||||
if (dw > 0 && dh > 0) {
|
||||
// 相同尺寸窗口,恢复选区
|
||||
x = (width * dx) / imgWidth
|
||||
y = (height * dy) / imgHeight
|
||||
width = (width * dw) / imgWidth
|
||||
height = (height * dh) / imgHeight
|
||||
}
|
||||
|
||||
return `<div class='ml_selection'
|
||||
data-tag="${tag}"
|
||||
style='position:absolute;
|
||||
border: 2px dashed red;
|
||||
pointer-events: none;
|
||||
background-image: url("${base64Data}");
|
||||
background-repeat: no-repeat;
|
||||
background-size: cover;
|
||||
left:${x}px;
|
||||
top:${y}px;
|
||||
width:${width}px;
|
||||
height:${height}px;
|
||||
'></div>`
|
||||
})}
|
||||
<div class="mx_close"> X </div>
|
||||
</div>`
|
||||
// document.body.querySelector('#ml_overlay')
|
||||
document.body.appendChild(div)
|
||||
|
||||
const tags = Array.from(grids, g => g.label)
|
||||
let select = createSelectWithOptions(tags)
|
||||
document.body.appendChild(select)
|
||||
|
||||
let img = div.querySelector('#ml_video')
|
||||
// let overlay = div.querySelector('#ml_overlay')
|
||||
let selections = [...div.querySelectorAll('.ml_selection')]
|
||||
|
||||
let selection = selections.filter(
|
||||
s => s.getAttribute('data-tag') === select.value
|
||||
)[0]
|
||||
|
||||
select.addEventListener('change', e => {
|
||||
selection = selections.filter(
|
||||
s => s.getAttribute('data-tag') === select.value
|
||||
)[0]
|
||||
})
|
||||
|
||||
// console.log(select.value,selection)
|
||||
let close = div.querySelector('.mx_close')
|
||||
let startX, startY, endX, endY
|
||||
let start = false
|
||||
let setDone = false
|
||||
// Set video source
|
||||
img.src = base64
|
||||
// canvas.toDataURL();
|
||||
close.style = `cursor: pointer;
|
||||
position: fixed;
|
||||
left: 12px;
|
||||
top: 12px;
|
||||
z-index: 99999999;
|
||||
background: black;
|
||||
width: 44px;
|
||||
height: 44px;
|
||||
text-align: center;
|
||||
line-height: 44px;`
|
||||
|
||||
// Add mouse events
|
||||
img.addEventListener('mousedown', startSelection)
|
||||
img.addEventListener('mousemove', updateSelection)
|
||||
img.addEventListener('mouseup', endSelection)
|
||||
|
||||
const removeDiv = () => {
|
||||
div.remove()
|
||||
select?.remove()
|
||||
close.removeEventListener('click', removeDiv)
|
||||
img.removeEventListener('mousedown', startSelection)
|
||||
img.removeEventListener('mousemove', updateSelection)
|
||||
img.removeEventListener('mouseup', endSelection)
|
||||
img.removeEventListener('mousedown', setDoneCheck)
|
||||
}
|
||||
close.addEventListener('click', removeDiv)
|
||||
|
||||
const setDoneCheck = event => {
|
||||
console.log(setDone)
|
||||
if (setDone) {
|
||||
img.addEventListener('mousedown', startSelection)
|
||||
img.addEventListener('mousemove', updateSelection)
|
||||
img.addEventListener('mouseup', endSelection)
|
||||
setDone = false
|
||||
start = false
|
||||
startX = event.clientX
|
||||
startY = event.clientY
|
||||
}
|
||||
}
|
||||
img.addEventListener('mousedown', setDoneCheck)
|
||||
|
||||
function remove () {
|
||||
img.removeEventListener('mousedown', startSelection)
|
||||
img.removeEventListener('mousemove', updateSelection)
|
||||
img.removeEventListener('mouseup', endSelection)
|
||||
setDone = true
|
||||
// select?.remove()
|
||||
}
|
||||
|
||||
function startSelection (event) {
|
||||
if (start == false) {
|
||||
startX = event.clientX
|
||||
startY = event.clientY
|
||||
updateSelection(event)
|
||||
start = true
|
||||
} else {
|
||||
}
|
||||
}
|
||||
|
||||
function updateSelection (event) {
|
||||
endX = event.clientX
|
||||
endY = event.clientY
|
||||
|
||||
// Calculate width, height, and coordinates
|
||||
let width = Math.abs(endX - startX)
|
||||
let height = Math.abs(endY - startY)
|
||||
let left = Math.min(startX, endX)
|
||||
let top = Math.min(startY, endY)
|
||||
|
||||
// Set selection style
|
||||
selection.style.left = left + 'px'
|
||||
selection.style.top = top + 'px'
|
||||
selection.style.width = width + 'px'
|
||||
selection.style.height = height + 'px'
|
||||
}
|
||||
|
||||
function endSelection (event) {
|
||||
endX = event.clientX
|
||||
endY = event.clientY
|
||||
|
||||
// 获取img元素的真实宽度和高度
|
||||
let imgWidth = img.naturalWidth
|
||||
let imgHeight = img.naturalHeight
|
||||
|
||||
// 换算起始坐标
|
||||
let realStartX = (startX / img.offsetWidth) * imgWidth
|
||||
let realStartY = (startY / img.offsetHeight) * imgHeight
|
||||
|
||||
// 换算起始坐标
|
||||
let realEndX = (endX / img.offsetWidth) * imgWidth
|
||||
let realEndY = (endY / img.offsetHeight) * imgHeight
|
||||
|
||||
startX = realStartX
|
||||
startY = realStartY
|
||||
endX = realEndX
|
||||
endY = realEndY
|
||||
// Calculate width, height, and coordinates
|
||||
let width = Math.round(Math.abs(endX - startX))
|
||||
let height = Math.round(Math.abs(endY - startY))
|
||||
let left = Math.round(Math.min(startX, endX))
|
||||
let top = Math.round(Math.min(startY, endY))
|
||||
|
||||
if (width <= 0 && height <= 0) return remove()
|
||||
|
||||
if (!!fn) fn(select.value, left, top, width, height)
|
||||
|
||||
remove()
|
||||
}
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.layer.ShowLayer',
|
||||
async getCustomWidgets (app) {
|
||||
@@ -598,8 +847,8 @@ app.registerExtension({
|
||||
}
|
||||
try {
|
||||
console.log('this.inputs', this.id)
|
||||
let imgs=findImages(this.id)
|
||||
|
||||
let imgs = findImages(this.id)
|
||||
|
||||
// let topLinkId = this.inputs[0].link
|
||||
// let topNodeId = app.graph.links[topLinkId].origin_id
|
||||
let topIm = imgs[0]
|
||||
@@ -607,9 +856,9 @@ app.registerExtension({
|
||||
let linkId = this.inputs[3].link
|
||||
let nodeId = app.graph.links[linkId].origin_id
|
||||
// console.log(linkId,this.inputs)
|
||||
let imgs2=findImages(nodeId)
|
||||
let imgs2 = findImages(nodeId)
|
||||
let im = imgs2[0]
|
||||
console.log(topIm,im)
|
||||
console.log(topIm, im)
|
||||
// let src = im.src
|
||||
setArea(
|
||||
im.naturalWidth,
|
||||
@@ -641,3 +890,306 @@ app.registerExtension({
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.layer.GridInput',
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'GridInput') {
|
||||
const orig_nodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = async function () {
|
||||
orig_nodeCreated?.apply(this, arguments)
|
||||
|
||||
const grids_widget = this.widgets.filter(w => w.name == 'grids')[0]
|
||||
|
||||
const widget = {
|
||||
type: 'div',
|
||||
name: 'upload',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, widget_width, y, node.size[1]),
|
||||
{
|
||||
justifyContent: 'flex-start'
|
||||
}
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
widget.div = $el('div', {})
|
||||
|
||||
const addBtn = document.createElement('button')
|
||||
addBtn.innerText = 'Add Box'
|
||||
addBtn.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;
|
||||
`
|
||||
|
||||
const vbtn = document.createElement('button')
|
||||
vbtn.innerText = 'Set Box'
|
||||
vbtn.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;
|
||||
`
|
||||
|
||||
const btn = document.createElement('button')
|
||||
btn.innerText = 'Upload JSON'
|
||||
|
||||
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;
|
||||
`
|
||||
|
||||
addBtn.addEventListener('click', () => {
|
||||
const { width, height, grids } = JSON.parse(grids_widget.value)
|
||||
grids.push({
|
||||
label: 'background',
|
||||
grid: [12, 12, width - 24, height - 24]
|
||||
})
|
||||
grids_widget.value = JSON.stringify(
|
||||
{
|
||||
width,
|
||||
height,
|
||||
grids
|
||||
},
|
||||
null,
|
||||
2
|
||||
)
|
||||
})
|
||||
|
||||
vbtn.addEventListener('click', () => {
|
||||
const { width, height, grids } = JSON.parse(grids_widget.value)
|
||||
|
||||
setAreaTags(width, height, grids, (tag, x, y, w, h) => {
|
||||
grids_widget.value = JSON.stringify(
|
||||
{
|
||||
width,
|
||||
height,
|
||||
grids: Array.from(grids, g => {
|
||||
if (g.label === tag) {
|
||||
g.grid = [x, y, w, h]
|
||||
}
|
||||
return g
|
||||
})
|
||||
},
|
||||
null,
|
||||
2
|
||||
)
|
||||
})
|
||||
})
|
||||
|
||||
btn.addEventListener('click', () => {
|
||||
let inp = document.createElement('input')
|
||||
inp.type = 'file'
|
||||
inp.accept = '.json'
|
||||
inp.click()
|
||||
inp.addEventListener('change', event => {
|
||||
// 获取选择的文件
|
||||
const file = event.target.files[0]
|
||||
this.title = file.name.split('.')[0]
|
||||
|
||||
// console.log(file.name.split('.')[0])
|
||||
// 创建文件读取器
|
||||
const reader = new FileReader()
|
||||
|
||||
// 定义读取完成事件的回调函数
|
||||
reader.onload = event => {
|
||||
// 读取完成后的文本内容
|
||||
const fileContent = JSON.parse(event.target.result)
|
||||
const grids = fileContent
|
||||
grids_widget.value = JSON.stringify(grids, null, 2)
|
||||
// widget.value = grids
|
||||
|
||||
inp.remove()
|
||||
}
|
||||
|
||||
// 以文本方式读取文件
|
||||
reader.readAsText(file)
|
||||
})
|
||||
})
|
||||
|
||||
widget.div.appendChild(addBtn)
|
||||
widget.div.appendChild(vbtn)
|
||||
widget.div.appendChild(btn)
|
||||
document.body.appendChild(widget.div)
|
||||
this.addCustomWidget(widget)
|
||||
|
||||
const onExecuted = nodeType.prototype.onExecuted
|
||||
nodeType.prototype.onExecuted = function (message) {
|
||||
const r = onExecuted?.apply?.(this, arguments)
|
||||
|
||||
let json = message.json
|
||||
if (json) {
|
||||
json = {
|
||||
width: json[0],
|
||||
height: json[1],
|
||||
grids: json[2]
|
||||
}
|
||||
grids_widget.value = JSON.stringify(json, null, 2)
|
||||
// widget.value = json
|
||||
}
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
const onRemoved = this.onRemoved
|
||||
this.onRemoved = () => {
|
||||
widget.div.remove()
|
||||
return onRemoved?.()
|
||||
}
|
||||
|
||||
if (this.onResize) {
|
||||
this.onResize(this.size)
|
||||
}
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
}
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
if (node.type === 'GridInput') {
|
||||
try {
|
||||
const grids_widget = node.widgets.filter(w => w.name == 'grids')[0]
|
||||
const { width, height, grids } = JSON.parse(grids_widget.value)
|
||||
console.log('#GridInput', node, grids)
|
||||
|
||||
const div = node.widgets.filter(w => w.name == 'upload')[0]
|
||||
div.div.querySelector('select').innerHTML = Array.from(
|
||||
grids,
|
||||
g => `<option value="${g.label}">${g.label}</option>`
|
||||
).join('')
|
||||
} catch (error) {}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
app.registerExtension({
|
||||
name: 'Mixlab.layer.GridDisplayAndSave',
|
||||
async beforeRegisterNodeDef (nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass == 'GridDisplayAndSave') {
|
||||
const orig_nodeCreated = nodeType.prototype.onNodeCreated
|
||||
nodeType.prototype.onNodeCreated = async function () {
|
||||
orig_nodeCreated?.apply(this, arguments)
|
||||
|
||||
const grids_widget = this.widgets.filter(w => w.name == 'grids')[0]
|
||||
console.log('GridDisplayAndSave', grids_widget)
|
||||
const widget = {
|
||||
type: 'div',
|
||||
name: 'save_json',
|
||||
draw (ctx, node, widget_width, y, widget_height) {
|
||||
Object.assign(
|
||||
this.div.style,
|
||||
get_position_style(ctx, widget_width, y, node.size[1]),
|
||||
{
|
||||
justifyContent: 'flex-start',
|
||||
flexDirection: 'column'
|
||||
}
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
widget.div = $el('div', {})
|
||||
|
||||
const btn = document.createElement('button')
|
||||
btn.innerText = 'Save JSON'
|
||||
|
||||
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;
|
||||
max-width: 122px;
|
||||
`
|
||||
|
||||
btn.addEventListener('click', () => {
|
||||
if (window._mixlab_grid)
|
||||
downloadJsonFile(
|
||||
window._mixlab_grid,
|
||||
this.widgets.filter(w => w.name == 'filename_prefix')[0]?.value +
|
||||
'_grid.json'
|
||||
)
|
||||
})
|
||||
|
||||
widget.div.appendChild(btn)
|
||||
document.body.appendChild(widget.div)
|
||||
this.addCustomWidget(widget)
|
||||
|
||||
const onExecuted = nodeType.prototype.onExecuted
|
||||
nodeType.prototype.onExecuted = function (message) {
|
||||
const r = onExecuted?.apply?.(this, arguments)
|
||||
let save_json = this.widgets.filter(d => d.name == 'save_json')[0]
|
||||
let div = save_json?.div
|
||||
// console.log('Test',message)
|
||||
|
||||
let image = message.image[0]
|
||||
let json = message.json
|
||||
if (image) {
|
||||
const { filename, subfolder, type } = image
|
||||
|
||||
if (!div.querySelector('img')) {
|
||||
let im = new Image()
|
||||
div.appendChild(im)
|
||||
im.style.width = '100%'
|
||||
}
|
||||
div.querySelector('img').src = api.apiURL(
|
||||
`/view?filename=${encodeURIComponent(
|
||||
filename
|
||||
)}&type=${type}&subfolder=${subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`
|
||||
)
|
||||
|
||||
window._mixlab_grid = {
|
||||
width: json[0],
|
||||
height: json[1],
|
||||
grids: json[2]
|
||||
}
|
||||
// console.log(src)
|
||||
}
|
||||
|
||||
this.onResize?.(this.size)
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
const onRemoved = this.onRemoved
|
||||
this.onRemoved = () => {
|
||||
widget.div.remove()
|
||||
return onRemoved?.()
|
||||
}
|
||||
|
||||
if (this.onResize) {
|
||||
this.onResize(this.size)
|
||||
}
|
||||
|
||||
this.serialize_widgets = true //需要保存参数
|
||||
}
|
||||
}
|
||||
},
|
||||
async loadedGraphNode (node, app) {
|
||||
if (node.type === 'GridDisplayAndSave') {
|
||||
try {
|
||||
let grids_widget = node.widgets.filter(w => w.name === 'grids')[0]
|
||||
// let ks = getLocalData(`_mixlab_PromptSlide`)
|
||||
let uploadWidget = node.widgets.filter(w => w.name == 'upload')[0]
|
||||
// console.log('##widget', uploadWidget.value)
|
||||
let grids = JSON.parse(uploadWidget.value)
|
||||
} catch (error) {}
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
@@ -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)
|
||||
// }
|
||||
// }
|
||||
// }
|
||||
}
|
||||
})
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
+121
-13
@@ -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
|
||||
}
|
||||
|
||||
@@ -943,6 +944,20 @@ app.registerExtension({
|
||||
}
|
||||
]
|
||||
|
||||
if (node.widgets) {
|
||||
// let text_widget = node.widgets.filter(
|
||||
// w => w.name === 'text' && typeof w.value == 'string'
|
||||
// )
|
||||
// if (text_widget && text_widget.length == 1) {
|
||||
// opts.push({
|
||||
// content: 'Text-to-Text ♾️Mixlab', // with a name
|
||||
// callback: () => {
|
||||
// LGraphCanvas.prototype.text2text(node)
|
||||
// } // and the callback
|
||||
// })
|
||||
// }
|
||||
}
|
||||
|
||||
opts = addSmartMenu(opts, node)
|
||||
|
||||
// if (node.type == 'CLIPTextEncode') {
|
||||
@@ -1123,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 () => {
|
||||
@@ -1132,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) {
|
||||
@@ -1145,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 => {
|
||||
@@ -1497,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
|
||||
@@ -1573,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)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
})()
|
||||
Vendored
+81
-64
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
Reference in New Issue
Block a user