modify comfyui workflow
This commit is contained in:
@@ -1,7 +1,5 @@
|
||||
<p align="center">
|
||||
|
||||
|
||||
|
||||
<h2 align="center"><img src="assets/figures/icon.png" height=16> ++: Instruction-Based Image Creation and Editing <br> via Context-Aware Content Filling </h2>
|
||||
|
||||
<p align="center">
|
||||
@@ -41,13 +39,11 @@
|
||||
|
||||
The original intention behind the design of ACE++ was to unify reference image generation, local editing,
|
||||
and controllable generation into a single framework, and to enable one model to adapt to a wider range of tasks.
|
||||
A more versatile model is often capable of handling more complex tasks. We have already released three LoRA models,
|
||||
focusing on portraits, objects, and regional editing, with the expectation that each would demonstrate strong adaptability
|
||||
within their respective domains. Undoubtedly, this presents certain challenges.
|
||||
|
||||
We are currently training a fully fine-tuned model, which has now entered the final stage of quality tuning.
|
||||
We are confident it will be released soon. This model will support a broader range of capabilities and is
|
||||
expected to empower community developers to build even more interesting applications.
|
||||
A more versatile model is often capable of handling more complex tasks. "We have released three LoRA models for
|
||||
specific vertical domains and a more versatile FFT model. Users can flexibly utilize these models and their
|
||||
combinations for their own scenarios. Furthermore, many community members have found that using them
|
||||
in conjunction with Redux modules significantly improves performance. We believe there are many more
|
||||
use cases to explore.
|
||||
|
||||
## 📢 News
|
||||
- [x] **[2025.01.06]** Release the code and models of ACE++.
|
||||
@@ -55,7 +51,22 @@ expected to empower community developers to build even more interesting applicat
|
||||
- [x] **[2025.01.16]** Release the training code for lora.
|
||||
- [x] **[2025.02.15]** Collection of workflows in Comfyui.
|
||||
- [x] **[2025.02.15]** Release the config for fully fine-tuning.
|
||||
- [x] **[2025.03.03]** Release a unified fft model for ACE++, support more image to image tasks. [HuggingFace](https://huggingface.co/ali-vilab/ACE_Plus/tree/main)
|
||||
- [x] **[2025.03.03]** Release a unified fft model for ACE++, support more image to image tasks.
|
||||
- [x] **[2025.03.11]** Release the comfyui workflow and the fp8
|
||||
version stored in [ms](https://www.modelscope.cn/models/iic/ACE_Plus/file/view/master?fileName=ace_plus_fft_fp8.safetensors&status=2) and [hf](https://huggingface.co/ali-vilab/ACE_Plus/blob/main/ace_plus_fft_fp8.safetensors) for ACE++ FFT model.
|
||||
|
||||
- We sincerely apologize
|
||||
for the delayed responses and updates regarding ACE++ issues.
|
||||
Further development of the ACE model through post-training on the FLUX model must be suspended.
|
||||
We have identified several significant challenges in post-training on the FLUX foundation.
|
||||
The primary issue is the high degree of heterogeneity between the training dataset and the FLUX model,
|
||||
which results in highly unstable training. Moreover, FLUX-Dev is a distilled model, and the influence of its original negative prompts on its final performance is uncertain.
|
||||
As a result, subsequent efforts will be focused on post-training the ACE model using the Wan series of foundational models.
|
||||
|
||||
- We have been busy with other projects recently. Our new work in the video domain, [VACE](https://ali-vilab.github.io/VACE-Page/),
|
||||
has also been released, and we welcome you to continue following our work.
|
||||
|
||||
|
||||
|
||||
## 🔥The unified fft model for ACE++
|
||||
Fully finetuning a composite model with ACE’s data to support various editing and reference generation tasks through an instructive approach.
|
||||
@@ -66,6 +77,42 @@ To address this issue, we introduced 64 additional channels in the channel dimen
|
||||
|
||||
One issue with this approach is that it changes the input channel number of the FLUX-Fill-Dev model from 384 to 448. The specific configuration can be referenced in the [configuration file](config/ace_plus_fft.yaml).
|
||||
|
||||
|
||||
We used tools from [stella](https://gist.github.com/Stella2211/10f5bd870387ec1ddb9932235321068e)(this is really a great work) to convert the fft-fp16 model to fft-fp8. The updated models are available on [ms](https://www.modelscope.cn/models/iic/ACE_Plus/file/view/master?fileName=ace_plus_fft_fp8.safetensors&status=2) and [hf](https://huggingface.co/ali-vilab/ACE_Plus/blob/main/ace_plus_fft_fp8.safetensors). The results of the fp8 model will differ from the fp16 model. We have not yet performed a rigorous comparison, so users should be aware of this.
|
||||
|
||||
### ComfyUI Workflow
|
||||
|
||||
Copy the workflow/ComfyUI-ACE_Plus folder into ComfyUI’s custom_nodes directory. Launch ComfyUI, and we have provided four example workflows in workflow_example_fft with the following explanations.
|
||||
|
||||
We provide a parameter to adjust the GPU memory usage. As shown in the figure below, max_seq_length controls the length of the token sequence during inference, thereby controlling the model's inference memory consumption. The range of this value is from 1024 to 5120, and it correspondingly affects the clarity of the generated image. The smaller the value, the lower the image clarity.
|
||||
|
||||
<img src="./assets/comfyui/snapshot.jpg" width="800">
|
||||
|
||||
<table><tbody>
|
||||
<tr>
|
||||
<td>Workflow</td>
|
||||
<td>Description</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>ACE_Plus_FFT_workflow_no_preprocess.json</td>
|
||||
<td>Use the preprocessed images, such as depth and contour, as input, or the super-resolution.</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>ACE_Plus_FFT_workflow_controlpreprocess.json</td>
|
||||
<td>Controllable image-to-image translation capability.</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>ACE_Plus_FFT_workflow_reference_generation.json</td>
|
||||
<td>Reference image generation capability for portrait or subject.</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>ACE_Plus_FFT_workflow_referenceediting_generation.json</td>
|
||||
<td>Reference image editing capability</td>
|
||||
</tr>
|
||||
<tbody>
|
||||
<table>
|
||||
|
||||
|
||||
### Examples
|
||||
<table><tbody>
|
||||
<tr>
|
||||
@@ -270,10 +317,6 @@ Additionally, many bloggers have published tutorials on how to use it, which are
|
||||
</table>
|
||||
|
||||
|
||||
## 🔥 ACE Models
|
||||
ACE++ provides a comprehensive toolkit for image editing and generation to support various applications. We encourage developers to choose the appropriate model based on their own scenarios and to fine-tune their models using data from their specific scenarios to achieve more stable results.
|
||||
|
||||
|
||||
|
||||
|
||||
### ACE++ Portrait
|
||||
@@ -435,9 +478,7 @@ python infer_lora.py
|
||||
The relevant commands for fft models are as follows:
|
||||
```bash
|
||||
export FLUX_FILL_PATH="hf://black-forest-labs/FLUX.1-Fill-dev"
|
||||
export ACE_PLUS_FFT_MODEL="ms://iic/ACE_Plus@ace_plus_fft.safetensors.safetensors"
|
||||
# Use the model from huggingface
|
||||
# export "ACE_PLUS_FFT_MODEL=hf://ali-vilab/ACE_Plus@ace_plus_fft.safetensors"
|
||||
export ACE_PLUS_FFT_MODEL="ms://iic/ACE_Plus@ace_plus_fft.safetensors.safetensors"
|
||||
python infer_fft.py
|
||||
```
|
||||
|
||||
@@ -456,11 +497,11 @@ The required fields include the following six, with their explanations as follow
|
||||
All parameters related to training are stored in 'train_config/ace_plus_lora.yaml'. To run the training code, execute the following command.
|
||||
|
||||
```bash
|
||||
export FLUX_FILL_PATH="hf://black-forest-labs/FLUX.1-Fill-dev"
|
||||
export FLUX_FILL_PATH="{path to FLUX.1-Fill-dev}"
|
||||
python run_train.py --cfg train_config/ace_plus_lora.yaml
|
||||
# Training from fft model
|
||||
export FLUX_FILL_PATH="hf://black-forest-labs/FLUX.1-Fill-dev"
|
||||
export ACE_PLUS_FFT_MODEL="ms://iic/ACE_Plus@ace_plus_fft.safetensors.safetensors"
|
||||
export FLUX_FILL_PATH="{path to FLUX.1-Fill-dev}"
|
||||
export ACE_PLUS_FFT_MODEL="path to ace_plus_fft.safetensors.safetensors"
|
||||
python run_train.py --cfg train_config/ace_plus_fft.yaml
|
||||
```
|
||||
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 47 KiB |
@@ -72,7 +72,7 @@ MODEL:
|
||||
FIRST_STAGE_MODEL:
|
||||
NAME: AutoencoderKLFlux
|
||||
EMBED_DIM: 16
|
||||
PRETRAINED_MODEL: ${FLUX_FILL_PATH}@ae.safetensors
|
||||
PRETRAINED_MODEL: ${FLUX_FILL_PATH}/ae.safetensors
|
||||
IGNORE_KEYS: [ ]
|
||||
BATCH_SIZE: 8
|
||||
USE_CONV: False
|
||||
@@ -116,11 +116,11 @@ MODEL:
|
||||
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_MODEL_CLS: T5EncoderModel
|
||||
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
|
||||
MODEL_PATH: ${FLUX_FILL_PATH}@text_encoder_2/
|
||||
MODEL_PATH: ${FLUX_FILL_PATH}/text_encoder_2/
|
||||
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_TOKENIZER_CLS: T5Tokenizer
|
||||
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
|
||||
TOKENIZER_PATH: ${FLUX_FILL_PATH}@tokenizer_2/
|
||||
TOKENIZER_PATH: ${FLUX_FILL_PATH}/tokenizer_2/
|
||||
ADDED_IDENTIFIER: [ '<img>','{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
|
||||
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
|
||||
MAX_LENGTH: 512
|
||||
@@ -138,11 +138,11 @@ MODEL:
|
||||
# HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_MODEL_CLS: CLIPTextModel
|
||||
# MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None
|
||||
MODEL_PATH: ${FLUX_FILL_PATH}@text_encoder/
|
||||
MODEL_PATH: ${FLUX_FILL_PATH}/text_encoder/
|
||||
# HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None
|
||||
HF_TOKENIZER_CLS: CLIPTokenizer
|
||||
# TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None
|
||||
TOKENIZER_PATH: ${FLUX_FILL_PATH}@tokenizer/
|
||||
TOKENIZER_PATH: ${FLUX_FILL_PATH}/tokenizer/
|
||||
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
|
||||
MAX_LENGTH: 77
|
||||
# OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state'
|
||||
|
||||
@@ -0,0 +1,194 @@
|
||||
# Copy from https://gist.github.com/Stella2211/10f5bd870387ec1ddb9932235321068e
|
||||
# This is a great work.
|
||||
import json
|
||||
from pathlib import Path
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
import struct
|
||||
from typing import Dict, Any
|
||||
import sys
|
||||
|
||||
# input file
|
||||
if(len(sys.argv) < 3):
|
||||
print("Usage: mem_eff_fp8_convert.py {fp16 model path} {output path}")
|
||||
sys.exit(1)
|
||||
|
||||
path = sys.argv[1]
|
||||
output =sys.argv[2]
|
||||
model_file = Path(path)
|
||||
|
||||
class MemoryEfficientSafeOpen:
|
||||
# does not support metadata loading
|
||||
def __init__(self, filename):
|
||||
self.filename = filename
|
||||
self.header, self.header_size = self._read_header()
|
||||
self.file = open(filename, "rb")
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
self.file.close()
|
||||
|
||||
def keys(self):
|
||||
return [k for k in self.header.keys() if k != "__metadata__"]
|
||||
|
||||
def get_tensor(self, key):
|
||||
if key not in self.header:
|
||||
raise KeyError(f"Tensor '{key}' not found in the file")
|
||||
|
||||
metadata = self.header[key]
|
||||
offset_start, offset_end = metadata["data_offsets"]
|
||||
|
||||
if offset_start == offset_end:
|
||||
tensor_bytes = None
|
||||
else:
|
||||
# adjust offset by header size
|
||||
self.file.seek(self.header_size + 8 + offset_start)
|
||||
tensor_bytes = self.file.read(offset_end - offset_start)
|
||||
|
||||
return self._deserialize_tensor(tensor_bytes, metadata)
|
||||
|
||||
def _read_header(self):
|
||||
with open(self.filename, "rb") as f:
|
||||
header_size = struct.unpack("<Q", f.read(8))[0]
|
||||
header_json = f.read(header_size).decode("utf-8")
|
||||
return json.loads(header_json), header_size
|
||||
|
||||
def _deserialize_tensor(self, tensor_bytes, metadata):
|
||||
dtype = self._get_torch_dtype(metadata["dtype"])
|
||||
shape = metadata["shape"]
|
||||
|
||||
if tensor_bytes is None:
|
||||
byte_tensor = torch.empty(0, dtype=torch.uint8)
|
||||
else:
|
||||
tensor_bytes = bytearray(tensor_bytes) # make it writable
|
||||
byte_tensor = torch.frombuffer(tensor_bytes, dtype=torch.uint8)
|
||||
|
||||
# process float8 types
|
||||
if metadata["dtype"] in ["F8_E5M2", "F8_E4M3"]:
|
||||
return self._convert_float8(byte_tensor, metadata["dtype"], shape)
|
||||
|
||||
# convert to the target dtype and reshape
|
||||
return byte_tensor.view(dtype).reshape(shape)
|
||||
|
||||
@staticmethod
|
||||
def _get_torch_dtype(dtype_str):
|
||||
dtype_map = {
|
||||
"F64": torch.float64,
|
||||
"F32": torch.float32,
|
||||
"F16": torch.float16,
|
||||
"BF16": torch.bfloat16,
|
||||
"I64": torch.int64,
|
||||
"I32": torch.int32,
|
||||
"I16": torch.int16,
|
||||
"I8": torch.int8,
|
||||
"U8": torch.uint8,
|
||||
"BOOL": torch.bool,
|
||||
}
|
||||
# add float8 types if available
|
||||
if hasattr(torch, "float8_e5m2"):
|
||||
dtype_map["F8_E5M2"] = torch.float8_e5m2
|
||||
if hasattr(torch, "float8_e4m3fn"):
|
||||
dtype_map["F8_E4M3"] = torch.float8_e4m3fn
|
||||
return dtype_map.get(dtype_str)
|
||||
|
||||
@staticmethod
|
||||
def _convert_float8(byte_tensor, dtype_str, shape):
|
||||
if dtype_str == "F8_E5M2" and hasattr(torch, "float8_e5m2"):
|
||||
return byte_tensor.view(torch.float8_e5m2).reshape(shape)
|
||||
elif dtype_str == "F8_E4M3" and hasattr(torch, "float8_e4m3fn"):
|
||||
return byte_tensor.view(torch.float8_e4m3fn).reshape(shape)
|
||||
else:
|
||||
# # convert to float16 if float8 is not supported
|
||||
# print(f"Warning: {dtype_str} is not supported in this PyTorch version. Converting to float16.")
|
||||
# return byte_tensor.view(torch.uint8).to(torch.float16).reshape(shape)
|
||||
raise ValueError(f"Unsupported float8 type: {dtype_str} (upgrade PyTorch to support float8 types)")
|
||||
|
||||
|
||||
def mem_eff_save_file(tensors: Dict[str, torch.Tensor], filename: str, metadata: Dict[str, Any] = None):
|
||||
_TYPES = {
|
||||
torch.float64: "F64",
|
||||
torch.float32: "F32",
|
||||
torch.float16: "F16",
|
||||
torch.bfloat16: "BF16",
|
||||
torch.int64: "I64",
|
||||
torch.int32: "I32",
|
||||
torch.int16: "I16",
|
||||
torch.int8: "I8",
|
||||
torch.uint8: "U8",
|
||||
torch.bool: "BOOL",
|
||||
getattr(torch, "float8_e5m2", None): "F8_E5M2",
|
||||
getattr(torch, "float8_e4m3fn", None): "F8_E4M3",
|
||||
}
|
||||
_ALIGN = 256
|
||||
|
||||
def validate_metadata(metadata: Dict[str, Any]) -> Dict[str, str]:
|
||||
validated = {}
|
||||
for key, value in metadata.items():
|
||||
if not isinstance(key, str):
|
||||
raise ValueError(f"Metadata key must be a string, got {type(key)}")
|
||||
if not isinstance(value, str):
|
||||
print(f"Warning: Metadata value for key '{key}' is not a string. Converting to string.")
|
||||
validated[key] = str(value)
|
||||
else:
|
||||
validated[key] = value
|
||||
return validated
|
||||
|
||||
header = {}
|
||||
offset = 0
|
||||
if metadata:
|
||||
header["__metadata__"] = validate_metadata(metadata)
|
||||
for k, v in tensors.items():
|
||||
if v.numel() == 0: # empty tensor
|
||||
header[k] = {"dtype": _TYPES[v.dtype], "shape": list(v.shape), "data_offsets": [offset, offset]}
|
||||
else:
|
||||
size = v.numel() * v.element_size()
|
||||
header[k] = {"dtype": _TYPES[v.dtype], "shape": list(v.shape), "data_offsets": [offset, offset + size]}
|
||||
offset += size
|
||||
|
||||
hjson = json.dumps(header).encode("utf-8")
|
||||
hjson += b" " * (-(len(hjson) + 8) % _ALIGN)
|
||||
|
||||
with open(filename, "wb") as f:
|
||||
f.write(struct.pack("<Q", len(hjson)))
|
||||
f.write(hjson)
|
||||
|
||||
for k, v in tensors.items():
|
||||
if v.numel() == 0:
|
||||
continue
|
||||
if v.is_cuda:
|
||||
# Direct GPU to disk save
|
||||
with torch.cuda.device(v.device):
|
||||
if v.dim() == 0: # if scalar, need to add a dimension to work with view
|
||||
v = v.unsqueeze(0)
|
||||
tensor_bytes = v.contiguous().view(torch.uint8)
|
||||
tensor_bytes.cpu().numpy().tofile(f)
|
||||
else:
|
||||
# CPU tensor save
|
||||
if v.dim() == 0: # if scalar, need to add a dimension to work with view
|
||||
v = v.unsqueeze(0)
|
||||
v.contiguous().view(torch.uint8).numpy().tofile(f)
|
||||
|
||||
|
||||
# read safetensors metadata
|
||||
def read_safetensors_metadata(path: str):
|
||||
with open(path, 'rb') as f:
|
||||
header_size = int.from_bytes(f.read(8), 'little')
|
||||
header_json = f.read(header_size).decode('utf-8')
|
||||
header = json.loads(header_json)
|
||||
metadata = header.get('__metadata__', {})
|
||||
return metadata
|
||||
|
||||
metadata = read_safetensors_metadata(path)
|
||||
print(json.dumps(metadata, indent=4)) #show metadata
|
||||
|
||||
sd_pruned = dict() #initialize empty dict
|
||||
|
||||
with MemoryEfficientSafeOpen(path) as reader:
|
||||
keys = reader.keys()
|
||||
for key in tqdm(keys): #for each key in the safetensors file
|
||||
sd_pruned[key] = reader.get_tensor(key).to(torch.float8_e4m3fn) #convert to fp8
|
||||
|
||||
# save the pruned safetensors file
|
||||
mem_eff_save_file(sd_pruned, output, metadata={"format": "pt", **metadata})
|
||||
@@ -0,0 +1,14 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from .ace_plus_fft_node import ACEPlusFFTLoader, ACEPlusFFTConditioning, AcePlusFFTProcessor
|
||||
|
||||
NODE_MAPPINGS = {
|
||||
'ACEPlusLoader': ('ACEPlusFFTLoader~', ACEPlusFFTLoader),
|
||||
'ACEPlusConditioning': ('ACEPlusFFTConditioning~', ACEPlusFFTConditioning),
|
||||
'ACEPlusFFTProcessor': ('ACEPlusFFTProcessor~', AcePlusFFTProcessor)
|
||||
}
|
||||
|
||||
NODE_CLASS_MAPPINGS = {k: v[1] for k, v in NODE_MAPPINGS.items()}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {k: v[0] for k, v in NODE_MAPPINGS.items()}
|
||||
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
@@ -0,0 +1,348 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import folder_paths, os
|
||||
from comfy.supported_models import FluxInpaint, models
|
||||
from nodes import UNETLoader
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
try:
|
||||
from scepter.modules.annotator.registry import ANNOTATORS
|
||||
from scepter.modules.utils.config import Config
|
||||
|
||||
fs_list = [
|
||||
Config(cfg_dict={"NAME": "HuggingfaceFs", "TEMP_DIR": os.environ["TEMP_DIR"]}, load=False),
|
||||
Config(cfg_dict={"NAME": "ModelscopeFs", "TEMP_DIR": os.environ["TEMP_DIR"]}, load=False),
|
||||
Config(cfg_dict={"NAME": "HttpFs", "TEMP_DIR": os.environ["TEMP_DIR"]}, load=False),
|
||||
Config(cfg_dict={"NAME": "LocalFs", "TEMP_DIR": os.environ["TEMP_DIR"]}, load=False),
|
||||
Config(cfg_dict={"NAME": "LocalFs", "TEMP_DIR": os.environ["TEMP_DIR"]}, load=False)
|
||||
]
|
||||
|
||||
for one_fs in fs_list:
|
||||
FS.init_fs_client(one_fs)
|
||||
SCEPTER = True
|
||||
except:
|
||||
SCEPTER = False
|
||||
|
||||
class ACEPlus(FluxInpaint):
|
||||
unet_config = {
|
||||
"image_model": "flux",
|
||||
"guidance_embed": True,
|
||||
"in_channels": 112,
|
||||
}
|
||||
|
||||
|
||||
class ACEPlusFFTLoader(UNETLoader):
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {"unet_name": (folder_paths.get_filename_list("diffusion_models"), ),
|
||||
"weight_dtype": (["default", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e5m2"],)
|
||||
}}
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "load_unet"
|
||||
CATEGORY = "ComfyUI-ACE_Plus"
|
||||
|
||||
def load_unet(self, unet_name, weight_dtype):
|
||||
models.append(ACEPlus)
|
||||
return super().load_unet(unet_name, weight_dtype)
|
||||
|
||||
|
||||
import torch
|
||||
import node_helpers
|
||||
|
||||
|
||||
class ACEPlusFFTConditioning:
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {"positive": ("CONDITIONING", ),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"vae": ("VAE", ),
|
||||
"ucpixels": ("IMAGE", ),
|
||||
"cpixels": ("IMAGE", ),
|
||||
"mask": ("MASK", ),
|
||||
"noise_mask": ("BOOLEAN", {"default": True, "tooltip": "Add a noise mask to the latent "
|
||||
"so sampling will only happen "
|
||||
"within the mask. Might improve "
|
||||
"results or completely break "
|
||||
"things depending on the model."}),
|
||||
}}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING", "CONDITIONING", "LATENT")
|
||||
RETURN_NAMES = ("positive", "negative", "latent")
|
||||
FUNCTION = "encode"
|
||||
|
||||
CATEGORY = "ComfyUI-ACE_Plus"
|
||||
|
||||
def encode(self,
|
||||
positive,
|
||||
negative,
|
||||
vae,
|
||||
ucpixels,
|
||||
cpixels,
|
||||
mask,
|
||||
noise_mask=True):
|
||||
x = (ucpixels.shape[1] // 8) * 8
|
||||
y = (ucpixels.shape[2] // 8) * 8
|
||||
mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])),
|
||||
size=(ucpixels.shape[1], ucpixels.shape[2]), mode="bilinear")
|
||||
|
||||
orig_pixels = ucpixels
|
||||
pixels = orig_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]
|
||||
|
||||
orig_c_pixels = cpixels
|
||||
c_pixels = orig_c_pixels.clone()
|
||||
if orig_c_pixels.shape[1] != x or orig_c_pixels.shape[2] != y:
|
||||
x_offset = (orig_c_pixels.shape[1] % 8) // 2
|
||||
y_offset = (orig_c_pixels.shape[2] % 8) // 2
|
||||
c_pixels = orig_c_pixels[:, x_offset:x + x_offset, y_offset:y + y_offset, :]
|
||||
|
||||
concat_latent = vae.encode(pixels)
|
||||
orig_latent = vae.encode(orig_pixels)
|
||||
c_concat_latent = vae.encode(c_pixels)
|
||||
|
||||
out_latent = {"samples": orig_latent}
|
||||
if noise_mask:
|
||||
out_latent["noise_mask"] = mask
|
||||
|
||||
out = []
|
||||
for conditioning in [positive, negative]:
|
||||
c = node_helpers.conditioning_set_values(conditioning, {
|
||||
"concat_latent_image": torch.cat([concat_latent, c_concat_latent], dim=1),
|
||||
"concat_mask": mask})
|
||||
out.append(c)
|
||||
|
||||
return (out[0], out[1], out_latent)
|
||||
|
||||
|
||||
import torch
|
||||
import math
|
||||
import os
|
||||
import yaml
|
||||
import torchvision.transforms as T
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
|
||||
class AcePlusFFTProcessor:
|
||||
def __init__(self,
|
||||
max_aspect_ratio=4,
|
||||
d=16,
|
||||
max_seq_len=1024):
|
||||
self.max_aspect_ratio = max_aspect_ratio
|
||||
self.max_seq_len = max_seq_len
|
||||
self.d = d
|
||||
self.processor_cfg = self.load_yaml(os.path.join('custom_nodes/ComfyUI-ACE_Plus/',
|
||||
'config/ace_plus_fft_processor.yaml'))
|
||||
self.task_list = {}
|
||||
for task in self.processor_cfg['PREPROCESSOR']:
|
||||
self.task_list[task['TYPE']] = task
|
||||
self.transforms = T.Compose([
|
||||
T.ToTensor(),
|
||||
T.Normalize(mean=[0, 0, 0], std=[1.0, 1.0, 1.0])
|
||||
])
|
||||
|
||||
CATEGORY = 'ComfyUI-ACE_Plus'
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
'required': {
|
||||
'use_reference': ('BOOLEAN', {'default': True}),
|
||||
'height': ('INT', {
|
||||
'default': 1024,
|
||||
'min': 256,
|
||||
'max': 1436,
|
||||
'step': 16
|
||||
}),
|
||||
'width': ('INT', {
|
||||
'default': 1024,
|
||||
'min': 256,
|
||||
'max': 1436,
|
||||
'step': 16
|
||||
}),
|
||||
'task_type': (list(s().task_list.keys()),),
|
||||
'keep_pixels_rate': ('FLOAT', {
|
||||
'default': 0.8,
|
||||
'min': 0,
|
||||
'max': 1,
|
||||
'step': 0.01
|
||||
}),
|
||||
'max_seq_length': ('INT', {
|
||||
'default': 3072,
|
||||
'min': 1024,
|
||||
'max': 5120,
|
||||
'step': 0.01
|
||||
}),
|
||||
},
|
||||
'optional': {
|
||||
'reference_image': ('IMAGE',),
|
||||
'edit_image': ('IMAGE',),
|
||||
'edit_mask': ('MASK',),
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
OUTPUT_NODE = True
|
||||
RETURN_TYPES = ('IMAGE', 'IMAGE', 'MASK', 'INT', 'INT', 'INT')
|
||||
RETURN_NAMES = ('UC_IMAGE', 'C_IMAGE', 'MASK', 'OUT_H', 'OUT_W', 'SLICE_W')
|
||||
FUNCTION = 'preprocess'
|
||||
|
||||
def load_yaml(self, cfg_file):
|
||||
with open(cfg_file, 'r') as f:
|
||||
cfg = yaml.load(f.read(), Loader=yaml.SafeLoader)
|
||||
return cfg
|
||||
|
||||
def image_check(self, image):
|
||||
if image is None:
|
||||
return image
|
||||
# preprocess
|
||||
H, W = image.shape[1: 3]
|
||||
image = image.permute(0, 3, 1, 2)
|
||||
if H / W > self.max_aspect_ratio:
|
||||
image[0] = T.CenterCrop([int(self.max_aspect_ratio * W), W])(image[0])
|
||||
elif W / H > self.max_aspect_ratio:
|
||||
image[0] = T.CenterCrop([H, int(self.max_aspect_ratio * H)])(image[0])
|
||||
return image[0]
|
||||
|
||||
def trans_pil_tensor(self, pil_image):
|
||||
transform = T.Compose([
|
||||
T.ToTensor()
|
||||
])
|
||||
tensor_image = transform(pil_image)
|
||||
return tensor_image
|
||||
|
||||
def edit_preprocess(self, processor, device, edit_image, edit_mask):
|
||||
|
||||
if edit_image is None or processor is None:
|
||||
return edit_image
|
||||
if not SCEPTER:
|
||||
raise ImportError(f'Please install scepter to use edit processor {processor} by '
|
||||
f'runing "pip install scepter" in the conda env')
|
||||
processor = Config(cfg_dict=processor, load=False)
|
||||
processor = ANNOTATORS.build(processor).to(device)
|
||||
edit_image = Image.fromarray(np.array(edit_image[0] * 255).astype(np.uint8)).convert('RGB')
|
||||
new_edit_image = processor(np.asarray(edit_image))
|
||||
|
||||
del processor
|
||||
new_edit_image = Image.fromarray(new_edit_image)
|
||||
if edit_mask is not None:
|
||||
edit_mask = np.where(edit_mask > 0.5, 1, 0) * 255
|
||||
edit_mask = Image.fromarray(np.array(edit_mask[0]).astype(np.uint8)).convert('L')
|
||||
|
||||
if new_edit_image.size != edit_image.size:
|
||||
edit_image = T.Resize((edit_image.size[1], edit_image.size[0]),
|
||||
interpolation=T.InterpolationMode.BILINEAR,
|
||||
antialias=True)(new_edit_image)
|
||||
|
||||
image = Image.composite(new_edit_image, edit_image, edit_mask)
|
||||
|
||||
return self.trans_pil_tensor(image).unsqueeze(0).permute(0, 2, 3, 1)
|
||||
|
||||
def preprocess(self,
|
||||
reference_image=None,
|
||||
edit_image=None,
|
||||
edit_mask=None,
|
||||
use_reference=True,
|
||||
task_type=None,
|
||||
height=1024,
|
||||
width=1024,
|
||||
keep_pixels_rate=0.8,
|
||||
max_seq_length=4096):
|
||||
self.max_seq_len = max_seq_length
|
||||
if not use_reference and edit_image is not None:
|
||||
reference_image = None
|
||||
if edit_mask is not None and edit_image is not None:
|
||||
iH, iW = edit_image.shape[1:3]
|
||||
mH, mW = edit_mask.shape[1:3]
|
||||
if iH != mH or iW != mW:
|
||||
edit_mask = torch.ones(edit_image.shape[:3])
|
||||
|
||||
if task_type != 'repainting':
|
||||
repainting_scale = 0
|
||||
else:
|
||||
repainting_scale = 1
|
||||
if task_type in self.task_list:
|
||||
edit_image = self.edit_preprocess(self.task_list[task_type]['ANNOTATOR'], 0,
|
||||
edit_image, edit_mask)
|
||||
if reference_image is not None:
|
||||
reference_image = self.image_check(reference_image) - 0.5
|
||||
if edit_image is not None:
|
||||
edit_image = self.image_check(edit_image) - 0.5
|
||||
# for reference generation
|
||||
if edit_image is None:
|
||||
edit_image = torch.zeros([3, height, width])
|
||||
edit_mask = torch.ones([1, height, width])
|
||||
else:
|
||||
if edit_mask is None:
|
||||
_, eH, eW = edit_image.shape
|
||||
edit_mask = np.ones((eH, eW))
|
||||
else:
|
||||
edit_mask = np.asarray(edit_mask)[0]
|
||||
edit_mask = np.where(edit_mask > 0.5, 1, 0)
|
||||
edit_mask = edit_mask.astype(
|
||||
np.float32) if np.any(edit_mask) else np.ones_like(edit_mask).astype(
|
||||
np.float32)
|
||||
edit_mask = torch.tensor(edit_mask).unsqueeze(0)
|
||||
|
||||
edit_image = edit_image * (1 - edit_mask * repainting_scale)
|
||||
|
||||
out_h, out_w = edit_image.shape[-2:]
|
||||
|
||||
assert edit_mask is not None
|
||||
if reference_image is not None:
|
||||
_, H, W = reference_image.shape
|
||||
_, eH, eW = edit_image.shape
|
||||
if not True:
|
||||
# align height with edit_image
|
||||
scale = eH / H
|
||||
tH, tW = eH, int(W * scale)
|
||||
reference_image = T.Resize((tH, tW), interpolation=T.InterpolationMode.BILINEAR, antialias=True)(
|
||||
reference_image)
|
||||
else:
|
||||
# padding
|
||||
if H >= keep_pixels_rate * eH:
|
||||
tH = int(eH * keep_pixels_rate)
|
||||
scale = tH / H
|
||||
tW = int(W * scale)
|
||||
reference_image = T.Resize((tH, tW), interpolation=T.InterpolationMode.BILINEAR, antialias=True)(
|
||||
reference_image)
|
||||
rH, rW = reference_image.shape[-2:]
|
||||
delta_w = 0
|
||||
delta_h = eH - rH
|
||||
padding = (delta_w // 2, delta_h // 2, delta_w - (delta_w // 2), delta_h - (delta_h // 2))
|
||||
reference_image = T.Pad(padding, fill=0, padding_mode="constant")(reference_image)
|
||||
edit_image = torch.cat([reference_image, edit_image], dim=-1)
|
||||
edit_mask = torch.cat([torch.zeros([1, reference_image.shape[1], reference_image.shape[2]]), edit_mask],
|
||||
dim=-1)
|
||||
slice_w = reference_image.shape[-1]
|
||||
else:
|
||||
slice_w = 0
|
||||
|
||||
H, W = edit_image.shape[-2:]
|
||||
scale = min(1.0, math.sqrt(self.max_seq_len * 2 / ((H / self.d) * (W / self.d))))
|
||||
rH = int(H * scale) // self.d * self.d
|
||||
rW = int(W * scale) // self.d * self.d
|
||||
slice_w = int(slice_w * scale) // self.d * self.d
|
||||
|
||||
edit_image = T.Resize((rH, rW), interpolation=T.InterpolationMode.NEAREST_EXACT, antialias=True)(edit_image)
|
||||
edit_mask = T.Resize((rH, rW), interpolation=T.InterpolationMode.NEAREST_EXACT, antialias=True)(edit_mask)
|
||||
|
||||
change_image = edit_image * edit_mask
|
||||
edit_image = edit_image * (1 - edit_mask)
|
||||
edit_image = edit_image.unsqueeze(0).permute(0, 2, 3, 1)
|
||||
change_image = change_image.unsqueeze(0).permute(0, 2, 3, 1)
|
||||
slice_w = slice_w if slice_w < 30 else slice_w + 30
|
||||
|
||||
return edit_image + 0.5, change_image + 0.5, edit_mask, out_h, out_w, slice_w
|
||||
@@ -0,0 +1,25 @@
|
||||
PREPROCESSOR:
|
||||
- TYPE: repainting
|
||||
REPAINTING_SCALE: 1.0
|
||||
ANNOTATOR:
|
||||
- TYPE: no_preprocess
|
||||
REPAINTING_SCALE: 0.0
|
||||
ANNOTATOR:
|
||||
- TYPE: contour_repainting
|
||||
REPAINTING_SCALE: 0.0
|
||||
ANNOTATOR:
|
||||
NAME: InfoDrawContourAnnotator
|
||||
INPUT_NC: 3
|
||||
OUTPUT_NC: 1
|
||||
N_RESIDUAL_BLOCKS: 3
|
||||
SIGMOID: True
|
||||
PRETRAINED_MODEL: "ms://iic/scepter_annotator@annotator/ckpts/informative_drawing_contour_style.pth"
|
||||
- TYPE: depth_repainting
|
||||
REPAINTING_SCALE: 0.0
|
||||
ANNOTATOR:
|
||||
NAME: MidasDetector
|
||||
PRETRAINED_MODEL: "ms://iic/scepter_annotator@annotator/ckpts/dpt_hybrid-midas-501f0c75.pt"
|
||||
- TYPE: recolorizing
|
||||
REPAINTING_SCALE: 0.0
|
||||
ANNOTATOR:
|
||||
NAME: GrayAnnotator
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user