modify demo fft

This commit is contained in:
皓童
2025-03-04 09:37:13 +08:00
parent aa0c15d8b6
commit 16d682baa5
11 changed files with 548 additions and 813 deletions
+138 -8
View File
@@ -53,9 +53,128 @@ 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.
- [] **[ToDo]** Release a unified fft model for ACE++, support more image to image tasks.
- [x] **[2025.03.03]** Release a unified fft model for ACE++, support more image to image tasks.
## 🔥 Comfyui Workflows in community
## 🔥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.
We found that there are conflicts between the repainting task and the editing task during the experimental process. This is because the edited image is concatenated with noise in the channel dimension, whereas the repainting task modifies the region using zero pixel values in the VAE's latent space. The editing task uses RGB pixel values in the modified region through the VAE's latent space, which is similar to the distribution of the non-modified part of the repainting task, making it a challenge for the model to distinguish between the two tasks.
To address this issue, we introduced 64 additional channels in the channel dimension to differentiate between these two tasks. In these channels, we place the latent representation of the pixel space from the edited image, while keeping other channels consistent with the repainting task. This approach significantly enhances the model's adaptability to different tasks.
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).
### Examples
<table><tbody>
<tr>
<td>Input Reference Image</td>
<td>Input Edit Image</td>
<td>Input Edit Mask</td>
<td>Output</td>
<td>Instruction</td>
<td>Function</td>
</tr>
<tr>
<td><img src="./assets/samples/portrait/human_1.jpg" width="200"></td>
<td></td>
<td></td>
<td><img src="./assets/samples/portrait/human_1_fft.webp" width="200"></td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"Maintain the facial features, A girl is wearing a neat police uniform and sporting a badge. She is smiling with a friendly and confident demeanor. The background is blurred, featuring a cartoon logo."</td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"Character ID Consistency Generation"</td>
</tr>
<tr>
<td><img src="./assets/samples/subject/subject_1.jpg" width="200"></td>
<td></td>
<td></td>
<td><img src="./assets/samples/subject/subject_1_fft.webp" width="200"></td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"Display the logo in a minimalist style printed in white on a matte black ceramic coffee mug, alongside a steaming cup of coffee on a cozy cafe table."</td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"Subject Consistency Generation"</td>
</tr>
<tr>
<td><img src="./assets/samples/application/photo_editing/1_ref.png" width="200"></td>
<td><img src="./assets/samples/application/photo_editing/1_2_edit.jpg" width="200"></td>
<td><img src="./assets/samples/application/photo_editing/1_2_m.webp" width="200"></td>
<td><img src="./assets/samples/application/photo_editing/1_2_fft.webp" width="200"></td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"The item is put on the table."</td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"Subject Consistency Editing"</td>
</tr>
<tr>
<td><img src="./assets/samples/application/logo_paste/1_ref.png" width="200"></td>
<td><img src="./assets/samples/application/logo_paste/1_1_edit.png" width="200"></td>
<td><img src="./assets/samples/application/logo_paste/1_1_m.png" width="200"></td>
<td><img src="./assets/samples/application/logo_paste/1_1_fft.webp" width="200"></td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"The logo is printed on the headphones."</td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"Subject Consistency Editing"</td>
</tr>
<tr>
<td><img src="./assets/samples/application/try_on/1_ref.png" width="200"></td>
<td><img src="./assets/samples/application/try_on/1_1_edit.png" width="200"></td>
<td><img src="./assets/samples/application/try_on/1_1_m.png" width="200"></td>
<td><img src="./assets/samples/application/try_on/1_1_fft.webp" width="200"></td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"The woman dresses this skirt."</td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"Try On"</td>
</tr>
<tr>
<td><img src="./assets/samples/application/movie_poster/1_ref.png" width="200"></td>
<td><img src="./assets/samples/portrait/human_1.jpg" width="200"></td>
<td><img src="./assets/samples/application/movie_poster/1_2_m.webp" width="200"></td>
<td><img src="./assets/samples/application/movie_poster/1_1_fft.webp" width="200"></td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"{image}, the man faces the camera."</td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"Face swap"</td>
</tr>
<tr>
<td></td>
<td><img src="./assets/samples/application/sr/sr_tiger.png" width="200"></td>
<td><img src="./assets/samples/application/sr/sr_tiger_m.webp" width="200"></td>
<td><img src="./assets/samples/application/sr/sr_tiger_fft.webp" width="200"></td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"{image} features a close-up of a young, furry tiger cub on a rock. The tiger, which appears to be quite young, has distinctive orange, black, and white striped fur, typical of tigers. The cub's eyes have a bright and curious expression, and its ears are perked up, indicating alertness. The cub seems to be in the act of climbing or resting on the rock. The background is a blurred grassland with trees, but the focus is on the cub, which is vividly colored while the rest of the image is in grayscale, drawing attention to the tiger's details. The photo captures a moment in the wild, depicting the charming and tenacious nature of this young tiger, as well as its typical interaction with the environment."</td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"Super-resolution"</td>
</tr>
<tr>
<td></td>
<td><img src="./assets/samples/application/photo_editing/1_ref.png" width="200"></td>
<td><img src="./assets/samples/application/photo_editing/1_1_orm.webp" width="200"></td>
<td><img src="./assets/samples/application/regional_editing/1_1_fft.webp" width="200"></td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"a blue hand"</td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"Regional Editing"</td>
</tr>
<tr>
<td></td>
<td><img src="./assets/samples/application/photo_editing/1_ref.png" width="200"></td>
<td><img src="./assets/samples/application/photo_editing/1_1_rm.webp" width="200"></td>
<td><img src="./assets/samples/application/regional_editing/1_2_fft.webp" width="200"></td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"Mechanical hands like a robot"</td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"Regional Editing"</td>
</tr>
<tr>
<td></td>
<td><img src="./assets/samples/control/1_1_recolor.webp" width="200"></td>
<td><img src="./assets/samples/control/1_1_m.webp" width="200"></td>
<td><img src="./assets/samples/control/1_1_fft_recolor.webp" width="200"></td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"{image} Beautiful female portrait, Robot with smooth White transparent carbon shell, rococo detailing, Natural lighting, Highly detailed, Cinematic, 4K."</td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"Recolorizing"</td>
</tr>
<tr>
<td></td>
<td><img src="./assets/samples/control/1_1_depth.webp" width="200"></td>
<td><img src="./assets/samples/control/1_1_m.webp" width="200"></td>
<td><img src="./assets/samples/control/1_1_fft_depth.webp" width="200"></td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"{image} Beautiful female portrait, Robot with smooth White transparent carbon shell, rococo detailing, Natural lighting, Highly detailed, Cinematic, 4K."</td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"Depth Guided Generation"</td>
</tr>
<tr>
<td></td>
<td><img src="./assets/samples/control/1_1_contourc.webp" width="200"></td>
<td><img src="./assets/samples/control/1_1_m.webp" width="200"></td>
<td><img src="./assets/samples/control/1_1_fft_contour.webp" width="200"></td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"{image} Beautiful female portrait, Robot with smooth White transparent carbon shell, rococo detailing, Natural lighting, Highly detailed, Cinematic, 4K."</td>
<td style="word-wrap:break-word;word-break:break-all;" width="250px";>"Contour Guided Generation"</td>
</tr>
</tbody>
</table>
## Comfyui Workflows in community
We are deeply grateful to the community developers for building many fascinating applications based on the ACE++ series of models.
During this process, we have received valuable feedback, particularly regarding artifacts in generated images and the stability of the results.
In response to these issues, many developers have proposed creative solutions, which have greatly inspired us, and we pay tribute to them.
@@ -230,9 +349,6 @@ Models' scepter_path:
- **ModelScope:** ms://iic/ACE_Plus@local_editing/xxxx.safetensors
- **HuggingFace:** hf://ali-vilab/ACE_Plus@local_editing/xxxx.safetensors
### ACE++ Fully [Coming soon]
Fully finetuning a composite model with ACE’s data to support various editing and reference generation tasks through an instructive approach.
## 🔥 Applications
The ACE++ model supports a wide range of downstream tasks through simple adaptations. Here are some examples, and we look forward to seeing the community explore even more exciting applications utilizing the ACE++ model.
@@ -302,7 +418,7 @@ For model preparation, we provide three methods for downloading the model. The s
## 🚀 Inference
Under the condition that the environment variables defined in [Installation](#-installation), users can run examples and test your own samples by executing infer.py.
The relevant commands are as follows:
The relevant commands for lora models are as follows:
```bash
export FLUX_FILL_PATH="hf://black-forest-labs/FLUX.1-Fill-dev"
export PORTRAIT_MODEL_PATH="ms://iic/ACE_Plus@portrait/comfyui_portrait_lora64.safetensors"
@@ -312,7 +428,13 @@ export LOCAL_MODEL_PATH="ms://iic/ACE_Plus@local_editing/comfyui_local_lora16.sa
# export PORTRAIT_MODEL_PATH="hf://ali-vilab/ACE_Plus@portrait/comfyui_portrait_lora64.safetensors"
# export SUBJECT_MODEL_PATH="hf://ali-vilab/ACE_Plus@subject/comfyui_subject_lora16.safetensors"
# export LOCAL_MODEL_PATH="hf://ali-vilab/ACE_Plus@local_editing/comfyui_local_lora16.safetensors"
python infer.py
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"
python infer_fft.py
```
## 🚀 Train
@@ -332,6 +454,10 @@ All parameters related to training are stored in 'train_config/ace_plus_lora.yam
```bash
export FLUX_FILL_PATH="hf://black-forest-labs/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"
python run_train.py --cfg train_config/ace_plus_fft.yaml
```
The models trained by ACE++ can be found in ./examples/exp_example/xxxx/checkpoints/xxxx/0_SwiftLoRA/comfyui_model.safetensors.
@@ -348,7 +474,11 @@ export LOCAL_MODEL_PATH="ms://iic/ACE_Plus@local_editing/comfyui_local_lora16.sa
# export PORTRAIT_MODEL_PATH="hf://ali-vilab/ACE_Plus@portrait/comfyui_portrait_lora64.safetensors"
# export SUBJECT_MODEL_PATH="hf://ali-vilab/ACE_Plus@subject/comfyui_subject_lora16.safetensors"
# export LOCAL_MODEL_PATH="hf://ali-vilab/ACE_Plus@local_editing/comfyui_local_lora16.safetensors"
python demo.py
python demo_lora.py
# Use the 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"
python demo_fft.py
```
## 📚 Limitations
-524
View File
@@ -1,524 +0,0 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import argparse
import csv
import glob
import os
import sys
import threading
import time
import gradio as gr
import numpy as np
import torch, importlib
from PIL import Image
from scepter.modules.transform.io import pillow_convert
from scepter.modules.utils.config import Config
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
if os.path.exists('__init__.py'):
package_name = 'scepter_ext'
spec = importlib.util.spec_from_file_location(package_name, '__init__.py')
package = importlib.util.module_from_spec(spec)
sys.modules[package_name] = package
spec.loader.exec_module(package)
from inference.ace_plus_diffusers import ACEPlusDiffuserInference
from inference.utils import edit_preprocess
from examples.examples import all_examples
inference_dict = {
"ACE_DIFFUSER_PLUS": ACEPlusDiffuserInference
}
fs_list = [
Config(cfg_dict={"NAME": "HuggingfaceFs", "TEMP_DIR": "./cache"}, load=False),
Config(cfg_dict={"NAME": "ModelscopeFs", "TEMP_DIR": "./cache"}, load=False),
Config(cfg_dict={"NAME": "HttpFs", "TEMP_DIR": "./cache"}, load=False),
Config(cfg_dict={"NAME": "LocalFs", "TEMP_DIR": "./cache"}, load=False),
]
for one_fs in fs_list:
FS.init_fs_client(one_fs)
csv.field_size_limit(sys.maxsize)
refresh_sty = '\U0001f504' # 🔄
clear_sty = '\U0001f5d1' # 🗑️
upload_sty = '\U0001f5bc' # 🖼️
sync_sty = '\U0001f4be' # 💾
chat_sty = '\U0001F4AC' # 💬
video_sty = '\U0001f3a5' # 🎥
lock = threading.Lock()
class DemoUI(object):
def __init__(self,
infer_dir = "./config",
model_list='./models/model_zoo.yaml'
):
self.model_yamls = glob.glob(os.path.join(infer_dir,
'*.yaml'))
self.model_choices = dict()
self.default_model_name = ''
for i in self.model_yamls:
model_cfg = Config(load=True, cfg_file=i)
model_name = model_cfg.NAME
if model_cfg.IS_DEFAULT: self.default_model_name = model_name
self.model_choices[model_name] = model_cfg
print('Models: ', self.model_choices.keys())
assert len(self.model_choices) > 0
if self.default_model_name == "": self.default_model_name = list(self.model_choices.keys())[0]
self.model_name = self.default_model_name
pipe_cfg = self.model_choices[self.default_model_name]
infer_name = pipe_cfg.get("INFERENCE_TYPE", "ACE")
self.pipe = inference_dict[infer_name]()
self.pipe.init_from_cfg(pipe_cfg)
# choose different model
self.task_model_cfg = Config(load=True, cfg_file=model_list)
self.task_model = {}
self.task_model_list = []
self.edit_type_dict = {"repainting": None}
self.edit_type_list = ["repainting"]
for task_name, task_model in self.task_model_cfg.MODEL.items():
self.task_model[task_name.lower()] = task_model
self.task_model_list.append(task_name.lower())
for preprocessor in task_model.get("PREPROCESSOR", []):
if preprocessor["TYPE"] in self.edit_type_dict:
continue
preprocessor["REPAINTING_SCALE"] = task_model.get("REPAINTING_SCALE", 1.0)
self.edit_type_dict[preprocessor["TYPE"]] = preprocessor
self.max_msgs = 20
# reformat examples
self.all_examples = [
[
one_example["task_type"], one_example["edit_type"], one_example["instruction"],
one_example["input_reference_image"], one_example["input_image"],
one_example["input_mask"], one_example["output_h"],
one_example["output_w"], one_example["seed"]
]
for one_example in all_examples
]
def construct_edit_image(self, edit_image, edit_mask):
if edit_image is not None and edit_mask is not None:
edit_image_rgb = pillow_convert(edit_image, "RGB")
edit_image_rgba = pillow_convert(edit_image, "RGBA")
edit_mask = pillow_convert(edit_mask, "L")
arr1 = np.array(edit_image_rgb)
arr2 = np.array(edit_mask)[:, :, np.newaxis]
result_array = np.concatenate((arr1, arr2), axis=2)
layer = Image.fromarray(result_array)
ret_data = {
"background": edit_image_rgba,
"composite": edit_image_rgba,
"layers": [layer]
}
return ret_data
else:
return None
def create_ui(self):
with gr.Row(equal_height=True, visible=True):
with gr.Column(scale=2):
self.gallery_image = gr.Image(
height=600,
interactive=False,
type='pil',
elem_id='Reference_image'
)
with gr.Column(scale=1, visible=True) as self.edit_preprocess_panel:
with gr.Row():
with gr.Accordion(label='Related Input Image', open=False):
self.edit_preprocess_preview = gr.Image(
height=600,
interactive=False,
type='pil',
elem_id='preprocess_image'
)
self.edit_preprocess_mask_preview = gr.Image(
height=600,
interactive=False,
type='pil',
elem_id='preprocess_image_mask'
)
with gr.Row():
instruction = """
**Instruction**:
1. Please choose the Task Type based on the scenario of the generation task. We provide three types of generation capabilities: Portrait ID Preservation Generation(portrait),
Object ID Preservation Generation(subject), and Local Controlled Generation(local editing), which can be selected from the task dropdown menu.
2. When uploading images in the Reference Image section, the generated image will reference the ID information of that image. Please ensure that the ID information is clear.
In the Edit Image section, the uploaded image will maintain its structural and content information, and you must draw a mask area to specify the region to be regenerated.
3. When the task type is local editing, there are various editing types to choose from. Users can select different information preserving dimensions, such as edge information,
color information, and more. The pre-processing information can be viewed in the 'related input image' tab.
"""
self.instruction = gr.Markdown(value=instruction)
with gr.Row():
self.model_name_dd = gr.Dropdown(
choices=self.model_choices,
value=self.default_model_name,
label='Model Version')
self.task_type = gr.Dropdown(choices=self.task_model_list,
interactive=True,
value=self.task_model_list[0],
label='Task Type')
self.edit_type = gr.Dropdown(choices=self.edit_type_list,
interactive=True,
value=self.edit_type_list[0],
label='Edit Type')
with gr.Row():
self.generation_info_preview = gr.Markdown(
label='System Log.',
show_label=True)
with gr.Row(variant='panel',
equal_height=True,
show_progress=False):
with gr.Column(scale=10, min_width=500):
self.text = gr.Textbox(
placeholder='Input "@" find history of image',
label='Instruction',
container=False,
lines = 1)
with gr.Column(scale=2, min_width=100):
with gr.Row():
with gr.Column(scale=1, min_width=100):
self.chat_btn = gr.Button(value='Generate', variant = "primary")
with gr.Accordion(label='Advance', open=True):
with gr.Row(visible=True):
with gr.Column():
self.reference_image = gr.Image(
height=1000,
interactive=True,
image_mode='RGB',
type='pil',
label='Reference Image',
elem_id='reference_image'
)
with gr.Column():
self.edit_image = gr.ImageMask(
height=1000,
interactive=True,
value=None,
sources=['upload'],
type='pil',
layers=False,
label='Edit Image',
elem_id='image_editor',
show_fullscreen_button=True,
format="png"
)
with gr.Row():
self.step = gr.Slider(minimum=1,
maximum=1000,
value=self.pipe.input.get("sample_steps", 20),
visible=self.pipe.input.get("sample_steps", None) is not None,
label='Sample Step')
self.cfg_scale = gr.Slider(
minimum=1.0,
maximum=100.0,
value=self.pipe.input.get("guide_scale", 4.5),
visible=self.pipe.input.get("guide_scale", None) is not None,
label='Guidance Scale')
self.seed = gr.Slider(minimum=-1,
maximum=10000000,
value=-1,
label='Seed')
self.output_height = gr.Slider(
minimum=256,
maximum=1440,
value=self.pipe.input.get("output_height", 1024),
visible=self.pipe.input.get("output_height", None) is not None,
label='Output Height')
self.output_width = gr.Slider(
minimum=256,
maximum=1440,
value=self.pipe.input.get("output_width", 1024),
visible=self.pipe.input.get("output_width", None) is not None,
label='Output Width')
self.repainting_scale = gr.Slider(
minimum=0.0,
maximum=1.0,
value=self.pipe.input.get("repainting_scale", 1.0),
visible=True,
label='Repainting Scale')
with gr.Row():
self.eg = gr.Column(visible=True)
def set_callbacks(self, *args, **kwargs):
########################################
def change_model(model_name):
if model_name not in self.model_choices:
gr.Info('The provided model name is not a valid choice!')
return model_name, gr.update(), gr.update()
if model_name != self.model_name:
lock.acquire()
del self.pipe
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
pipe_cfg = self.model_choices[model_name]
infer_name = pipe_cfg.get("INFERENCE_TYPE", "ACE")
self.pipe = inference_dict[infer_name]()
self.pipe.init_from_cfg(pipe_cfg)
self.model_name = model_name
lock.release()
return (model_name, gr.update(),
gr.Slider(
value=self.pipe.input.get("sample_steps", 20),
visible=self.pipe.input.get("sample_steps", None) is not None),
gr.Slider(
value=self.pipe.input.get("guide_scale", 4.5),
visible=self.pipe.input.get("guide_scale", None) is not None),
gr.Slider(
value=self.pipe.input.get("output_height", 1024),
visible=self.pipe.input.get("output_height", None) is not None),
gr.Slider(
value=self.pipe.input.get("output_width", 1024),
visible=self.pipe.input.get("output_width", None) is not None),
gr.Slider(value=self.pipe.input.get("repainting_scale", 1.0))
)
self.model_name_dd.change(
change_model,
inputs=[self.model_name_dd],
outputs=[
self.model_name_dd, self.text,
self.step,
self.cfg_scale,
self.output_height,
self.output_width,
self.repainting_scale])
def change_task_type(task_type):
task_info = self.task_model[task_type]
edit_type_list = [self.edit_type_list[0]]
for preprocessor in task_info.get("PREPROCESSOR", []):
preprocessor["REPAINTING_SCALE"] = task_info.get("REPAINTING_SCALE", 1.0)
self.edit_type_dict[preprocessor["TYPE"]] = preprocessor
edit_type_list.append(preprocessor["TYPE"])
return gr.update(choices=edit_type_list, value=edit_type_list[0])
self.task_type.change(change_task_type, inputs=[self.task_type], outputs=[self.edit_type])
def change_edit_type(edit_type):
edit_info = self.edit_type_dict[edit_type]
edit_info = edit_info or {}
repainting_scale = edit_info.get("REPAINTING_SCALE", 1.0)
if edit_type == self.edit_type_list[0]:
return gr.Slider(value=1.0)
else:
return gr.Slider(
value=repainting_scale)
self.edit_type.change(change_edit_type, inputs=[self.edit_type], outputs=[self.repainting_scale])
def preprocess_input(ref_image, edit_image_dict, preprocess = None):
err_msg = ""
is_suc = True
if ref_image is not None:
ref_image = pillow_convert(ref_image, "RGB")
if edit_image_dict is None:
edit_image = None
edit_mask = None
else:
edit_image = edit_image_dict["background"]
edit_mask = np.array(edit_image_dict["layers"][0])[:, :, 3]
if np.sum(np.array(edit_image)) < 1:
edit_image = None
edit_mask = None
elif np.sum(np.array(edit_mask)) < 1:
err_msg = "You must draw the repainting area for the edited image."
return None, None, None, False, err_msg
else:
edit_image = pillow_convert(edit_image, "RGB")
edit_mask = Image.fromarray(edit_mask).convert('L')
if ref_image is None and edit_image is None:
err_msg = "Please provide the reference image or edited image."
return None, None, None, False, err_msg
return edit_image, edit_mask, ref_image, is_suc, err_msg
def run_chat(
prompt,
ref_image,
edit_image,
task_type,
edit_type,
cfg_scale,
step,
seed,
output_h,
output_w,
repainting_scale,
progress=gr.Progress(track_tqdm=True)
):
model_path = self.task_model[task_type]["MODEL_PATH"]
edit_info = self.edit_type_dict[edit_type]
if task_type in ["portrait", "subject"] and ref_image is None:
err_msg = "<mark>Please provide the reference image.</mark>"
return (gr.Image(), gr.Column(visible=True),
gr.Image(),
gr.Image(),
gr.Text(value=err_msg))
pre_edit_image, pre_edit_mask, pre_ref_image, is_suc, err_msg = preprocess_input(ref_image, edit_image)
if not is_suc:
err_msg = f"<mark>{err_msg}</mark>"
return (gr.Image(), gr.Column(visible=True),
gr.Image(),
gr.Image(),
gr.Text(value=err_msg))
pre_edit_image = edit_preprocess(edit_info, we.device_id, pre_edit_image, pre_edit_mask)
# edit_image["background"] = pre_edit_image
st = time.time()
image, seed = self.pipe(
reference_image=pre_ref_image,
edit_image=pre_edit_image,
edit_mask=pre_edit_mask,
prompt=prompt,
output_height=output_h,
output_width=output_w,
sampler='flow_euler',
sample_steps=step,
guide_scale=cfg_scale,
seed=seed,
repainting_scale=repainting_scale,
lora_path = model_path
)
et = time.time()
msg = f"prompt: {prompt}; seed: {seed}; cost time: {et - st}s; repaiting scale: {repainting_scale}"
return (gr.Image(value=image), gr.Column(visible=True),
gr.Image(value=pre_edit_image if pre_edit_image is not None else pre_ref_image),
gr.Image(value=pre_edit_mask if pre_edit_mask is not None else None),
gr.Text(value=msg))
chat_inputs = [
self.reference_image,
self.edit_image,
self.task_type,
self.edit_type,
self.cfg_scale,
self.step,
self.seed,
self.output_height,
self.output_width,
self.repainting_scale
]
chat_outputs = [
self.gallery_image, self.edit_preprocess_panel, self.edit_preprocess_preview,
self.edit_preprocess_mask_preview, self.generation_info_preview
]
self.chat_btn.click(run_chat,
inputs=[self.text] + chat_inputs,
outputs=chat_outputs,
queue=True)
self.text.submit(run_chat,
inputs=[self.text] + chat_inputs,
outputs=chat_outputs,
queue=True)
def run_example(task_type, edit_type, prompt, ref_image, edit_image, edit_mask,
output_h, output_w, seed, progress=gr.Progress(track_tqdm=True)):
model_path = self.task_model[task_type]["MODEL_PATH"]
step = self.pipe.input.get("sample_steps", 20)
cfg_scale = self.pipe.input.get("guide_scale", 20)
edit_info = self.edit_type_dict[edit_type]
edit_image = self.construct_edit_image(edit_image, edit_mask)
pre_edit_image, pre_edit_mask, pre_ref_image, _, _ = preprocess_input(ref_image, edit_image)
pre_edit_image = edit_preprocess(edit_info, we.device_id, pre_edit_image, pre_edit_mask)
edit_info = edit_info or {}
repainting_scale = edit_info.get("REPAINTING_SCALE", 1.0)
st = time.time()
image, seed = self.pipe(
reference_image=pre_ref_image,
edit_image=pre_edit_image,
edit_mask=pre_edit_mask,
prompt=prompt,
output_height=output_h,
output_width=output_w,
sampler='flow_euler',
sample_steps=step,
guide_scale=cfg_scale,
seed=seed,
repainting_scale=repainting_scale,
lora_path=model_path
)
et = time.time()
msg = f"prompt: {prompt}; seed: {seed}; cost time: {et - st}s; repaiting scale: {repainting_scale}"
if pre_edit_image is not None:
ret_image = Image.composite(Image.new("RGB", pre_edit_image.size, (0, 0, 0)), pre_edit_image, pre_edit_mask)
else:
ret_image = None
return (gr.Image(value=image), gr.Column(visible=True),
gr.Image(value=pre_edit_image if pre_edit_image is not None else pre_ref_image),
gr.Image(value=pre_edit_mask if pre_edit_mask is not None else None),
gr.Text(value=msg),
gr.update(value=ret_image))
with self.eg:
self.example_edit_image = gr.Image(label='Edit Image',
type='pil',
image_mode='RGB',
visible=False)
self.example_edit_mask = gr.Image(label='Edit Image Mask',
type='pil',
image_mode='L',
visible=False)
self.examples = gr.Examples(
fn=run_example,
examples=self.all_examples,
inputs=[
self.task_type, self.edit_type, self.text, self.reference_image, self.example_edit_image,
self.example_edit_mask, self.output_height, self.output_width, self.seed
],
outputs=[self.gallery_image, self.edit_preprocess_panel, self.edit_preprocess_preview,
self.edit_preprocess_mask_preview, self.generation_info_preview, self.edit_image],
examples_per_page=6,
cache_examples=False,
run_on_click=True)
def run_gr(cfg):
with gr.Blocks() as demo:
chatbot = DemoUI()
chatbot.create_ui()
chatbot.set_callbacks()
demo.launch(server_name='0.0.0.0',
server_port=cfg.args.server_port,
root_path=cfg.args.root_path)
if __name__ == '__main__':
parser = argparse.ArgumentParser(description='Argparser for Scepter:\n')
parser.add_argument('--server_port',
dest='server_port',
help='',
type=int,
default=2345)
parser.add_argument('--root_path', dest='root_path', help='', default='')
cfg = Config(load=True, parser_ins=parser)
run_gr(cfg)
+152
View File
@@ -78,4 +78,156 @@ all_examples = [
"edit_type": "repainting"
}
]
fft_examples = [
{
"input_image": None,
"input_mask": None,
"input_reference_image": "./assets/samples/portrait/human_1.jpg",
"save_path": "examples/outputs/portrait_human_1.jpg",
"instruction": "Maintain the facial features, A girl is wearing a neat police uniform and sporting a badge. She is smiling with a friendly and confident demeanor. The background is blurred, featuring a cartoon logo.",
"output_h": 1024,
"output_w": 1024,
"seed": 10000000,
"repainting_scale": 1.0,
"edit_type": "repainting"
},
{
"input_image": None,
"input_mask": None,
"input_reference_image": "./assets/samples/subject/subject_1.jpg",
"save_path": "examples/outputs/subject_subject_1.jpg",
"instruction": "Display the logo in a minimalist style printed in white on a matte black ceramic coffee mug, alongside a steaming cup of coffee on a cozy cafe table.",
"output_h": 1024,
"output_w": 1024,
"seed": 10000000,
"repainting_scale": 1.0,
"edit_type": "repainting"
},
{
"input_image": "./assets/samples/application/photo_editing/1_2_edit.jpg",
"input_mask": "./assets/samples/application/photo_editing/1_2_m.webp",
"input_reference_image": "./assets/samples/application/photo_editing/1_ref.png",
"save_path": "examples/outputs/photo_editing_1.jpg",
"instruction": "The item is put on the table.",
"output_h": 1024,
"output_w": 1024,
"seed": 8006019,
"repainting_scale": 1.0,
"edit_type": "repainting"
},
{
"input_image": "./assets/samples/application/logo_paste/1_1_edit.png",
"input_mask": "./assets/samples/application/logo_paste/1_1_m.png",
"input_reference_image": "assets/samples/application/logo_paste/1_ref.png",
"save_path": "examples/outputs/logo_paste_1.jpg",
"instruction": "The logo is printed on the headphones.",
"output_h": 1024,
"output_w": 1024,
"seed": 934582264,
"repainting_scale": 1.0,
"edit_type": "repainting"
},
{
"input_image": "./assets/samples/application/try_on/1_1_edit.png",
"input_mask": "./assets/samples/application/try_on/1_1_m.png",
"input_reference_image": "assets/samples/application/try_on/1_ref.png",
"save_path": "examples/outputs/try_on_1.jpg",
"instruction": "The woman dresses this skirt.",
"output_h": 1024,
"output_w": 1024,
"seed": 934582264,
"repainting_scale": 1.0,
"edit_type": "repainting"
},
{
"input_image": "./assets/samples/portrait/human_1.jpg",
"input_mask": "assets/samples/application/movie_poster/1_2_m.webp",
"input_reference_image": "assets/samples/application/movie_poster/1_ref.png",
"save_path": "examples/outputs/movie_poster_1.jpg",
"instruction": "{image}, the man faces the camera.",
"output_h": 1024,
"output_w": 1024,
"seed": 3999647,
"repainting_scale": 1.0,
"edit_type": "repainting"
},
{
"input_image": "./assets/samples/application/sr/sr_tiger.png",
"input_mask": "./assets/samples/application/sr/sr_tiger_m.webp",
"input_reference_image": None,
"save_path": "examples/outputs/mario_recolorizing_1.jpg",
"instruction": "{image} features a close-up of a young, furry tiger cub on a rock. The tiger, which appears to be quite young, has distinctive orange, "
"black, and white striped fur, typical of tigers. The cub's eyes have a bright and curious expression, and its ears are perked up, "
"indicating alertness. The cub seems to be in the act of climbing or resting on the rock. The background is a blurred grassland with trees, "
"but the focus is on the cub, which is vividly colored while the rest of the image is in grayscale, drawing attention to the tiger's details."
" The photo captures a moment in the wild, depicting the charming and tenacious nature of this young tiger,"
" as well as its typical interaction with the environment.",
"output_h": 1024,
"output_w": 1024,
"seed": 199999,
"repainting_scale": 0.0,
"edit_type": "no_preprocess"
},
{
"input_image": "./assets/samples/application/photo_editing/1_ref.png",
"input_mask": "./assets/samples/application/photo_editing/1_1_orm.webp",
"input_reference_image": None,
"save_path": "examples/outputs/mario_repainting_1.jpg",
"instruction": "a blue hand",
"output_h": 1024,
"output_w": 1024,
"seed": 63401,
"repainting_scale": 1.0,
"edit_type": "repainting"
},
{
"input_image": "./assets/samples/application/photo_editing/1_ref.png",
"input_mask": "./assets/samples/application/photo_editing/1_1_rm.webp",
"input_reference_image": None,
"save_path": "examples/outputs/mario_repainting_2.jpg",
"instruction": "Mechanical hands like a robot",
"output_h": 1024,
"output_w": 1024,
"seed": 59107,
"repainting_scale": 1.0,
"edit_type": "repainting"
},
{
"input_image": "./assets/samples/control/1_1.webp",
"input_mask": "./assets/samples/control/1_1_m.webp",
"input_reference_image": None,
"save_path": "examples/outputs/control_recolorizing.jpg",
"instruction": "{image} Beautiful female portrait, Robot with smooth White transparent carbon shell, rococo detailing, Natural lighting, Highly detailed, Cinematic, 4K.",
"output_h": 1024,
"output_w": 1024,
"seed": 9652101,
"repainting_scale": 0.0,
"edit_type": "recolorizing"
},
{
"input_image": "./assets/samples/control/1_1.webp",
"input_mask": "./assets/samples/control/1_1_m.webp",
"input_reference_image": None,
"save_path": "examples/outputs/control_depth.jpg",
"instruction": "{image} Beautiful female portrait, Robot with smooth White transparent carbon shell, rococo detailing, Natural lighting, Highly detailed, Cinematic, 4K.",
"output_h": 1024,
"output_w": 1024,
"seed": 14979476,
"repainting_scale": 0.0,
"edit_type": "depth_repainting"
},
{
"input_image": "./assets/samples/control/1_1.webp",
"input_mask": "./assets/samples/control/1_1_m.webp",
"input_reference_image": None,
"save_path": "examples/outputs/control_contour.jpg",
"instruction": "{image} Beautiful female portrait, Robot with smooth White transparent carbon shell, rococo detailing, Natural lighting, Highly detailed, Cinematic, 4K.",
"output_h": 1024,
"output_w": 1024,
"seed": 4227292472,
"repainting_scale": 0.0,
"edit_type": "contour_repainting"
}
]
-228
View File
@@ -1,228 +0,0 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import argparse
import glob
import io
import os
from PIL import Image
from scepter.modules.transform.io import pillow_convert
from scepter.modules.utils.config import Config
from scepter.modules.utils.file_system import FS
from examples.examples import all_examples
from inference.ace_plus_diffusers import ACEPlusDiffuserInference
inference_dict = {
"ACE_DIFFUSER_PLUS": ACEPlusDiffuserInference
}
fs_list = [
Config(cfg_dict={"NAME": "HuggingfaceFs", "TEMP_DIR": "./cache"}, load=False),
Config(cfg_dict={"NAME": "ModelscopeFs", "TEMP_DIR": "./cache"}, load=False),
Config(cfg_dict={"NAME": "HttpFs", "TEMP_DIR": "./cache"}, load=False),
Config(cfg_dict={"NAME": "LocalFs", "TEMP_DIR": "./cache"}, load=False),
]
for one_fs in fs_list:
FS.init_fs_client(one_fs)
def run_one_case(pipe,
input_image = None,
input_mask = None,
input_reference_image = None,
save_path = "examples/output/example.png",
instruction = "",
output_h = 1024,
output_w = 1024,
seed = -1,
sample_steps = None,
guide_scale = None,
repainting_scale = None,
model_path = None,
**kwargs):
if input_image is not None:
input_image = Image.open(io.BytesIO(FS.get_object(input_image)))
input_image = pillow_convert(input_image, "RGB")
if input_mask is not None:
input_mask = Image.open(io.BytesIO(FS.get_object(input_mask)))
input_mask = pillow_convert(input_mask, "L")
if input_reference_image is not None:
input_reference_image = Image.open(io.BytesIO(FS.get_object(input_reference_image)))
input_reference_image = pillow_convert(input_reference_image, "RGB")
image, seed = pipe(
reference_image=input_reference_image,
edit_image=input_image,
edit_mask=input_mask,
prompt=instruction,
output_height=output_h,
output_width=output_w,
sampler='flow_euler',
sample_steps=sample_steps or pipe.input.get("sample_steps", 28),
guide_scale=guide_scale or pipe.input.get("guide_scale", 50),
seed=seed,
repainting_scale=repainting_scale or pipe.input.get("repainting_scale", 1.0),
lora_path = model_path
)
with FS.put_to(save_path) as local_path:
image.save(local_path)
return local_path, seed
def run():
parser = argparse.ArgumentParser(description='Argparser for Scepter:\n')
parser.add_argument('--instruction',
dest='instruction',
help='The instruction for editing or generating!',
default="")
parser.add_argument('--output_h',
dest='output_h',
help='The height of output image for generation tasks!',
type=int,
default=1024)
parser.add_argument('--output_w',
dest='output_w',
help='The width of output image for generation tasks!',
type=int,
default=1024)
parser.add_argument('--input_reference_image',
dest='input_reference_image',
help='The input reference image!',
default=None
)
parser.add_argument('--input_image',
dest='input_image',
help='The input image!',
default=None
)
parser.add_argument('--input_mask',
dest='input_mask',
help='The input mask!',
default=None
)
parser.add_argument('--save_path',
dest='save_path',
help='The save path for output image!',
default='examples/output_images/output.png'
)
parser.add_argument('--seed',
dest='seed',
help='The seed for generation!',
type=int,
default=-1)
parser.add_argument('--step',
dest='step',
help='The sample step for generation!',
type=int,
default=None)
parser.add_argument('--guide_scale',
dest='guide_scale',
help='The guide scale for generation!',
type=int,
default=None)
parser.add_argument('--repainting_scale',
dest='repainting_scale',
help='The repainting scale for content filling generation!',
type=int,
default=None)
parser.add_argument('--task_type',
dest='task_type',
choices=['portrait', 'subject', 'local_editing'],
help="Choose the task type.",
default='')
parser.add_argument('--task_model',
dest='task_model',
help='The models list for different tasks!',
default="./models/model_zoo.yaml")
parser.add_argument('--infer_type',
dest='infer_type',
choices=['diffusers'],
default='diffusers',
help="Choose the inference scripts. 'native' refers to using the official implementation of ace++, "
"while 'diffusers' refers to using the adaptation for diffusers")
parser.add_argument('--cfg_folder',
dest='cfg_folder',
help='The inference config!',
default="./config")
cfg = Config(load=True, parser_ins=parser)
model_yamls = glob.glob(os.path.join(cfg.args.cfg_folder, '*.yaml'))
model_choices = dict()
for i in model_yamls:
model_cfg = Config(load=True, cfg_file=i)
model_name = model_cfg.NAME
model_choices[model_name] = model_cfg
if cfg.args.infer_type == "native":
infer_name = "ace_plus_native_infer"
elif cfg.args.infer_type == "diffusers":
infer_name = "ace_plus_diffuser_infer"
else:
raise ValueError("infer_type should be native or diffusers")
assert infer_name in model_choices
# choose different model
task_model_cfg = Config(load=True, cfg_file=cfg.args.task_model)
task_model_dict = {}
for task_name, task_model in task_model_cfg.MODEL.items():
task_model_dict[task_name] = task_model
# choose the inference scripts.
pipe_cfg = model_choices[infer_name]
infer_name = pipe_cfg.get("INFERENCE_TYPE", "ACE_PLUS")
pipe = inference_dict[infer_name]()
pipe.init_from_cfg(pipe_cfg)
if cfg.args.instruction == "" and cfg.args.input_image is None and cfg.args.input_reference_image is None:
params = {
"output_h": cfg.args.output_h,
"output_w": cfg.args.output_w,
"sample_steps": cfg.args.step,
"guide_scale": cfg.args.guide_scale
}
# run examples
for example in all_examples:
example["model_path"] = FS.get_from(task_model_dict[example["task_type"].upper()]["MODEL_PATH"])
example.update(params)
if example["edit_type"] == "repainting":
example["repainting_scale"] = 1.0
else:
example["repainting_scale"] = task_model_dict[example["task_type"].upper()].get("REPAINTING_SCALE", 1.0)
print(example)
local_path, seed = run_one_case(pipe, **example)
else:
assert cfg.args.task_type.upper() in task_model_cfg
params = {
"input_image": cfg.args.input_image,
"input_mask": cfg.args.input_mask,
"input_reference_image": cfg.args.input_reference_image,
"save_path": cfg.args.save_path,
"instruction": cfg.args.instruction,
"output_h": cfg.args.output_h,
"output_w": cfg.args.output_w,
"sample_steps": cfg.args.step,
"guide_scale": cfg.args.guide_scale,
"repainting_scale": cfg.args.repainting_scale,
"model_path": FS.get_from(task_model_dict[cfg.args.task_type.upper()]["MODEL_PATH"])
}
local_path, seed = run_one_case(pipe, **params)
print(local_path, seed)
if __name__ == '__main__':
run()
+2 -2
View File
@@ -1,6 +1,6 @@
from .flux import FluxMRACEPlus
from .flux import FluxMRACEPlus, FluxMRModiACEPlus
from .ace_plus_dataset import ACEPlusDataset
from .ace_plus_ldm import LatentDiffusionACEPlus
from .ace_plus_solver import ACEPlusSolver
from .ace_plus_solver import FormalACEPlusSolver
from .embedder import ACEHFEmbedder, T5ACEPlusClipFluxEmbedder
from .checkpoint import ACECheckpointHook, ACEBackwardHook
+52 -14
View File
@@ -38,6 +38,30 @@ def ensure_src_align_target_h_mode(src_image, size, image_id, interpolation=Inte
ret_image.append(T.Resize((tH, tW), interpolation=interpolation, antialias=True)(edit_image))
return ret_image
def ensure_src_align_target_padding_mode(src_image, size, image_id, size_h = [], interpolation=InterpolationMode.BILINEAR):
# padding mode
H, W = size
ret_data = []
ret_h = []
for idx, one_id in enumerate(image_id):
if len(size_h) < 1:
rH = random.randint(int(H / 3), int(H))
else:
rH = size_h[idx]
ret_h.append(rH)
edit_image = src_image[one_id]
_, eH, eW = edit_image.shape
scale = rH/eH
tH, tW = rH, int(eW * scale)
edit_image = T.Resize((tH, tW), interpolation=interpolation, antialias=True)(edit_image)
# padding
delta_w = 0
delta_h = H - tH
padding = (delta_w // 2, delta_h // 2, delta_w - (delta_w // 2), delta_h - (delta_h // 2))
ret_data.append(T.Pad(padding, fill=0, padding_mode="constant")(edit_image).float())
return ret_data, ret_h
def ensure_limit_sequence(image, max_seq_len = 4096, d = 16, interpolation=InterpolationMode.BILINEAR):
# resize image for max_seq_len, while keep the aspect ratio
H, W = image.shape[-2:]
@@ -83,6 +107,7 @@ class ACEPlusDataset(BaseDataset):
fields = cfg.get("FIELDS", [])
prefix = cfg.get("PATH_PREFIX", "")
edit_type_list = cfg.get("EDIT_TYPE_LIST", [])
self.modify_mode = cfg.get("MODIFY_MODE", True)
self.max_seq_len = cfg.get("MAX_SEQ_LEN", 4096)
self.repaiting_scale = cfg.get("REPAINTING_SCALE", 0.5)
self.d = cfg.get("D", 16)
@@ -135,6 +160,7 @@ class ACEPlusDataset(BaseDataset):
def _get(self, index):
# normalize
sample_id = index%len(self)
index = self.items[index%len(self)]
prefix = index.get("prefix", "")
edit_image = index.get("edit_image", "")
@@ -152,7 +178,7 @@ class ACEPlusDataset(BaseDataset):
edit_id, ref_id, src_image_list, src_mask_list = [], [], [], []
# parse editing image
if edit_image is None:
edit_image = Image.new("RGB", target_image.size, 255)
edit_image = Image.new("RGB", target_image.size, (255, 255, 255))
edit_mask = Image.new("L", edit_image.size, 255)
elif edit_mask is None:
edit_mask = Image.new("L", edit_image.size, 255)
@@ -163,7 +189,7 @@ class ACEPlusDataset(BaseDataset):
if ref_image is not None:
src_image_list.append(ref_image)
ref_id.append(1)
src_mask_list.append(Image.new("L", ref_image.size, 255))
src_mask_list.append(Image.new("L", ref_image.size, 0))
image = transform_image(torch.tensor(np.array(target_image).astype(np.float32)))
if edit_mask is not None:
@@ -183,23 +209,24 @@ class ACEPlusDataset(BaseDataset):
repainting_scale = self.repaiting_scale
for e_i in edit_id:
src_image_list[e_i] = src_image_list[e_i] * (1 - repainting_scale * src_mask_list[e_i])
# use fill mode(cat img, not align)
# ensure the height of ref image is aligned with that of target image
size = image.shape[1:]
ref_image_list = ensure_src_align_target_h_mode(src_image_list, size,
image_id=ref_id,
interpolation=InterpolationMode.BILINEAR)
ref_mask_list = ensure_src_align_target_h_mode(src_mask_list, size,
image_id=ref_id,
interpolation=InterpolationMode.NEAREST_EXACT)
ref_image_list, ret_h = ensure_src_align_target_padding_mode(src_image_list, size,
image_id=ref_id,
interpolation=InterpolationMode.NEAREST_EXACT)
ref_mask_list, ret_h = ensure_src_align_target_padding_mode(src_mask_list, size,
size_h=ret_h,
image_id=ref_id,
interpolation=InterpolationMode.NEAREST_EXACT)
edit_image_list = ensure_src_align_target_h_mode(src_image_list, size,
image_id=edit_id,
interpolation=InterpolationMode.BILINEAR)
interpolation=InterpolationMode.NEAREST_EXACT)
edit_mask_list = ensure_src_align_target_h_mode(src_mask_list, size,
image_id=edit_id,
interpolation=InterpolationMode.NEAREST_EXACT)
src_image_list = [torch.cat(ref_image_list + edit_image_list, dim=-1)]
src_mask_list = [torch.cat(ref_mask_list + edit_mask_list, dim=-1)]
image = torch.cat(ref_image_list + [image], dim=-1)
@@ -214,16 +241,27 @@ class ACEPlusDataset(BaseDataset):
d = self.d, interpolation=InterpolationMode.BILINEAR) for i in src_image_list]
src_mask_list = [ensure_limit_sequence(i, max_seq_len = self.max_seq_len,
d = self.d, interpolation=InterpolationMode.NEAREST_EXACT) for i in src_mask_list]
# print(src_image_list[0].shape, src_mask_list[0].shape, image.shape, image_mask.shape)
if self.modify_mode:
# To be modified regions according to mask
modify_image_list = [ii * im for ii, im in zip(src_image_list, src_mask_list)]
# To be edited regions according to mask
src_image_list = [ii * (1 - im) for ii, im in zip(src_image_list, src_mask_list)]
else:
src_image_list = src_image_list
modify_image_list = src_image_list
item = {
"src_image_list": src_image_list,
"src_mask_list": src_mask_list,
"modify_image_list": modify_image_list,
"image": image,
"image_mask": image_mask,
"edit_id": edit_id,
"ref_id": ref_id,
"prompt": prompt,
"edit_key": index["edit_key"] if "edit_key" in index else ""
"edit_key": index["edit_key"] if "edit_key" in index else "",
"sample_id": sample_id
}
return item
+51 -23
View File
@@ -127,11 +127,13 @@ class LatentDiffusionACEPlus(LatentDiffusion):
if x is None: return x
return F.interpolate(x.unsqueeze(0), size = size, mode='nearest-exact')
def parse_ref_and_edit(self, src_image,
modify_image,
src_image_mask,
text_embedding,
#text_mask,
edit_id):
edit_image = []
modi_image = []
edit_mask = []
ref_image = []
ref_mask = []
@@ -140,11 +142,14 @@ class LatentDiffusionACEPlus(LatentDiffusion):
ref_id = []
txt = []
txt_y = []
for sample_id, (one_src, one_src_mask,
for sample_id, (one_src,
one_modify,
one_src_mask,
one_text_embedding,
one_text_y,
# one_text_mask,
one_edit_id) in enumerate(zip(src_image,
modify_image,
src_image_mask,
text_embedding["context"],
text_embedding["y"],
@@ -160,10 +165,22 @@ class LatentDiffusionACEPlus(LatentDiffusion):
# process edit image & edit image mask
current_edit_image = to_device([one_src[i] for i in one_edit_id], strict=False)
current_edit_image = [v.squeeze(0) for v in self.encode_first_stage(current_edit_image)]
current_edit_image_mask = to_device([one_src_mask[i] for i in one_edit_id], strict=False)
current_edit_image_mask = [self.reshape_func(m).squeeze(0) for m in current_edit_image_mask]
# process modi image
current_modify_image = to_device([one_modify[i] for i in one_edit_id],
strict=False)
current_modify_image = [
v.squeeze(0)
for v in self.encode_first_stage(current_modify_image)
]
current_edit_image_mask = to_device(
[one_src_mask[i] for i in one_edit_id], strict=False)
current_edit_image_mask = [
self.reshape_func(m).squeeze(0)
for m in current_edit_image_mask
]
edit_image.append(current_edit_image)
modi_image.append(current_modify_image)
edit_mask.append(current_edit_image_mask)
ref_context.append(one_text_embedding[:len(ref_id[-1])])
ref_y.append(one_text_y[:len(ref_id[-1])])
@@ -177,6 +194,7 @@ class LatentDiffusionACEPlus(LatentDiffusion):
txt_y.append(one_text_y[-1])
return {
"edit": edit_image,
'modify': modi_image,
"edit_mask": edit_mask,
"edit_id": edit_id,
"ref_context": ref_context,
@@ -201,8 +219,9 @@ class LatentDiffusionACEPlus(LatentDiffusion):
return mask
def forward_train(self,
src_image_list =[],
src_mask_list =[],
src_image_list=[],
modify_image_list=[],
src_mask_list=[],
edit_id=[],
image=None,
image_mask=None,
@@ -210,18 +229,19 @@ class LatentDiffusionACEPlus(LatentDiffusion):
prompt=[],
**kwargs):
'''
Args:
src_image: list of list of src_image
src_image_mask: list of list of src_image_mask
image: target image
image_mask: target image mask
noise: default is None, generate automaticly
ref_prompt: list of list of text
prompt: list of text
**kwargs:
Returns:
'''
assert check_list_of_list(src_image_list) and check_list_of_list(src_mask_list)
Args:
src_image: list of list of src_image
src_image_mask: list of list of src_image_mask
image: target image
image_mask: target image mask
noise: default is None, generate automaticly
ref_prompt: list of list of text
prompt: list of text
**kwargs:
Returns:
'''
assert check_list_of_list(src_image_list) and check_list_of_list(
src_mask_list)
assert self.cond_stage_model is not None
gc_seg = kwargs.pop("gc_seg", [])
@@ -263,7 +283,7 @@ class LatentDiffusionACEPlus(LatentDiffusion):
# process image mask
context['x_mask'] = x_mask
ref_edit_context = self.parse_ref_and_edit(src_image_list, src_mask_list, context, edit_id)
ref_edit_context = self.parse_ref_and_edit(src_image_list, modify_image_list, src_mask_list, context, edit_id)
context.update(ref_edit_context)
teacher_context = copy.deepcopy(context)
@@ -284,6 +304,7 @@ class LatentDiffusionACEPlus(LatentDiffusion):
@torch.no_grad()
def forward_test(self,
src_image_list=[],
modify_image_list=[],
src_mask_list=[],
edit_id=[],
image=None,
@@ -300,6 +321,7 @@ class LatentDiffusionACEPlus(LatentDiffusion):
outputs = self.forward_editing(
src_image_list=src_image_list,
src_mask_list=src_mask_list,
modify_image_list=modify_image_list,
edit_id=edit_id,
image=image,
image_mask=image_mask,
@@ -318,6 +340,7 @@ class LatentDiffusionACEPlus(LatentDiffusion):
@torch.no_grad()
def forward_editing(self,
src_image_list=[],
modify_image_list=None,
src_mask_list=[],
edit_id=[],
image=None,
@@ -331,8 +354,8 @@ class LatentDiffusionACEPlus(LatentDiffusion):
**kwargs
):
# gc_seg is unused
prompt, image, image_mask, src_image, src_image_mask, edit_id = limit_batch_data(
[prompt, image, image_mask, src_image_list, src_mask_list, edit_id], log_num)
prompt, image, image_mask, src_image, modify_image, src_image_mask, edit_id = limit_batch_data(
[prompt, image, image_mask, src_image_list, modify_image_list, src_mask_list, edit_id], log_num)
assert check_list_of_list(src_image) and check_list_of_list(src_image_mask)
assert self.cond_stage_model is not None
align = kwargs.pop("align", [])
@@ -361,7 +384,7 @@ class LatentDiffusionACEPlus(LatentDiffusion):
image_mask = to_device(image_mask, strict=False)
x_mask = [self.reshape_func(i).squeeze(0) for i in image_mask]
context['x_mask'] = x_mask
ref_edit_context = self.parse_ref_and_edit(src_image, src_image_mask, context, edit_id)
ref_edit_context = self.parse_ref_and_edit(src_image, modify_image, src_image_mask, context, edit_id)
context.update(ref_edit_context)
# UNet use input n_prompt
# model = self.model_ema if self.use_ema and self.eval_ema else self.model
@@ -388,13 +411,17 @@ class LatentDiffusionACEPlus(LatentDiffusion):
for i in range(len(prompt)):
rec_img = torch.clamp((x_samples[i].float() + 1.0) / 2.0, min=0.0, max=1.0)
rec_img = rec_img.squeeze(0)
edit_imgs, edit_img_masks = [], []
edit_imgs, modify_imgs, edit_img_masks = [], [], []
if src_image is not None and src_image[i] is not None:
if src_image_mask[i] is None:
src_image_mask[i] = [None] * len(src_image[i])
for edit_img, edit_mask in zip(src_image[i], src_image_mask[i]):
for edit_img, modify_img, edit_mask in zip(src_image[i], modify_image_list[i], src_image_mask[i]):
edit_img = torch.clamp((edit_img.float() + 1.0) / 2.0, min=0.0, max=1.0)
edit_imgs.append(edit_img.squeeze(0))
modify_img = torch.clamp((modify_img.float() + 1.0) / 2.0,
min=0.0,
max=1.0)
modify_imgs.append(modify_img.squeeze(0))
if edit_mask is None:
edit_mask = torch.ones_like(edit_img[[0], :, :])
edit_img_masks.append(edit_mask)
@@ -402,6 +429,7 @@ class LatentDiffusionACEPlus(LatentDiffusion):
'reconstruct_image': rec_img,
'instruction': prompt[i],
'edit_image': edit_imgs if len(edit_imgs) > 0 else None,
'modify_image': modify_imgs if len(modify_imgs) > 0 else None,
'edit_mask': edit_img_masks if len(edit_imgs) > 0 else None
}
if image is not None:
+18 -1
View File
@@ -9,11 +9,12 @@ from scepter.modules.utils.distribute import we
from scepter.modules.utils.probe import ProbeData
from tqdm import tqdm
@SOLVERS.register_class()
class ACEPlusSolver(LatentDiffusionSolver):
class FormalACEPlusSolver(LatentDiffusionSolver):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.probe_prompt = cfg.get("PROBE_PROMPT", None)
self.probe_hw = cfg.get("PROBE_HW", [])
@torch.no_grad()
def run_eval(self):
self.eval_mode()
@@ -75,11 +76,24 @@ class ACEPlusSolver(LatentDiffusionSolver):
self.after_all_iter(self.hooks_dict[self._mode])
def run_step_val(self, batch_data, batch_idx=0, step=None, rank=None):
sample_id_list = batch_data['sample_id']
loss_dict = {}
with torch.autocast(device_type='cuda',
enabled=self.use_amp,
dtype=self.dtype):
results = self.model.forward_train(**batch_data)
loss = results['loss']
for sample_id in sample_id_list:
loss_dict[sample_id] = loss.detach().cpu().numpy()
return loss_dict
def save_results(self, results):
log_data, log_label = [], []
for result in results:
ret_images, ret_labels = [], []
edit_image = result.get('edit_image', None)
modify_image = result.get('modify_image', None)
edit_mask = result.get('edit_mask', None)
if edit_image is not None:
for i, edit_img in enumerate(result['edit_image']):
@@ -87,6 +101,8 @@ class ACEPlusSolver(LatentDiffusionSolver):
continue
ret_images.append((edit_img.permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
ret_labels.append(f'edit_image{i}; ')
ret_images.append((modify_image[i].permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
ret_labels.append(f'modify_image{i}; ')
if edit_mask is not None:
ret_images.append((edit_mask[i].permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
ret_labels.append(f'edit_mask{i}; ')
@@ -143,6 +159,7 @@ class ACEPlusSolver(LatentDiffusionSolver):
"image": [torch.zeros(3, self.probe_hw[0], self.probe_hw[1])],
"image_mask": [torch.ones(1, self.probe_hw[0], self.probe_hw[1])],
"src_image_list": [[]],
"modify_image_list": [[]],
"src_mask_list": [[]],
"edit_id": [[]],
"height": self.probe_hw[0],
+104 -7
View File
@@ -653,25 +653,122 @@ class FluxMRACEPlus(FluxMR):
def prepare_input(self, x, cond):
context, y = cond["context"], cond["y"]
batch_frames, batch_frames_ids = [], []
for ix, shape, imask, ie, ie_mask in zip(x, cond["x_shapes"], cond["x_mask"],
cond["edit"], cond["edit_mask"]):
for ix, shape, imask, ie, ie_mask in zip(x,
cond['x_shapes'],
cond['x_mask'],
cond['edit'],
cond['edit_mask']):
# unpack image from sequence
ix = ix[:, :shape[0] * shape[1]].view(-1, shape[0], shape[1])
imask = torch.ones_like(ix[[0], :, :]) if imask is None else imask.squeeze(0)
imask = torch.ones_like(
ix[[0], :, :]) if imask is None else imask.squeeze(0)
if len(ie) > 0:
ie = [iie.squeeze(0) for iie in ie]
ie_mask = [torch.ones((ix.shape[0] * 4, ix.shape[1], ix.shape[2])) if iime is None else iime.squeeze(0) for iime in ie_mask]
ie_mask = [
torch.ones(
(ix.shape[0] * 4, ix.shape[1],
ix.shape[2])) if iime is None else iime.squeeze(0)
for iime in ie_mask
]
ie = torch.cat(ie, dim=-1)
ie_mask = torch.cat(ie_mask, dim=-1)
else:
ie, ie_mask = torch.zeros_like(ix).to(x), torch.ones_like(imask).to(x)
ie, ie_mask = torch.zeros_like(ix).to(x), torch.ones_like(
imask).to(x),
ix = torch.cat([ix, ie, ie_mask], dim=0)
c, h, w = ix.shape
ix = rearrange(ix, "c (h ph) (w pw) -> (h w) (c ph pw)", ph=2, pw=2)
ix = rearrange(ix,
'c (h ph) (w pw) -> (h w) (c ph pw)',
ph=2,
pw=2)
ix_id = torch.zeros(h // 2, w // 2, 3)
ix_id[..., 1] = ix_id[..., 1] + torch.arange(h // 2)[:, None]
ix_id[..., 2] = ix_id[..., 2] + torch.arange(w // 2)[None, :]
ix_id = rearrange(ix_id, "h w c -> (h w) c")
ix_id = rearrange(ix_id, 'h w c -> (h w) c')
batch_frames.append([ix])
batch_frames_ids.append([ix_id])
x_list, x_id_list, mask_x_list, x_seq_length = [], [], [], []
for frames, frame_ids in zip(batch_frames, batch_frames_ids):
proj_frames = []
for idx, one_frame in enumerate(frames):
one_frame = self.img_in(one_frame)
proj_frames.append(one_frame)
ix = torch.cat(proj_frames, dim=0)
if_id = torch.cat(frame_ids, dim=0)
x_list.append(ix)
x_id_list.append(if_id)
mask_x_list.append(torch.ones(ix.shape[0]).to(ix.device, non_blocking=True).bool())
x_seq_length.append(ix.shape[0])
# if len(x_list) < 1: import pdb;pdb.set_trace()
x = pad_sequence(tuple(x_list), batch_first=True)
x_ids = pad_sequence(tuple(x_id_list), batch_first=True).to(x) # [b,pad_seq,2] pad (0.,0.) at dim2
mask_x = pad_sequence(tuple(mask_x_list), batch_first=True)
if isinstance(context, list):
txt_list, mask_txt_list, y_list = [], [], []
for sample_id, (ctx, yy) in enumerate(zip(context, y)):
txt_list.append(self.txt_in(ctx.to(x)))
mask_txt_list.append(torch.ones(txt_list[-1].shape[0]).to(ctx.device, non_blocking=True).bool())
y_list.append(yy.to(x))
txt = pad_sequence(tuple(txt_list), batch_first=True)
txt_ids = torch.zeros(txt.shape[0], txt.shape[1], 3).to(x)
mask_txt = pad_sequence(tuple(mask_txt_list), batch_first=True)
y = torch.cat(y_list, dim=0)
assert y.ndim == 2 and txt.ndim == 3
else:
txt = self.txt_in(context)
txt_ids = torch.zeros(context.shape[0], context.shape[1], 3).to(x)
mask_txt = torch.ones(context.shape[0], context.shape[1]).to(x.device, non_blocking=True).bool()
return x, x_ids, txt, txt_ids, y, mask_x, mask_txt, x_seq_length
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
FluxMRACEPlus.para_dict,
set_name=True)
@BACKBONES.register_class()
class FluxMRModiACEPlus(FluxMR):
def __init__(self, cfg, logger = None):
super().__init__(cfg, logger)
def prepare_input(self, x, cond):
context, y = cond["context"], cond["y"]
batch_frames, batch_frames_ids = [], []
for ix, shape, imask, ie, im, ie_mask in zip(x,
cond['x_shapes'],
cond['x_mask'],
cond['edit'],
cond['modify'],
cond['edit_mask']):
# unpack image from sequence
ix = ix[:, :shape[0] * shape[1]].view(-1, shape[0], shape[1])
imask = torch.ones_like(
ix[[0], :, :]) if imask is None else imask.squeeze(0)
if len(ie) > 0:
ie = [iie.squeeze(0) for iie in ie]
im = [iim.squeeze(0) for iim in im]
ie_mask = [
torch.ones(
(ix.shape[0] * 4, ix.shape[1],
ix.shape[2])) if iime is None else iime.squeeze(0)
for iime in ie_mask
]
im = torch.cat(im, dim=-1)
ie = torch.cat(ie, dim=-1)
ie_mask = torch.cat(ie_mask, dim=-1)
else:
ie, im, ie_mask = torch.zeros_like(ix).to(x), torch.zeros_like(ix).to(x), torch.ones_like(
imask).to(x),
ix = torch.cat([ix, ie, im, ie_mask], dim=0)
c, h, w = ix.shape
ix = rearrange(ix,
'c (h ph) (w pw) -> (h w) (c ph pw)',
ph=2,
pw=2)
ix_id = torch.zeros(h // 2, w // 2, 3)
ix_id[..., 1] = ix_id[..., 1] + torch.arange(h // 2)[:, None]
ix_id[..., 2] = ix_id[..., 2] + torch.arange(w // 2)[None, :]
ix_id = rearrange(ix_id, 'h w c -> (h w) c')
batch_frames.append([ix])
batch_frames_ids.append([ix_id])
x_list, x_id_list, mask_x_list, x_seq_length = [], [], [], []
+30 -5
View File
@@ -3,7 +3,7 @@ ENV:
SEED: 1999
SOLVER:
# NAME DESCRIPTION: TYPE: default: 'LatentUfitSolver'
NAME: ACEPlusSolver
NAME: FormalACEPlusSolver
# MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
MAX_STEPS: 100000
# USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False
@@ -79,10 +79,10 @@ SOLVER:
#
DIFFUSION_MODEL:
# NAME DESCRIPTION: TYPE: default: 'Flux'
NAME: FluxMRACEPlus
PRETRAINED_MODEL: ${FLUX_FILL_PATH}/flux1-fill-dev.safetensors
NAME: FluxMRModiACEPlus
PRETRAINED_MODEL: ${ACE_PLUS_FFT_MODEL}
# IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
IN_CHANNELS: 384
IN_CHANNELS: 448
# OUT_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
OUT_CHANNELS: 64
# HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
@@ -220,6 +220,7 @@ SOLVER:
NAME: ACEPlusDataset
MODE: train
DATA_LIST: data/train.csv
MODIFY_MODE: True
DELIMITER: "#;#"
# input_image, input_mask, input_reference_image, target_image, instruction, task_type
FIELDS: ["edit_image", "edit_mask", "ref_image", "target_image", "prompt", "data_type"]
@@ -229,7 +230,7 @@ SOLVER:
D: 16
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
NUM_WORKERS: 0
SAMPLER:
NAME: LoopSampler
@@ -237,6 +238,7 @@ SOLVER:
NAME: ACEPlusDataset
MODE: eval
DATA_LIST: data/train.csv
MODIFY_MODE: True
DELIMITER: "#;#"
# input_image, input_mask, input_reference_image, target_image, instruction, task_type
FIELDS: [ "edit_image", "edit_mask", "ref_image", "target_image", "prompt", "data_type" ]
@@ -257,6 +259,29 @@ SOLVER:
- NAME: ACECheckpointHook
INTERVAL: 250
PRIORITY: 200
- NAME: ValLossHook
VAL_INTERVAL: 250
VAL_LIMITATION_SIZE: 1000000
VAL_SEED: 42
META_FIELD: [ 'edit_key' ]
PRIORITY: 5
DATA:
NAME: ACEPlusDataset
MODE: eval
PIN_MEMORY: True
BATCH_SIZE: 1
USE_NUM: -1
NUM_WORKERS: 4
DATA_LIST: data/train.csv
MODIFY_MODE: True
DELIMITER: "#;#"
# input_image, input_mask, input_reference_image, target_image, instruction, task_type
FIELDS: [ "edit_image", "edit_mask", "ref_image", "target_image", "prompt", "data_type" ]
PATH_PREFIX: ""
EDIT_TYPE_LIST: [ ]
MAX_SEQ_LEN: 2048
D: 16
- NAME: ProbeDataHook
PROB_INTERVAL: 50
PRIORITY: 0
+1 -1
View File
@@ -3,7 +3,7 @@ ENV:
SEED: 1999
SOLVER:
# NAME DESCRIPTION: TYPE: default: 'LatentUfitSolver'
NAME: ACEPlusSolver
NAME: FormalACEPlusSolver
# MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000
MAX_STEPS: 100000
# USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False