modify comfyui workflow

This commit is contained in:
皓童
2025-03-11 22:53:41 +08:00
parent 80de810ad5
commit 6ad009403b
11 changed files with 4767 additions and 25 deletions
+61 -20
View File
@@ -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

+5 -5
View File
@@ -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'
+194
View File
@@ -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})
+14
View File
@@ -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