From 16d682baa5ee6dadc3a3bafd6c67d2498a3bc924 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=9A=93=E7=AB=A5?= Date: Tue, 4 Mar 2025 09:37:13 +0800 Subject: [PATCH] modify demo fft --- README.md | 146 ++++++++- demo.py | 524 -------------------------------- examples/examples.py | 152 +++++++++ infer.py | 228 -------------- modules/__init__.py | 4 +- modules/ace_plus_dataset.py | 66 +++- modules/ace_plus_ldm.py | 74 +++-- modules/ace_plus_solver.py | 19 +- modules/flux.py | 111 ++++++- train_config/ace_plus_fft.yaml | 35 ++- train_config/ace_plus_lora.yaml | 2 +- 11 files changed, 548 insertions(+), 813 deletions(-) delete mode 100644 demo.py delete mode 100644 infer.py diff --git a/README.md b/README.md index 95c6a4c..eba7ae6 100644 --- a/README.md +++ b/README.md @@ -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 + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
Input Reference ImageInput Edit ImageInput Edit MaskOutputInstructionFunction
"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.""Character ID Consistency Generation"
"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.""Subject Consistency Generation"
"The item is put on the table.""Subject Consistency Editing"
"The logo is printed on the headphones.""Subject Consistency Editing"
"The woman dresses this skirt.""Try On"
"{image}, the man faces the camera.""Face swap"
"{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.""Super-resolution"
"a blue hand""Regional Editing"
"Mechanical hands like a robot""Regional Editing"
"{image} Beautiful female portrait, Robot with smooth White transparent carbon shell, rococo detailing, Natural lighting, Highly detailed, Cinematic, 4K.""Recolorizing"
"{image} Beautiful female portrait, Robot with smooth White transparent carbon shell, rococo detailing, Natural lighting, Highly detailed, Cinematic, 4K.""Depth Guided Generation"
"{image} Beautiful female portrait, Robot with smooth White transparent carbon shell, rococo detailing, Natural lighting, Highly detailed, Cinematic, 4K.""Contour Guided Generation"
+ + +## 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 diff --git a/demo.py b/demo.py deleted file mode 100644 index a827aa7..0000000 --- a/demo.py +++ /dev/null @@ -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 = "Please provide the reference image." - 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"{err_msg}" - 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) diff --git a/examples/examples.py b/examples/examples.py index d66441d..8c1fa81 100644 --- a/examples/examples.py +++ b/examples/examples.py @@ -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" + } ] \ No newline at end of file diff --git a/infer.py b/infer.py deleted file mode 100644 index aa4f190..0000000 --- a/infer.py +++ /dev/null @@ -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() - diff --git a/modules/__init__.py b/modules/__init__.py index d8bc0cb..00b1696 100644 --- a/modules/__init__.py +++ b/modules/__init__.py @@ -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 \ No newline at end of file diff --git a/modules/ace_plus_dataset.py b/modules/ace_plus_dataset.py index b67e3b6..22cc9f9 100644 --- a/modules/ace_plus_dataset.py +++ b/modules/ace_plus_dataset.py @@ -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 diff --git a/modules/ace_plus_ldm.py b/modules/ace_plus_ldm.py index 138ae09..68b461b 100644 --- a/modules/ace_plus_ldm.py +++ b/modules/ace_plus_ldm.py @@ -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: diff --git a/modules/ace_plus_solver.py b/modules/ace_plus_solver.py index 90d3fc8..d09a40b 100644 --- a/modules/ace_plus_solver.py +++ b/modules/ace_plus_solver.py @@ -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], diff --git a/modules/flux.py b/modules/flux.py index 7ac764c..3c57fbe 100644 --- a/modules/flux.py +++ b/modules/flux.py @@ -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 = [], [], [], [] diff --git a/train_config/ace_plus_fft.yaml b/train_config/ace_plus_fft.yaml index c63f5fe..5740324 100644 --- a/train_config/ace_plus_fft.yaml +++ b/train_config/ace_plus_fft.yaml @@ -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 diff --git a/train_config/ace_plus_lora.yaml b/train_config/ace_plus_lora.yaml index e47615b..dc261b2 100644 --- a/train_config/ace_plus_lora.yaml +++ b/train_config/ace_plus_lora.yaml @@ -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