diff --git a/.gitignore b/.gitignore
new file mode 100644
index 0000000..723ef36
--- /dev/null
+++ b/.gitignore
@@ -0,0 +1 @@
+.idea
\ No newline at end of file
diff --git a/README.md b/README.md
index 9b679a4..76f83bd 100644
--- a/README.md
+++ b/README.md
@@ -1 +1,57 @@
-# ACE_plus
\ No newline at end of file
+
+
+
++: Instruction-Based Image Creation and Editing
via Context-Aware Content Filling
+
+
+
+
+
+
+
+
+
+ Chaojie Mao
+ ·
+ Jingfeng Zhang
+ ·
+ Yulin Pan
+ ·
+ Zeyinzi Jiang
+ ·
+ Zhen Han
+
+ ·
+ Yu Liu
+ ·
+ Jingren Zhou
+
+ Tongyi Lab, Alibaba Group
+
+
+
+
+
+ |
+
+
+
+## 📢 News
+* **[2025.01.03]** Release the paper of ACE++ on arxiv.
+
+## 🚀 Installation
+Install the necessary packages with `pip`:
+```bash
+pip install -r requirements.txt
+```
+
+
+## 📝 Citation
+
+```bibtex
+@article{,
+ title={},
+ author={},
+ journal={},
+ year={2025}
+}
+```
\ No newline at end of file
diff --git a/__init__.py b/__init__.py
new file mode 100644
index 0000000..422787a
--- /dev/null
+++ b/__init__.py
@@ -0,0 +1 @@
+import modules
\ No newline at end of file
diff --git a/assets/ace_method/method++.jpg b/assets/ace_method/method++.jpg
new file mode 100644
index 0000000..0719d89
Binary files /dev/null and b/assets/ace_method/method++.jpg differ
diff --git a/config/ace_plus_diffusers_infer.yaml b/config/ace_plus_diffusers_infer.yaml
new file mode 100644
index 0000000..b540d50
--- /dev/null
+++ b/config/ace_plus_diffusers_infer.yaml
@@ -0,0 +1,25 @@
+NAME: ace_plus_diffuser_infer
+IS_DEFAULT: True
+USE_DYNAMIC_MODEL: False
+INFERENCE_TYPE: ACE_DIFFUSER_PLUS
+DEFAULT_PARAS:
+ PARAS:
+ #
+ INPUT:
+ INPUT_IMAGE:
+ INPUT_MASK:
+ TASK:
+ PROMPT: ""
+ OUTPUT_HEIGHT: 1024
+ OUTPUT_WIDTH: 1024
+ SAMPLER: flow_euler
+ SAMPLE_STEPS: 28
+ GUIDE_SCALE: 50
+ SEED: 42
+ MAX_SEQ_LENGTH: 4096
+ OUTPUT:
+ LATENT:
+ IMAGES:
+ SEED:
+MODEL:
+ PRETRAINED_MODEL: ${FLUX_FILL_PATH}
\ No newline at end of file
diff --git a/demo.py b/demo.py
new file mode 100644
index 0000000..3f130cf
--- /dev/null
+++ b/demo.py
@@ -0,0 +1,397 @@
+# -*- 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_inference import ACEPlusInference
+from inference.ace_plus_diffusers import ACEPlusDiffuserInference
+from inference.utils import edit_preprocess
+
+inference_dict = {
+ "ACE_PLUS": ACEPlusInference,
+ "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
+
+ 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.Accordion(label='Related Input Image', open=False):
+ self.generation_info_preview = gr.Text(
+ lines=2,
+ )
+ 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(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.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.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')
+
+
+
+
+ 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
+ for preprocessor in task_info.get("PREPROCESSOR", []):
+ if preprocessor["TYPE"] in self.edit_type_dict:
+ continue
+ 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):
+ 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(edit_image) < 1:
+ edit_image = None
+ edit_mask = None
+ else:
+ edit_image = pillow_convert(edit_image, "RGB")
+ edit_mask = Image.fromarray(edit_mask).convert('L')
+
+ edit_image.save("debug/edit_image.png") if edit_image is not None else None
+ edit_mask.save("debug/edit_mask.png") if edit_mask is not None else None
+ ref_image.save("debug/ref_image.png") if ref_image is not None else None
+
+ return edit_image, edit_mask, ref_image
+
+ def run_chat(
+ prompt,
+ ref_image,
+ edit_image,
+ task_type,
+ edit_type,
+ cfg_scale,
+ step,
+ seed,
+ output_h,
+ output_w,
+ repainting_scale
+ ):
+ model_path = self.task_model[task_type]["MODEL_PATH"]
+ edit_info = self.edit_type_dict[edit_type]
+
+ 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_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"
+
+ 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_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/input/86e6357922a3c13cf3d0c177d2a5c382.jpg b/examples/input/86e6357922a3c13cf3d0c177d2a5c382.jpg
new file mode 100644
index 0000000..2638982
Binary files /dev/null and b/examples/input/86e6357922a3c13cf3d0c177d2a5c382.jpg differ
diff --git a/infer.py b/infer.py
new file mode 100644
index 0000000..88497c9
--- /dev/null
+++ b/infer.py
@@ -0,0 +1,221 @@
+# -*- 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 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
+ ):
+ 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", 0.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=['native', '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,
+ "repainting_scale": cfg.args.repainting_scale
+ }
+ # run examples
+ all_examples = [
+ ]
+ for example in all_examples:
+ example.update(params)
+ 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,
+ "lora_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/inference/__init__.py b/inference/__init__.py
new file mode 100644
index 0000000..e69de29
diff --git a/inference/ace_plus_diffusers.py b/inference/ace_plus_diffusers.py
new file mode 100644
index 0000000..67a2b6b
--- /dev/null
+++ b/inference/ace_plus_diffusers.py
@@ -0,0 +1,114 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+import random
+from collections import OrderedDict
+
+import torch, os
+from diffusers import FluxFillPipeline, FluxControlPipeline
+from scepter.modules.utils.config import Config
+from scepter.modules.utils.distribute import we
+from scepter.modules.utils.file_system import FS
+from scepter.modules.utils.logger import get_logger
+from transformers import T5TokenizerFast
+from .utils import ACEPlusImageProcessor
+
+
+class ACEPlusDiffuserInference():
+ def __init__(self, logger=None):
+ if logger is None:
+ logger = get_logger(name='ace_plus')
+ self.logger = logger
+ self.input = {}
+
+ def load_default(self, cfg):
+ if cfg is not None:
+ self.input_cfg = {k.lower(): v for k, v in cfg.INPUT.items()}
+ self.input = {k.lower(): dict(v).get('DEFAULT', None) if isinstance(v, (dict, OrderedDict, Config)) else v for k, v in cfg.INPUT.items()}
+ self.output = {k.lower(): v for k, v in cfg.OUTPUT.items()}
+
+ def init_from_cfg(self, cfg):
+ self.max_seq_len = cfg.get("MAX_SEQ_LEN", 4096)
+ self.image_processor = ACEPlusImageProcessor(max_seq_len=self.max_seq_len)
+
+ local_folder = FS.get_dir_to_local_dir(cfg.MODEL.PRETRAINED_MODEL)
+
+ self.pipe = FluxFillPipeline.from_pretrained(local_folder, torch_dtype=torch.bfloat16).to("cuda")
+
+ tokenizer_2 = T5TokenizerFast.from_pretrained(os.path.join(local_folder, "tokenizer_2"),
+ additional_special_tokens=["{image}"])
+ self.pipe.tokenizer_2 = tokenizer_2
+ self.load_default(cfg.DEFAULT_PARAS)
+
+
+ def prepare_input(self,
+ image,
+ mask,
+ batch_size=1,
+ dtype = torch.bfloat16,
+ num_images_per_prompt=1,
+ height=512,
+ width=512,
+ generator=None):
+ num_channels_latents = self.pipe.vae.config.latent_channels
+ # import pdb;pdb.set_trace()
+ mask, masked_image_latents = self.pipe.prepare_mask_latents(
+ mask.unsqueeze(0),
+ image.unsqueeze(0).to(we.device_id, dtype = dtype),
+ batch_size,
+ num_channels_latents,
+ num_images_per_prompt,
+ height,
+ width,
+ dtype,
+ we.device_id,
+ generator,
+ )
+ # import pdb;pdb.set_trace()
+ masked_image_latents = torch.cat((masked_image_latents, mask), dim=-1)
+ return masked_image_latents
+
+ @torch.no_grad()
+ def __call__(self,
+ reference_image=None,
+ edit_image=None,
+ edit_mask=None,
+ prompt='',
+ task=None,
+ output_height=1024,
+ output_width=1024,
+ sampler='flow_euler',
+ sample_steps=28,
+ guide_scale=50,
+ lora_path=None,
+ seed=-1,
+ tar_index=0,
+ align=0,
+ repainting_scale=0,
+ **kwargs):
+ if isinstance(prompt, str):
+ prompt = [prompt]
+ seed = seed if seed >= 0 else random.randint(0, 2 ** 32 - 1)
+ image, mask, out_h, out_w, slice_w = self.image_processor.preprocess(reference_image, edit_image, edit_mask)
+ h, w = image.shape[1:]
+ masked_image_latents = self.prepare_input(image, mask,
+ batch_size=len(prompt) , height=h, width=w)
+
+ if lora_path is not None:
+ with FS.get_from(lora_path) as local_path:
+ self.pipe.load_lora_weights(local_path)
+
+ image = self.pipe(
+ prompt=prompt,
+ masked_image_latents=masked_image_latents,
+ height=h,
+ width=w,
+ guidance_scale=guide_scale,
+ num_inference_steps=sample_steps,
+ max_sequence_length=512,
+ generator=torch.Generator("cpu").manual_seed(seed),
+ ).images[0]
+ return self.image_processor.postprocess(image, slice_w, out_w, out_h), seed
+
+
+if __name__ == '__main__':
+ pass
\ No newline at end of file
diff --git a/inference/utils.py b/inference/utils.py
new file mode 100644
index 0000000..2944b5b
--- /dev/null
+++ b/inference/utils.py
@@ -0,0 +1,105 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+import math
+
+import torch
+import torchvision.transforms as T
+import numpy as np
+from scepter.modules.annotator.registry import ANNOTATORS
+from scepter.modules.utils.config import Config
+from PIL import Image
+
+
+def edit_preprocess(processor, device, edit_image, edit_mask):
+ if edit_image is None or processor is None:
+ return edit_image
+ processor = Config(cfg_dict=processor, load=False)
+ processor = ANNOTATORS.build(processor).to(device)
+ new_edit_image = processor(np.asarray(edit_image))
+ processor = processor.to("cpu")
+ del processor
+ new_edit_image = Image.fromarray(new_edit_image)
+ return Image.composite(new_edit_image, edit_image, edit_mask)
+
+class ACEPlusImageProcessor():
+ def __init__(self, max_aspect_ratio=4, d=16, max_seq_len=1024):
+ self.max_aspect_ratio = max_aspect_ratio
+ self.d = d
+ self.max_seq_len = max_seq_len
+ self.transforms = T.Compose([
+ T.ToTensor(),
+ T.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
+ ])
+
+ def image_check(self, image):
+ if image is None:
+ return image
+ # preprocess
+ W, H = image.size
+ if H / W > self.max_aspect_ratio:
+ image = T.CenterCrop([int(self.max_aspect_ratio * W), W])(image)
+ elif W / H > self.max_aspect_ratio:
+ image = T.CenterCrop([H, int(self.max_aspect_ratio * H)])(image)
+ return self.transforms(image)
+
+
+ def preprocess(self,
+ reference_image=None,
+ edit_image=None,
+ edit_mask=None,
+ height=1024,
+ width=1024,
+ repainting_scale = 1.0):
+ reference_image = self.image_check(reference_image)
+ edit_image = self.image_check(edit_image)
+ # for reference generation
+ if edit_image is None:
+ edit_image = torch.zeros([3, height, width])
+ edit_mask = torch.ones([1, height, width])
+ else:
+ edit_mask = np.asarray(edit_mask)
+ edit_mask = np.where(edit_mask > 128, 1, 0)
+ edit_mask = edit_mask.astype(
+ np.float32) if np.any(edit_mask) else np.ones_like(edit_mask).astype(
+ np.float32)
+ edit_mask = torch.tensor(edit_mask).unsqueeze(0)
+
+ edit_image = edit_image * (1 - edit_mask * repainting_scale)
+
+
+ out_h, out_w = edit_image.shape[-2:]
+
+ assert edit_mask is not None
+ if reference_image is not None:
+ # align height with edit_image
+ _, H, W = reference_image.shape
+ _, eH, eW = edit_image.shape
+ scale = eH / H
+ tH, tW = eH, int(W * scale)
+ reference_image = T.Resize((tH, tW), interpolation=T.InterpolationMode.BILINEAR, antialias=True)(reference_image)
+ edit_image = torch.cat([reference_image, edit_image], dim=-1)
+ edit_mask = torch.cat([torch.zeros([1, reference_image.shape[1], reference_image.shape[2]]), edit_mask], dim=-1)
+ slice_w = reference_image.shape[-1]
+ else:
+ slice_w = 0
+
+ H, W = edit_image.shape[-2:]
+ scale = min(1.0, math.sqrt(self.max_seq_len * 2 / ((H / self.d) * (W / self.d))))
+ rH = int(H * scale) // self.d * self.d # ensure divisible by self.d
+ rW = int(W * scale) // self.d * self.d
+ slice_w = int(slice_w * scale) // self.d * self.d
+
+ edit_image = T.Resize((rH, rW), interpolation=T.InterpolationMode.BILINEAR, antialias=True)(edit_image)
+ edit_mask = T.Resize((rH, rW), interpolation=T.InterpolationMode.NEAREST_EXACT, antialias=True)(edit_mask)
+
+ return edit_image, edit_mask, out_h, out_w, slice_w
+
+
+ def postprocess(self, image, slice_w, out_w, out_h):
+ w, h = image.size
+ if slice_w > 0:
+ output_image = image.crop((slice_w + 20, 0, w, h))
+ output_image = output_image.resize((out_w, out_h))
+ else:
+ output_image = image
+ return output_image
\ No newline at end of file
diff --git a/models/model_zoo.yaml b/models/model_zoo.yaml
new file mode 100644
index 0000000..1c1044f
--- /dev/null
+++ b/models/model_zoo.yaml
@@ -0,0 +1,28 @@
+MODEL:
+ PORTRAIT:
+ MODEL_PATH: ${PORTRAIT_MODEL_PATH}
+ SUBJECT:
+ MODEL_PATH: ${SUBJECT_MODEL_PATH}
+ LOCAL_EDITING:
+ MODEL_PATH: ${LOCAL_MODEL_PATH}
+ REPAINTING_SCALE: 0.5
+ PREPROCESSOR:
+ - NAME: CannyAnnotator
+ TYPE: canny_repainting
+ LOW_THRESHOLD: 100
+ HIGH_THRESHOLD: 200
+ - NAME: ColorAnnotator
+ TYPE: mosaic_repainting
+ RATIO: 64
+ - NAME: InfoDrawContourAnnotator
+ TYPE: contour_repainting
+ INPUT_NC: 3
+ OUTPUT_NC: 1
+ N_RESIDUAL_BLOCKS: 3
+ SIGMOID: True
+ PRETRAINED_MODEL: "ms://iic/scepter_annotator@annotator/ckpts/informative_drawing_contour_style.pth"
+ - NAME: MidasDetector
+ PRETRAINED_MODEL: "ms://iic/scepter_annotator@annotator/ckpts/dpt_hybrid-midas-501f0c75.pt"
+ TYPE: depth_repainting
+ - NAME: GrayAnnotator
+ TYPE: recolorizing
\ No newline at end of file
diff --git a/modules/__init__.py b/modules/__init__.py
new file mode 100644
index 0000000..7d1bd5d
--- /dev/null
+++ b/modules/__init__.py
@@ -0,0 +1,2 @@
+from .flux import Flux, ACEPlus
+from .embedder import ACEHFEmbedder, T5ACEPlusClipFluxEmbedder
\ No newline at end of file
diff --git a/modules/embedder.py b/modules/embedder.py
new file mode 100644
index 0000000..f1beece
--- /dev/null
+++ b/modules/embedder.py
@@ -0,0 +1,383 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+import warnings
+from contextlib import nullcontext
+
+import torch
+import torch.nn.functional as F
+import torch.utils.dlpack
+import transformers
+from scepter.modules.model.embedder.base_embedder import BaseEmbedder
+from scepter.modules.model.registry import EMBEDDERS
+from scepter.modules.model.tokenizer.tokenizer_component import (
+ basic_clean, canonicalize, heavy_clean, whitespace_clean)
+from scepter.modules.utils.config import dict_to_yaml
+from scepter.modules.utils.distribute import we
+from scepter.modules.utils.file_system import FS
+
+try:
+ from transformers import AutoTokenizer, T5EncoderModel
+except Exception as e:
+ warnings.warn(
+ f'Import transformers error, please deal with this problem: {e}')
+
+
+@EMBEDDERS.register_class()
+class ACETextEmbedder(BaseEmbedder):
+ """
+ Uses the OpenCLIP transformer encoder for text
+ """
+ """
+ Uses the OpenCLIP transformer encoder for text
+ """
+ para_dict = {
+ 'PRETRAINED_MODEL': {
+ 'value':
+ 'google/umt5-small',
+ 'description':
+ 'Pretrained Model for umt5, modelcard path or local path.'
+ },
+ 'TOKENIZER_PATH': {
+ 'value': 'google/umt5-small',
+ 'description':
+ 'Tokenizer Path for umt5, modelcard path or local path.'
+ },
+ 'FREEZE': {
+ 'value': True,
+ 'description': ''
+ },
+ 'USE_GRAD': {
+ 'value': False,
+ 'description': 'Compute grad or not.'
+ },
+ 'CLEAN': {
+ 'value':
+ 'whitespace',
+ 'description':
+ 'Set the clean strtegy for tokenizer, used when TOKENIZER_PATH is not None.'
+ },
+ 'LAYER': {
+ 'value': 'last',
+ 'description': ''
+ },
+ 'LEGACY': {
+ 'value':
+ True,
+ 'description':
+ 'Whether use legacy returnd feature or not ,default True.'
+ }
+ }
+
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ pretrained_path = cfg.get('PRETRAINED_MODEL', None)
+ self.t5_dtype = cfg.get('T5_DTYPE', 'float32')
+ assert pretrained_path
+ with FS.get_dir_to_local_dir(pretrained_path,
+ wait_finish=True) as local_path:
+ self.model = T5EncoderModel.from_pretrained(
+ local_path,
+ torch_dtype=getattr(
+ torch,
+ 'float' if self.t5_dtype == 'float32' else self.t5_dtype))
+ tokenizer_path = cfg.get('TOKENIZER_PATH', None)
+ self.length = cfg.get('LENGTH', 77)
+
+ self.use_grad = cfg.get('USE_GRAD', False)
+ self.clean = cfg.get('CLEAN', 'whitespace')
+ self.added_identifier = cfg.get('ADDED_IDENTIFIER', None)
+ if tokenizer_path:
+ self.tokenize_kargs = {'return_tensors': 'pt'}
+ with FS.get_dir_to_local_dir(tokenizer_path,
+ wait_finish=True) as local_path:
+ if self.added_identifier is not None and isinstance(
+ self.added_identifier, list):
+ self.tokenizer = AutoTokenizer.from_pretrained(local_path)
+ else:
+ self.tokenizer = AutoTokenizer.from_pretrained(local_path)
+ if self.length is not None:
+ self.tokenize_kargs.update({
+ 'padding': 'max_length',
+ 'truncation': True,
+ 'max_length': self.length
+ })
+ self.eos_token = self.tokenizer(
+ self.tokenizer.eos_token)['input_ids'][0]
+ else:
+ self.tokenizer = None
+ self.tokenize_kargs = {}
+
+ self.use_grad = cfg.get('USE_GRAD', False)
+ self.clean = cfg.get('CLEAN', 'whitespace')
+
+ def freeze(self):
+ self.model = self.model.eval()
+ for param in self.parameters():
+ param.requires_grad = False
+
+ # encode && encode_text
+ def forward(self, tokens, return_mask=False, use_mask=True):
+ # tokenization
+ embedding_context = nullcontext if self.use_grad else torch.no_grad
+ with embedding_context():
+ if use_mask:
+ x = self.model(tokens.input_ids.to(we.device_id),
+ tokens.attention_mask.to(we.device_id))
+ else:
+ x = self.model(tokens.input_ids.to(we.device_id))
+ x = x.last_hidden_state
+
+ if return_mask:
+ return x.detach() + 0.0, tokens.attention_mask.to(we.device_id)
+ else:
+ return x.detach() + 0.0, None
+
+ def _clean(self, text):
+ if self.clean == 'whitespace':
+ text = whitespace_clean(basic_clean(text))
+ elif self.clean == 'lower':
+ text = whitespace_clean(basic_clean(text)).lower()
+ elif self.clean == 'canonicalize':
+ text = canonicalize(basic_clean(text))
+ elif self.clean == 'heavy':
+ text = heavy_clean(basic_clean(text))
+ return text
+
+ def encode(self, text, return_mask=False, use_mask=True):
+ if isinstance(text, str):
+ text = [text]
+ if self.clean:
+ text = [self._clean(u) for u in text]
+ assert self.tokenizer is not None
+ cont, mask = [], []
+ with torch.autocast(device_type='cuda',
+ enabled=self.t5_dtype in ('float16', 'bfloat16'),
+ dtype=getattr(torch, self.t5_dtype)):
+ for tt in text:
+ tokens = self.tokenizer([tt], **self.tokenize_kargs)
+ one_cont, one_mask = self(tokens,
+ return_mask=return_mask,
+ use_mask=use_mask)
+ cont.append(one_cont)
+ mask.append(one_mask)
+ if return_mask:
+ return torch.cat(cont, dim=0), torch.cat(mask, dim=0)
+ else:
+ return torch.cat(cont, dim=0)
+
+ def encode_list(self, text_list, return_mask=True):
+ cont_list = []
+ mask_list = []
+ for pp in text_list:
+ cont, cont_mask = self.encode(pp, return_mask=return_mask)
+ cont_list.append(cont)
+ mask_list.append(cont_mask)
+ if return_mask:
+ return cont_list, mask_list
+ else:
+ return cont_list
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('MODELS',
+ __class__.__name__,
+ ACETextEmbedder.para_dict,
+ set_name=True)
+
+@EMBEDDERS.register_class()
+class ACEHFEmbedder(BaseEmbedder):
+ para_dict = {
+ "HF_MODEL_CLS": {
+ "value": None,
+ "description": "huggingface cls in transfomer"
+ },
+ "MODEL_PATH": {
+ "value": None,
+ "description": "model folder path"
+ },
+ "HF_TOKENIZER_CLS": {
+ "value": None,
+ "description": "huggingface cls in transfomer"
+ },
+
+ "TOKENIZER_PATH": {
+ "value": None,
+ "description": "tokenizer folder path"
+ },
+ "MAX_LENGTH": {
+ "value": 77,
+ "description": "max length of input"
+ },
+ "OUTPUT_KEY": {
+ "value": "last_hidden_state",
+ "description": "output key"
+ },
+ "D_TYPE": {
+ "value": "float",
+ "description": "dtype"
+ },
+ "BATCH_INFER": {
+ "value": False,
+ "description": "batch infer"
+ }
+ }
+ para_dict.update(BaseEmbedder.para_dict)
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ hf_model_cls = cfg.get('HF_MODEL_CLS', None)
+ model_path = cfg.get("MODEL_PATH", None)
+ hf_tokenizer_cls = cfg.get('HF_TOKENIZER_CLS', None)
+ tokenizer_path = cfg.get('TOKENIZER_PATH', None)
+ self.max_length = cfg.get('MAX_LENGTH', 77)
+ self.output_key = cfg.get("OUTPUT_KEY", "last_hidden_state")
+ self.d_type = cfg.get("D_TYPE", "float")
+ self.clean = cfg.get("CLEAN", "whitespace")
+ self.batch_infer = cfg.get("BATCH_INFER", False)
+ self.added_identifier = cfg.get('ADDED_IDENTIFIER', None)
+ torch_dtype = getattr(torch, self.d_type)
+
+ assert hf_model_cls is not None and hf_tokenizer_cls is not None
+ assert model_path is not None and tokenizer_path is not None
+ with FS.get_dir_to_local_dir(tokenizer_path, wait_finish=True) as local_path:
+ self.tokenizer = getattr(transformers, hf_tokenizer_cls).from_pretrained(local_path,
+ max_length = self.max_length,
+ torch_dtype = torch_dtype,
+ additional_special_tokens=self.added_identifier)
+
+ with FS.get_dir_to_local_dir(model_path, wait_finish=True) as local_path:
+ self.hf_module = getattr(transformers, hf_model_cls).from_pretrained(local_path, torch_dtype = torch_dtype)
+
+
+ self.hf_module = self.hf_module.eval().requires_grad_(False)
+
+ def forward(self, text: list[str], return_mask = False):
+ batch_encoding = self.tokenizer(
+ text,
+ truncation=True,
+ max_length=self.max_length,
+ return_length=False,
+ return_overflowing_tokens=False,
+ padding="max_length",
+ return_tensors="pt",
+ )
+
+ outputs = self.hf_module(
+ input_ids=batch_encoding["input_ids"].to(self.hf_module.device),
+ attention_mask=None,
+ output_hidden_states=False,
+ )
+ if return_mask:
+ return outputs[self.output_key], batch_encoding['attention_mask'].to(self.hf_module.device)
+ else:
+ return outputs[self.output_key], None
+
+ def encode(self, text, return_mask = False):
+ if isinstance(text, str):
+ text = [text]
+ if self.clean:
+ text = [self._clean(u) for u in text]
+ if not self.batch_infer:
+ cont, mask = [], []
+ for tt in text:
+ one_cont, one_mask = self([tt], return_mask=return_mask)
+ cont.append(one_cont)
+ mask.append(one_mask)
+ if return_mask:
+ return torch.cat(cont, dim=0), torch.cat(mask, dim=0)
+ else:
+ return torch.cat(cont, dim=0)
+ else:
+ ret_data = self(text, return_mask = return_mask)
+ if return_mask:
+ return ret_data
+ else:
+ return ret_data[0]
+
+ def encode_list(self, text_list, return_mask=True):
+ cont_list = []
+ mask_list = []
+ for pp in text_list:
+ cont = self.encode(pp, return_mask=return_mask)
+ cont_list.append(cont[0]) if return_mask else cont_list.append(cont)
+ mask_list.append(cont[1]) if return_mask else mask_list.append(None)
+ if return_mask:
+ return cont_list, mask_list
+ else:
+ return cont_list
+
+ def encode_list_of_list(self, text_list, return_mask=True):
+ cont_list = []
+ mask_list = []
+ for pp in text_list:
+ cont = self.encode_list(pp, return_mask=return_mask)
+ cont_list.append(cont[0]) if return_mask else cont_list.append(cont)
+ mask_list.append(cont[1]) if return_mask else mask_list.append(None)
+ if return_mask:
+ return cont_list, mask_list
+ else:
+ return cont_list
+
+ def _clean(self, text):
+ if self.clean == 'whitespace':
+ text = whitespace_clean(basic_clean(text))
+ elif self.clean == 'lower':
+ text = whitespace_clean(basic_clean(text)).lower()
+ elif self.clean == 'canonicalize':
+ text = canonicalize(basic_clean(text))
+ return text
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('EMBEDDER',
+ __class__.__name__,
+ ACEHFEmbedder.para_dict,
+ set_name=True)
+
+@EMBEDDERS.register_class()
+class T5ACEPlusClipFluxEmbedder(BaseEmbedder):
+ """
+ Uses the OpenCLIP transformer encoder for text
+ """
+ para_dict = {
+ 'T5_MODEL': {},
+ 'CLIP_MODEL': {}
+ }
+
+ def __init__(self, cfg, logger=None):
+ super().__init__(cfg, logger=logger)
+ self.t5_model = EMBEDDERS.build(cfg.T5_MODEL, logger=logger)
+ self.clip_model = EMBEDDERS.build(cfg.CLIP_MODEL, logger=logger)
+
+ def encode(self, text, return_mask = False):
+ t5_embeds = self.t5_model.encode(text, return_mask = return_mask)
+ clip_embeds = self.clip_model.encode(text, return_mask = return_mask)
+ # change embedding strategy here
+ return {
+ 'context': t5_embeds,
+ 'y': clip_embeds,
+ }
+
+ def encode_list(self, text, return_mask = False):
+ t5_embeds = self.t5_model.encode_list(text, return_mask = return_mask)
+ clip_embeds = self.clip_model.encode_list(text, return_mask = return_mask)
+ # change embedding strategy here
+ return {
+ 'context': t5_embeds,
+ 'y': clip_embeds,
+ }
+
+ def encode_list_of_list(self, text, return_mask = False):
+ t5_embeds = self.t5_model.encode_list_of_list(text, return_mask = return_mask)
+ clip_embeds = self.clip_model.encode_list_of_list(text, return_mask = return_mask)
+ # change embedding strategy here
+ return {
+ 'context': t5_embeds,
+ 'y': clip_embeds,
+ }
+
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('EMBEDDER',
+ __class__.__name__,
+ T5ACEPlusClipFluxEmbedder.para_dict,
+ set_name=True)
\ No newline at end of file
diff --git a/modules/flux.py b/modules/flux.py
new file mode 100644
index 0000000..6d097a3
--- /dev/null
+++ b/modules/flux.py
@@ -0,0 +1,632 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+import math, torch
+from collections import OrderedDict
+from functools import partial
+from einops import rearrange, repeat
+from scepter.modules.model.base_model import BaseModel
+from scepter.modules.model.registry import BACKBONES
+from scepter.modules.utils.config import dict_to_yaml
+from scepter.modules.utils.distribute import we
+from scepter.modules.utils.file_system import FS
+from torch import Tensor, nn
+from torch.nn.utils.rnn import pad_sequence
+from torch.utils.checkpoint import checkpoint_sequential
+from .layers import (DoubleStreamBlock, EmbedND, LastLayer,
+ MLPEmbedder, SingleStreamBlock,
+ timestep_embedding)
+
+@BACKBONES.register_class()
+class Flux(BaseModel):
+ """
+ Transformer backbone Diffusion model with RoPE.
+ """
+ para_dict = {
+ "IN_CHANNELS": {
+ "value": 64,
+ "description": "model's input channels."
+ },
+ "OUT_CHANNELS": {
+ "value": 64,
+ "description": "model's output channels."
+ },
+ "HIDDEN_SIZE": {
+ "value": 1024,
+ "description": "model's hidden size."
+ },
+ "NUM_HEADS": {
+ "value": 16,
+ "description": "number of heads in the transformer."
+ },
+ "AXES_DIM": {
+ "value": [16, 56, 56],
+ "description": "dimensions of the axes of the positional encoding."
+ },
+ "THETA": {
+ "value": 10_000,
+ "description": "theta for positional encoding."
+ },
+ "VEC_IN_DIM": {
+ "value": 768,
+ "description": "dimension of the vector input."
+ },
+ "GUIDANCE_EMBED": {
+ "value": False,
+ "description": "whether to use guidance embedding."
+ },
+ "CONTEXT_IN_DIM": {
+ "value": 4096,
+ "description": "dimension of the context input."
+ },
+ "MLP_RATIO": {
+ "value": 4.0,
+ "description": "ratio of mlp hidden size to hidden size."
+ },
+ "QKV_BIAS": {
+ "value": True,
+ "description": "whether to use bias in qkv projection."
+ },
+ "DEPTH": {
+ "value": 19,
+ "description": "number of transformer blocks."
+ },
+ "DEPTH_SINGLE_BLOCKS": {
+ "value": 38,
+ "description": "number of transformer blocks in the single stream block."
+ },
+ "USE_GRAD_CHECKPOINT": {
+ "value": False,
+ "description": "whether to use gradient checkpointing."
+ },
+ "ATTN_BACKEND": {
+ "value": "pytorch",
+ "description": "backend for the transformer blocks, 'pytorch' or 'flash_attn'."
+ }
+ }
+ def __init__(
+ self,
+ cfg,
+ logger = None
+ ):
+ super().__init__(cfg, logger=logger)
+ self.in_channels = cfg.IN_CHANNELS
+ self.out_channels = cfg.get("OUT_CHANNELS", self.in_channels)
+ hidden_size = cfg.get("HIDDEN_SIZE", 1024)
+ num_heads = cfg.get("NUM_HEADS", 16)
+ axes_dim = cfg.AXES_DIM
+ theta = cfg.THETA
+ vec_in_dim = cfg.VEC_IN_DIM
+ self.guidance_embed = cfg.GUIDANCE_EMBED
+ context_in_dim = cfg.CONTEXT_IN_DIM
+ mlp_ratio = cfg.MLP_RATIO
+ qkv_bias = cfg.QKV_BIAS
+ depth = cfg.DEPTH
+ depth_single_blocks = cfg.DEPTH_SINGLE_BLOCKS
+ self.use_grad_checkpoint = cfg.get("USE_GRAD_CHECKPOINT", False)
+ self.attn_backend = cfg.get("ATTN_BACKEND", "pytorch")
+ self.cache_pretrain_model = cfg.get("CACHE_PRETRAIN_MODEL", False)
+ self.lora_model = cfg.get("DIFFUSERS_LORA_MODEL", None)
+ self.comfyui_lora_model = cfg.get("COMFYUI_LORA_MODEL", None)
+ self.swift_lora_model = cfg.get("SWIFT_LORA_MODEL", None)
+ self.blackforest_lora_model = cfg.get("BLACKFOREST_LORA_MODEL", None)
+ self.pretrain_adapter = cfg.get("PRETRAIN_ADAPTER", None)
+
+ if hidden_size % num_heads != 0:
+ raise ValueError(
+ f"Hidden size {hidden_size} must be divisible by num_heads {num_heads}"
+ )
+ pe_dim = hidden_size // num_heads
+ if sum(axes_dim) != pe_dim:
+ raise ValueError(f"Got {axes_dim} but expected positional dim {pe_dim}")
+ self.hidden_size = hidden_size
+ self.num_heads = num_heads
+ self.pe_embedder = EmbedND(dim=pe_dim, theta=theta, axes_dim= axes_dim)
+ self.img_in = nn.Linear(self.in_channels, self.hidden_size, bias=True)
+ self.time_in = MLPEmbedder(in_dim=256, hidden_dim=self.hidden_size)
+ self.vector_in = MLPEmbedder(vec_in_dim, self.hidden_size)
+ self.guidance_in = (
+ MLPEmbedder(in_dim=256, hidden_dim=self.hidden_size) if self.guidance_embed else nn.Identity()
+ )
+ self.txt_in = nn.Linear(context_in_dim, self.hidden_size)
+
+ self.double_blocks = nn.ModuleList(
+ [
+ DoubleStreamBlock(
+ self.hidden_size,
+ self.num_heads,
+ mlp_ratio=mlp_ratio,
+ qkv_bias=qkv_bias,
+ backend=self.attn_backend
+ )
+ for _ in range(depth)
+ ]
+ )
+
+ self.single_blocks = nn.ModuleList(
+ [
+ SingleStreamBlock(self.hidden_size, self.num_heads, mlp_ratio=mlp_ratio, backend=self.attn_backend)
+ for _ in range(depth_single_blocks)
+ ]
+ )
+
+ self.final_layer = LastLayer(self.hidden_size, 1, self.out_channels)
+ def merge_diffuser_lora(self, ori_sd, lora_sd, scale=1.0):
+ key_map = {
+ "single_blocks.{}.linear1.weight": {"key_list": [
+ ["transformer.single_transformer_blocks.{}.attn.to_q.lora_A.weight",
+ "transformer.single_transformer_blocks.{}.attn.to_q.lora_B.weight", [0, 3072]],
+ ["transformer.single_transformer_blocks.{}.attn.to_k.lora_A.weight",
+ "transformer.single_transformer_blocks.{}.attn.to_k.lora_B.weight", [3072, 6144]],
+ ["transformer.single_transformer_blocks.{}.attn.to_v.lora_A.weight",
+ "transformer.single_transformer_blocks.{}.attn.to_v.lora_B.weight", [6144, 9216]],
+ ["transformer.single_transformer_blocks.{}.proj_mlp.lora_A.weight",
+ "transformer.single_transformer_blocks.{}.proj_mlp.lora_B.weight", [9216, 21504]]
+ ], "num": 38},
+ "single_blocks.{}.modulation.lin.weight": {"key_list": [
+ ["transformer.single_transformer_blocks.{}.norm.linear.lora_A.weight",
+ "transformer.single_transformer_blocks.{}.norm.linear.lora_B.weight", [0, 9216]],
+ ], "num": 38},
+ "single_blocks.{}.linear2.weight": {"key_list": [
+ ["transformer.single_transformer_blocks.{}.proj_out.lora_A.weight",
+ "transformer.single_transformer_blocks.{}.proj_out.lora_B.weight", [0, 3072]],
+ ], "num": 38},
+ "double_blocks.{}.txt_attn.qkv.weight": {"key_list": [
+ ["transformer.transformer_blocks.{}.attn.add_q_proj.lora_A.weight",
+ "transformer.transformer_blocks.{}.attn.add_q_proj.lora_B.weight", [0, 3072]],
+ ["transformer.transformer_blocks.{}.attn.add_k_proj.lora_A.weight",
+ "transformer.transformer_blocks.{}.attn.add_k_proj.lora_B.weight", [3072, 6144]],
+ ["transformer.transformer_blocks.{}.attn.add_v_proj.lora_A.weight",
+ "transformer.transformer_blocks.{}.attn.add_v_proj.lora_B.weight", [6144, 9216]],
+ ], "num": 19},
+ "double_blocks.{}.img_attn.qkv.weight": {"key_list": [
+ ["transformer.transformer_blocks.{}.attn.to_q.lora_A.weight",
+ "transformer.transformer_blocks.{}.attn.to_q.lora_B.weight", [0, 3072]],
+ ["transformer.transformer_blocks.{}.attn.to_k.lora_A.weight",
+ "transformer.transformer_blocks.{}.attn.to_k.lora_B.weight", [3072, 6144]],
+ ["transformer.transformer_blocks.{}.attn.to_v.lora_A.weight",
+ "transformer.transformer_blocks.{}.attn.to_v.lora_B.weight", [6144, 9216]],
+ ], "num": 19},
+ "double_blocks.{}.img_attn.proj.weight": {"key_list": [
+ ["transformer.transformer_blocks.{}.attn.to_out.0.lora_A.weight",
+ "transformer.transformer_blocks.{}.attn.to_out.0.lora_B.weight", [0, 3072]]
+ ], "num": 19},
+ "double_blocks.{}.txt_attn.proj.weight": {"key_list": [
+ ["transformer.transformer_blocks.{}.attn.to_add_out.lora_A.weight",
+ "transformer.transformer_blocks.{}.attn.to_add_out.lora_B.weight", [0, 3072]]
+ ], "num": 19},
+ "double_blocks.{}.img_mlp.0.weight": {"key_list": [
+ ["transformer.transformer_blocks.{}.ff.net.0.proj.lora_A.weight",
+ "transformer.transformer_blocks.{}.ff.net.0.proj.lora_B.weight", [0, 12288]]
+ ], "num": 19},
+ "double_blocks.{}.img_mlp.2.weight": {"key_list": [
+ ["transformer.transformer_blocks.{}.ff.net.2.lora_A.weight",
+ "transformer.transformer_blocks.{}.ff.net.2.lora_B.weight", [0, 3072]]
+ ], "num": 19},
+ "double_blocks.{}.txt_mlp.0.weight": {"key_list": [
+ ["transformer.transformer_blocks.{}.ff_context.net.0.proj.lora_A.weight",
+ "transformer.transformer_blocks.{}.ff_context.net.0.proj.lora_B.weight", [0, 12288]]
+ ], "num": 19},
+ "double_blocks.{}.txt_mlp.2.weight": {"key_list": [
+ ["transformer.transformer_blocks.{}.ff_context.net.2.lora_A.weight",
+ "transformer.transformer_blocks.{}.ff_context.net.2.lora_B.weight", [0, 3072]]
+ ], "num": 19},
+ "double_blocks.{}.img_mod.lin.weight": {"key_list": [
+ ["transformer.transformer_blocks.{}.norm1.linear.lora_A.weight",
+ "transformer.transformer_blocks.{}.norm1.linear.lora_B.weight", [0, 18432]]
+ ], "num": 19},
+ "double_blocks.{}.txt_mod.lin.weight": {"key_list": [
+ ["transformer.transformer_blocks.{}.norm1_context.linear.lora_A.weight",
+ "transformer.transformer_blocks.{}.norm1_context.linear.lora_B.weight", [0, 18432]]
+ ], "num": 19}
+ }
+ cover_lora_keys = set()
+ cover_ori_keys = set()
+ for k, v in key_map.items():
+ key_list = v["key_list"]
+ block_num = v["num"]
+ for block_id in range(block_num):
+ for k_list in key_list:
+ if k_list[0].format(block_id) in lora_sd and k_list[1].format(block_id) in lora_sd:
+ cover_lora_keys.add(k_list[0].format(block_id))
+ cover_lora_keys.add(k_list[1].format(block_id))
+ current_weight = torch.matmul(lora_sd[k_list[0].format(block_id)].permute(1, 0),
+ lora_sd[k_list[1].format(block_id)].permute(1, 0)).permute(1, 0)
+ ori_sd[k.format(block_id)][k_list[2][0]:k_list[2][1], ...] += scale * current_weight
+ cover_ori_keys.add(k.format(block_id))
+ # lora_sd.pop(k_list[0].format(block_id))
+ # lora_sd.pop(k_list[1].format(block_id))
+ self.logger.info(f"merge_blackforest_lora loads lora'parameters lora-paras: \n"
+ f"cover-{len(cover_lora_keys)} vs total {len(lora_sd)} \n"
+ f"cover ori-{len(cover_ori_keys)} vs total {len(ori_sd)}")
+ return ori_sd
+
+ def merge_swift_lora(self, ori_sd, lora_sd, scale = 1.0):
+ have_lora_keys = {}
+ for k, v in lora_sd.items():
+ k = k[len("model."):] if k.startswith("model.") else k
+ ori_key = k.split("lora")[0] + "weight"
+ if ori_key not in ori_sd:
+ raise f"{ori_key} should in the original statedict"
+ if ori_key not in have_lora_keys:
+ have_lora_keys[ori_key] = {}
+ if "lora_A" in k:
+ have_lora_keys[ori_key]["lora_A"] = v
+ elif "lora_B" in k:
+ have_lora_keys[ori_key]["lora_B"] = v
+ else:
+ raise NotImplementedError
+ self.logger.info(f"merge_swift_lora loads lora'parameters {len(have_lora_keys)}")
+ for key, v in have_lora_keys.items():
+ current_weight = torch.matmul(v["lora_A"].permute(1, 0), v["lora_B"].permute(1, 0)).permute(1, 0)
+ ori_sd[key] += scale * current_weight
+ return ori_sd
+
+ def merge_blackforest_lora(self, ori_sd, lora_sd, scale = 1.0):
+ have_lora_keys = {}
+ cover_lora_keys = set()
+ cover_ori_keys = set()
+ for k, v in lora_sd.items():
+ if "lora" in k:
+ ori_key = k.split("lora")[0] + "weight"
+ if ori_key not in ori_sd:
+ raise f"{ori_key} should in the original statedict"
+ if ori_key not in have_lora_keys:
+ have_lora_keys[ori_key] = {}
+ if "lora_A" in k:
+ have_lora_keys[ori_key]["lora_A"] = v
+ cover_lora_keys.add(k)
+ cover_ori_keys.add(ori_key)
+ elif "lora_B" in k:
+ have_lora_keys[ori_key]["lora_B"] = v
+ cover_lora_keys.add(k)
+ cover_ori_keys.add(ori_key)
+ else:
+ if k in ori_sd:
+ ori_sd[k] = v
+ cover_lora_keys.add(k)
+ cover_ori_keys.add(k)
+ else:
+ print("unsurpport keys: ", k)
+ self.logger.info(f"merge_blackforest_lora loads lora'parameters lora-paras: \n"
+ f"cover-{len(cover_lora_keys)} vs total {len(lora_sd)} \n"
+ f"cover ori-{len(cover_ori_keys)} vs total {len(ori_sd)}")
+
+ for key, v in have_lora_keys.items():
+ current_weight = torch.matmul(v["lora_A"].permute(1, 0), v["lora_B"].permute(1, 0)).permute(1, 0)
+ # print(key, ori_sd[key].shape, current_weight.shape)
+ ori_sd[key] += scale * current_weight
+ return ori_sd
+
+ def merge_comfyui_lora(self, ori_sd, lora_sd, scale = 1.0):
+ ori_key_map = {key.replace("_", ".") : key for key in ori_sd.keys()}
+ parse_ckpt = OrderedDict()
+ for k, v in lora_sd.items():
+ if "alpha" in k:
+ continue
+ k = k.replace("lora_unet_", "").replace("_", ".")
+ map_k = ori_key_map[k.split(".lora")[0] + ".weight"]
+ if map_k not in parse_ckpt:
+ parse_ckpt[map_k] = {}
+ if "lora.up" in k:
+ parse_ckpt[map_k]["lora_up"] = v
+ elif "lora.down" in k:
+ parse_ckpt[map_k]["lora_down"] = v
+ if self.cache_pretrain_model:
+ self.lora_dict[self.comfyui_lora_model] = {}
+
+ for key, v in parse_ckpt.items():
+ current_weight = torch.matmul(v["lora_down"].permute(1, 0), v["lora_up"].permute(1, 0)).permute(1, 0)
+ self.lora_dict[self.comfyui_lora_model] = current_weight
+ ori_sd[key] += scale * current_weight
+ return ori_sd
+
+ def easy_lora_merge(self, ori_sd, lora_sd, scale = 1.0):
+ for key, v in lora_sd.items():
+ ori_sd[key] += scale * v
+ return ori_sd
+
+ def load_pretrained_model(self, pretrained_model, lora_scale = 1.0):
+ if next(self.parameters()).device.type == 'meta':
+ map_location = torch.device(we.device_id)
+ safe_device = we.device_id
+ # elif next(self.parameters()).device.type == 'cuda':
+ # map_location = torch.device(we.device_id)
+ # safe_device = we.device_id
+ else:
+ map_location = "cpu"
+ safe_device = "cpu"
+
+ if pretrained_model is not None:
+ if not hasattr(self, "ckpt"):
+ with FS.get_from(pretrained_model, wait_finish=True) as local_model:
+ if local_model.endswith('safetensors'):
+ from safetensors.torch import load_file as load_safetensors
+ ckpt = load_safetensors(local_model, device=safe_device)
+ else:
+ ckpt = torch.load(local_model, map_location=map_location, weights_only=True)
+ if "state_dict" in ckpt:
+ ckpt = ckpt["state_dict"]
+ if "model" in ckpt:
+ ckpt = ckpt["model"]["model"]
+ if self.cache_pretrain_model:
+ self.ckpt = ckpt
+ self.lora_dict = {}
+ else:
+ ckpt = self.ckpt
+
+ new_ckpt = OrderedDict()
+ for k, v in ckpt.items():
+ if k in ("img_in.weight"):
+ model_p = self.state_dict()[k]
+ if v.shape != model_p.shape:
+ expanded_state_dict_weight = torch.zeros_like(model_p, device=v.device)
+ slices = tuple(slice(0, dim) for dim in v.shape)
+ expanded_state_dict_weight[slices] = v
+ new_ckpt[k] = expanded_state_dict_weight
+ else:
+ new_ckpt[k] = v
+ else:
+ new_ckpt[k] = v
+
+
+ if self.lora_model is not None:
+ with FS.get_from(self.lora_model, wait_finish=True) as local_model:
+ if local_model.endswith('safetensors'):
+ from safetensors.torch import load_file as load_safetensors
+ lora_sd = load_safetensors(local_model, device=safe_device)
+ else:
+ lora_sd = torch.load(local_model, map_location=map_location, weights_only=True)
+ new_ckpt = self.merge_diffuser_lora(new_ckpt, lora_sd, scale=lora_scale)
+ if self.swift_lora_model is not None:
+ if not isinstance(self.swift_lora_model, list):
+ self.swift_lora_model = [(self.swift_lora_model, 1.0)]
+ for lora_model in self.swift_lora_model:
+ if isinstance(lora_model, str):
+ lora_model = (lora_model, 1.0/len(self.swift_lora_model))
+ print(lora_model)
+ self.logger.info(f"load swift lora model: {lora_model}")
+ with FS.get_from(lora_model[0], wait_finish=True) as local_model:
+ if local_model.endswith('safetensors'):
+ from safetensors.torch import load_file as load_safetensors
+ lora_sd = load_safetensors(local_model, device=safe_device)
+ else:
+ lora_sd = torch.load(local_model, map_location=map_location, weights_only=True)
+ new_ckpt = self.merge_swift_lora(new_ckpt, lora_sd, scale=lora_model[1])
+
+ if self.blackforest_lora_model is not None:
+ with FS.get_from(self.blackforest_lora_model, wait_finish=True) as local_model:
+ if local_model.endswith('safetensors'):
+ from safetensors.torch import load_file as load_safetensors
+ lora_sd = load_safetensors(local_model, device=safe_device)
+ else:
+ lora_sd = torch.load(local_model, map_location=map_location, weights_only=True)
+ new_ckpt = self.merge_blackforest_lora(new_ckpt, lora_sd, scale=lora_scale)
+
+ if self.comfyui_lora_model is not None:
+ if hasattr(self, "current_lora") and self.current_lora == self.comfyui_lora_model:
+ return
+ if hasattr(self, "lora_dict") and self.comfyui_lora_model in self.lora_dict:
+ new_ckpt = self.easy_lora_merge(new_ckpt, self.lora_dict[self.comfyui_lora_model], scale=lora_scale)
+ else:
+ with FS.get_from(self.comfyui_lora_model, wait_finish=True) as local_model:
+ if local_model.endswith('safetensors'):
+ from safetensors.torch import load_file as load_safetensors
+ lora_sd = load_safetensors(local_model, device=safe_device)
+ else:
+ lora_sd = torch.load(local_model, map_location=map_location, weights_only=True)
+ new_ckpt = self.merge_comfyui_lora(new_ckpt, lora_sd, scale=lora_scale)
+ if self.comfyui_lora_model:
+ self.current_lora = self.comfyui_lora_model
+
+
+ adapter_ckpt = {}
+ if self.pretrain_adapter is not None:
+ with FS.get_from(self.pretrain_adapter, wait_finish=True) as local_adapter:
+ if local_adapter.endswith('safetensors'):
+ from safetensors.torch import load_file as load_safetensors
+ adapter_ckpt = load_safetensors(local_adapter, device=safe_device)
+ else:
+ adapter_ckpt = torch.load(local_adapter, map_location=map_location, weights_only=True)
+ new_ckpt.update(adapter_ckpt)
+
+ missing, unexpected = self.load_state_dict(new_ckpt, strict=False, assign=True)
+ self.logger.info(
+ f'Restored from {pretrained_model} with {len(missing)} missing and {len(unexpected)} unexpected keys'
+ )
+ if len(missing) > 0:
+ self.logger.info(f'Missing Keys:\n {missing}')
+ if len(unexpected) > 0:
+ self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
+
+ def prepare_input(self, x, cond):
+ if isinstance(cond['context'], list):
+ context, y = torch.cat(cond["context"], dim=0).to(x), torch.cat(cond["y"], dim=0).to(x)
+ else:
+ context, y = cond['context'].to(x), cond['y'].to(x)
+ batch_frames, batch_frames_ids = [], []
+ for ix, shape in zip(x, cond["x_shapes"]):
+ # unpack image from sequence
+ ix = ix[:, :shape[0] * shape[1]].view(-1, shape[0], shape[1])
+ 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 = [], [], [], []
+ 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])
+ 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)
+
+ 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
+
+ def unpack(self, x: Tensor, cond: dict = None, x_seq_length: list = None) -> Tensor:
+ x_list = []
+ image_shapes = cond["x_shapes"]
+ for u, shape, seq_length in zip(x, image_shapes, x_seq_length):
+ height, width = shape
+ h, w = math.ceil(height / 2), math.ceil(width / 2)
+ u = rearrange(
+ u[seq_length-h*w:seq_length, ...],
+ "(h w) (c ph pw) -> (h ph w pw) c",
+ h=h,
+ w=w,
+ ph=2,
+ pw=2,
+ )
+ x_list.append(u)
+ x = pad_sequence(tuple(x_list), batch_first=True).permute(0, 2, 1)
+ return x
+
+ def forward(
+ self,
+ x: Tensor,
+ t: Tensor,
+ cond: dict = {},
+ guidance: Tensor | None = None,
+ gc_seg: int = 0,
+ **kwargs
+ ) -> Tensor:
+ x, x_ids, txt, txt_ids, y, mask_x, mask_txt, seq_length_list = self.prepare_input(x, cond)
+ # running on sequences img
+ vec = self.time_in(timestep_embedding(t, 256))
+ if self.guidance_embed and guidance[-1] >= 0:
+ if guidance is None:
+ raise ValueError("Didn't get guidance strength for guidance distilled model.")
+ vec = vec + self.guidance_in(timestep_embedding(guidance, 256))
+ vec = vec + self.vector_in(y)
+ ids = torch.cat((txt_ids, x_ids), dim=1)
+ pe = self.pe_embedder(ids)
+
+ mask_aside = torch.cat((mask_txt, mask_x), dim=1)
+ mask = mask_aside[:, None, :] * mask_aside[:, :, None]
+
+ kwargs = dict(
+ vec=vec,
+ pe=pe,
+ mask=mask,
+ txt_length = txt.shape[1],
+ )
+ x = torch.cat((txt, x), 1)
+ if self.use_grad_checkpoint and gc_seg >= 0:
+ x = checkpoint_sequential(
+ functions=[partial(block, **kwargs) for block in self.double_blocks],
+ segments=gc_seg if gc_seg > 0 else len(self.double_blocks),
+ input=x,
+ use_reentrant=False
+ )
+ else:
+ for block in self.double_blocks:
+ x = block(x, **kwargs)
+
+ kwargs = dict(
+ vec=vec,
+ pe=pe,
+ mask=mask,
+ )
+
+ if self.use_grad_checkpoint and gc_seg >= 0:
+ x = checkpoint_sequential(
+ functions=[partial(block, **kwargs) for block in self.single_blocks],
+ segments=gc_seg if gc_seg > 0 else len(self.single_blocks),
+ input=x,
+ use_reentrant=False
+ )
+ else:
+ for block in self.single_blocks:
+ x = block(x, **kwargs)
+ x = x[:, txt.shape[1]:, ...]
+ x = self.final_layer(x, vec) # (N, T, patch_size ** 2 * out_channels) 6 64 64
+ x = self.unpack(x, cond, seq_length_list)
+ return x
+
+ @staticmethod
+ def get_config_template():
+ return dict_to_yaml('MODEL',
+ __class__.__name__,
+ Flux.para_dict,
+ set_name=True)
+@BACKBONES.register_class()
+class ACEPlus(Flux):
+ 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, 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)
+ if len(ie) > 0:
+ ie = ie[0].squeeze(0)
+ ie_mask = torch.ones((ix.shape[0] * 4, ix.shape[1], ix.shape[2])) if ie_mask is None else ie_mask[0].squeeze(0)
+ else:
+ 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_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 = [], [], [], []
+ 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)
+ # import pdb;pdb.set_trace()
+ 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__,
+ ACEPlus.para_dict,
+ set_name=True)
\ No newline at end of file
diff --git a/modules/layers.py b/modules/layers.py
new file mode 100644
index 0000000..6c5dcce
--- /dev/null
+++ b/modules/layers.py
@@ -0,0 +1,519 @@
+# -*- coding: utf-8 -*-
+# Copyright (c) Alibaba, Inc. and its affiliates.
+from __future__ import annotations
+
+import math
+from dataclasses import dataclass
+from torch import Tensor, nn
+import torch
+from einops import rearrange, repeat
+from torch import Tensor
+from torch.nn.utils.rnn import pad_sequence
+
+try:
+ from flash_attn import (
+ flash_attn_varlen_func
+ )
+ FLASHATTN_IS_AVAILABLE = True
+except ImportError:
+ FLASHATTN_IS_AVAILABLE = False
+ flash_attn_varlen_func = None
+
+def attention(q: Tensor, k: Tensor, v: Tensor, pe: Tensor, mask: Tensor | None = None, backend = 'pytorch') -> Tensor:
+ q, k = apply_rope(q, k, pe)
+ if backend == 'pytorch':
+ if mask is not None and mask.dtype == torch.bool:
+ mask = torch.zeros_like(mask).to(q).masked_fill_(mask.logical_not(), -1e20)
+ x = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
+ # x = torch.nan_to_num(x, nan=0.0, posinf=1e10, neginf=-1e10)
+ x = rearrange(x, "B H L D -> B L (H D)")
+ elif backend == 'flash_attn':
+ # q: (B, H, L, D)
+ # k: (B, H, S, D) now L = S
+ # v: (B, H, S, D)
+ b, h, lq, d = q.shape
+ _, _, lk, _ = k.shape
+ q = rearrange(q, "B H L D -> B L H D")
+ k = rearrange(k, "B H S D -> B S H D")
+ v = rearrange(v, "B H S D -> B S H D")
+ if mask is None:
+ q_lens = torch.tensor([lq] * b, dtype=torch.int32).to(q.device, non_blocking=True)
+ k_lens = torch.tensor([lk] * b, dtype=torch.int32).to(k.device, non_blocking=True)
+ else:
+ q_lens = torch.sum(mask[:, 0, :, 0], dim=1).int()
+ k_lens = torch.sum(mask[:, 0, 0, :], dim=1).int()
+ q = torch.cat([q_v[:q_l] for q_v, q_l in zip(q, q_lens)])
+ k = torch.cat([k_v[:k_l] for k_v, k_l in zip(k, k_lens)])
+ v = torch.cat([v_v[:v_l] for v_v, v_l in zip(v, k_lens)])
+ cu_seqlens_q = torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(0, dtype=torch.int32)
+ cu_seqlens_k = torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(0, dtype=torch.int32)
+ max_seqlen_q = q_lens.max()
+ max_seqlen_k = k_lens.max()
+
+ x = flash_attn_varlen_func(
+ q,
+ k,
+ v,
+ cu_seqlens_q=cu_seqlens_q,
+ cu_seqlens_k=cu_seqlens_k,
+ max_seqlen_q=max_seqlen_q,
+ max_seqlen_k=max_seqlen_k
+ )
+ x_list = [x[cu_seqlens_q[i]:cu_seqlens_q[i+1]] for i in range(b)]
+ x = pad_sequence(tuple(x_list), batch_first=True)
+ x = rearrange(x, "B L H D -> B L (H D)")
+ else:
+ raise NotImplementedError
+ return x
+
+
+def rope(pos: Tensor, dim: int, theta: int) -> Tensor:
+ assert dim % 2 == 0
+ scale = torch.arange(0, dim, 2, dtype=torch.float64, device=pos.device) / dim
+ omega = 1.0 / (theta**scale)
+ out = torch.einsum("...n,d->...nd", pos, omega)
+ out = torch.stack([torch.cos(out), -torch.sin(out), torch.sin(out), torch.cos(out)], dim=-1)
+ out = rearrange(out, "b n d (i j) -> b n d i j", i=2, j=2)
+ return out.float()
+
+
+def apply_rope(xq: Tensor, xk: Tensor, freqs_cis: Tensor) -> tuple[Tensor, Tensor]:
+ xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
+ xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)
+ xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]
+ xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]
+ return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)
+
+class EmbedND(nn.Module):
+ def __init__(self, dim: int, theta: int, axes_dim: list[int]):
+ super().__init__()
+ self.dim = dim
+ self.theta = theta
+ self.axes_dim = axes_dim
+
+ def forward(self, ids: Tensor) -> Tensor:
+ n_axes = ids.shape[-1]
+ emb = torch.cat(
+ [rope(ids[..., i], self.axes_dim[i], self.theta) for i in range(n_axes)],
+ dim=-3,
+ )
+
+ return emb.unsqueeze(1)
+
+
+def timestep_embedding(t: Tensor, dim, max_period=10000, time_factor: float = 1000.0):
+ """
+ Create sinusoidal timestep embeddings.
+ :param t: a 1-D Tensor of N indices, one per batch element.
+ These may be fractional.
+ :param dim: the dimension of the output.
+ :param max_period: controls the minimum frequency of the embeddings.
+ :return: an (N, D) Tensor of positional embeddings.
+ """
+ t = time_factor * t
+ half = dim // 2
+ freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half).to(
+ t.device
+ )
+
+ args = t[:, None].float() * freqs[None]
+ embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
+ if dim % 2:
+ embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
+ if torch.is_floating_point(t):
+ embedding = embedding.to(t)
+ return embedding
+
+
+class MLPEmbedder(nn.Module):
+ def __init__(self, in_dim: int, hidden_dim: int):
+ super().__init__()
+ self.in_layer = nn.Linear(in_dim, hidden_dim, bias=True)
+ self.silu = nn.SiLU()
+ self.out_layer = nn.Linear(hidden_dim, hidden_dim, bias=True)
+
+ def forward(self, x: Tensor) -> Tensor:
+ return self.out_layer(self.silu(self.in_layer(x)))
+
+
+class RMSNorm(torch.nn.Module):
+ def __init__(self, dim: int):
+ super().__init__()
+ self.scale = nn.Parameter(torch.ones(dim))
+
+ def forward(self, x: Tensor):
+ x_dtype = x.dtype
+ x = x.float()
+ rrms = torch.rsqrt(torch.mean(x**2, dim=-1, keepdim=True) + 1e-6)
+ return (x * rrms).to(dtype=x_dtype) * self.scale
+
+
+class QKNorm(torch.nn.Module):
+ def __init__(self, dim: int):
+ super().__init__()
+ self.query_norm = RMSNorm(dim)
+ self.key_norm = RMSNorm(dim)
+
+ def forward(self, q: Tensor, k: Tensor, v: Tensor) -> tuple[Tensor, Tensor]:
+ q = self.query_norm(q)
+ k = self.key_norm(k)
+ return q.to(v), k.to(v)
+
+
+class SelfAttention(nn.Module):
+ def __init__(self, dim: int, num_heads: int = 8, qkv_bias: bool = False):
+ super().__init__()
+ self.num_heads = num_heads
+ head_dim = dim // num_heads
+
+ self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
+ self.norm = QKNorm(head_dim)
+ self.proj = nn.Linear(dim, dim)
+
+ def forward(self, x: Tensor, pe: Tensor, mask: Tensor | None = None) -> Tensor:
+ qkv = self.qkv(x)
+ q, k, v = rearrange(qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
+ q, k = self.norm(q, k, v)
+ x = attention(q, k, v, pe=pe, mask=mask)
+ x = self.proj(x)
+ return x
+
+class CrossAttention(nn.Module):
+ def __init__(self, dim: int, context_dim: int, num_heads: int = 8, qkv_bias: bool = False):
+ super().__init__()
+ self.num_heads = num_heads
+ head_dim = dim // num_heads
+ self.q = nn.Linear(dim, dim, bias=qkv_bias)
+ self.kv = nn.Linear(dim, context_dim * 2, bias=qkv_bias)
+ self.norm = QKNorm(head_dim)
+ self.proj = nn.Linear(dim, dim)
+
+ def forward(self, x: Tensor, context: Tensor, pe: Tensor, mask: Tensor | None = None) -> Tensor:
+ qkv = self.qkv(x)
+ q, k, v = rearrange(qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
+ q, k = self.norm(q, k, v)
+ x = attention(q, k, v, pe=pe, mask=mask)
+ x = self.proj(x)
+ return x
+
+
+@dataclass
+class ModulationOut:
+ shift: Tensor
+ scale: Tensor
+ gate: Tensor
+
+
+class Modulation(nn.Module):
+ def __init__(self, dim: int, double: bool):
+ super().__init__()
+ self.is_double = double
+ self.multiplier = 6 if double else 3
+ self.lin = nn.Linear(dim, self.multiplier * dim, bias=True)
+
+ def forward(self, vec: Tensor) -> tuple[ModulationOut, ModulationOut | None]:
+ out = self.lin(nn.functional.silu(vec))[:, None, :].chunk(self.multiplier, dim=-1)
+
+ return (
+ ModulationOut(*out[:3]),
+ ModulationOut(*out[3:]) if self.is_double else None,
+ )
+
+
+class DoubleStreamBlock(nn.Module):
+ def __init__(self, hidden_size: int, num_heads: int, mlp_ratio: float, qkv_bias: bool = False, backend = 'pytorch'):
+ super().__init__()
+
+ mlp_hidden_dim = int(hidden_size * mlp_ratio)
+ self.num_heads = num_heads
+ self.hidden_size = hidden_size
+ self.img_mod = Modulation(hidden_size, double=True)
+ self.img_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.img_attn = SelfAttention(dim=hidden_size, num_heads=num_heads, qkv_bias=qkv_bias)
+
+ self.img_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.img_mlp = nn.Sequential(
+ nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
+ nn.GELU(approximate="tanh"),
+ nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
+ )
+
+ self.backend = backend
+
+ self.txt_mod = Modulation(hidden_size, double=True)
+ self.txt_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.txt_attn = SelfAttention(dim=hidden_size, num_heads=num_heads, qkv_bias=qkv_bias)
+
+ self.txt_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.txt_mlp = nn.Sequential(
+ nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
+ nn.GELU(approximate="tanh"),
+ nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
+ )
+
+
+
+
+ def forward(self, x: Tensor, vec: Tensor, pe: Tensor, mask: Tensor = None, txt_length = None):
+ img_mod1, img_mod2 = self.img_mod(vec)
+ txt_mod1, txt_mod2 = self.txt_mod(vec)
+
+ txt, img = x[:, :txt_length], x[:, txt_length:]
+
+ # prepare image for attention
+ img_modulated = self.img_norm1(img)
+ img_modulated = (1 + img_mod1.scale) * img_modulated + img_mod1.shift
+ img_qkv = self.img_attn.qkv(img_modulated)
+ img_q, img_k, img_v = rearrange(img_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
+ img_q, img_k = self.img_attn.norm(img_q, img_k, img_v)
+ # prepare txt for attention
+ txt_modulated = self.txt_norm1(txt)
+ txt_modulated = (1 + txt_mod1.scale) * txt_modulated + txt_mod1.shift
+ txt_qkv = self.txt_attn.qkv(txt_modulated)
+ txt_q, txt_k, txt_v = rearrange(txt_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
+ txt_q, txt_k = self.txt_attn.norm(txt_q, txt_k, txt_v)
+
+ # run actual attention
+ q = torch.cat((txt_q, img_q), dim=2)
+ k = torch.cat((txt_k, img_k), dim=2)
+ v = torch.cat((txt_v, img_v), dim=2)
+ if mask is not None:
+ mask = repeat(mask, 'B L S-> B H L S', H=self.num_heads)
+ attn = attention(q, k, v, pe=pe, mask = mask, backend = self.backend)
+ txt_attn, img_attn = attn[:, : txt.shape[1]], attn[:, txt.shape[1] :]
+
+ # calculate the img bloks
+ img = img + img_mod1.gate * self.img_attn.proj(img_attn)
+ img = img + img_mod2.gate * self.img_mlp((1 + img_mod2.scale) * self.img_norm2(img) + img_mod2.shift)
+
+ # calculate the txt bloks
+ txt = txt + txt_mod1.gate * self.txt_attn.proj(txt_attn)
+ txt = txt + txt_mod2.gate * self.txt_mlp((1 + txt_mod2.scale) * self.txt_norm2(txt) + txt_mod2.shift)
+ x = torch.cat((txt, img), 1)
+ return x
+
+
+class SingleStreamBlock(nn.Module):
+ """
+ A DiT block with parallel linear layers as described in
+ https://arxiv.org/abs/2302.05442 and adapted modulation interface.
+ """
+
+ def __init__(
+ self,
+ hidden_size: int,
+ num_heads: int,
+ mlp_ratio: float = 4.0,
+ qk_scale: float | None = None,
+ backend='pytorch'
+ ):
+ super().__init__()
+ self.hidden_dim = hidden_size
+ self.num_heads = num_heads
+ head_dim = hidden_size // num_heads
+ self.scale = qk_scale or head_dim**-0.5
+
+ self.mlp_hidden_dim = int(hidden_size * mlp_ratio)
+ # qkv and mlp_in
+ self.linear1 = nn.Linear(hidden_size, hidden_size * 3 + self.mlp_hidden_dim)
+ # proj and mlp_out
+ self.linear2 = nn.Linear(hidden_size + self.mlp_hidden_dim, hidden_size)
+
+ self.norm = QKNorm(head_dim)
+
+ self.hidden_size = hidden_size
+ self.pre_norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+
+ self.mlp_act = nn.GELU(approximate="tanh")
+ self.modulation = Modulation(hidden_size, double=False)
+ self.backend = backend
+
+ def forward(self, x: Tensor, vec: Tensor, pe: Tensor, mask: Tensor = None) -> Tensor:
+ mod, _ = self.modulation(vec)
+ x_mod = (1 + mod.scale) * self.pre_norm(x) + mod.shift
+ qkv, mlp = torch.split(self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1)
+
+ q, k, v = rearrange(qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
+ q, k = self.norm(q, k, v)
+ if mask is not None:
+ mask = repeat(mask, 'B L S-> B H L S', H=self.num_heads)
+ # compute attention
+ attn = attention(q, k, v, pe=pe, mask = mask, backend=self.backend)
+ # compute activation in mlp stream, cat again and run second linear layer
+ output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
+ return x + mod.gate * output
+
+
+class DoubleStreamBlockC(DoubleStreamBlock):
+ """
+ A DiT block with parallel linear layers as described in
+ https://arxiv.org/abs/2302.05442 and adapted modulation interface.
+ """
+
+ def __init__(self, hidden_size: int, num_heads: int, mlp_ratio: float,
+ qkv_bias: bool = False, backend='pytorch',
+ abondon_cond = False):
+ super().__init__(hidden_size, num_heads, mlp_ratio,
+ qkv_bias, backend)
+ self.abondon_cond = abondon_cond
+
+ def forward(self, x: Tensor, vec: Tensor,
+ pe: Tensor, mask: Tensor = None,
+ txt_length=None,
+ uncondi_length=None,
+ uncondi_pe = None,
+ mask_uncond = None):
+ # pad_sequence(tuple(x_list), batch_first=True)
+ if self.abondon_cond:
+ x = [ix[:u_l, :] for ix, u_l in zip(x, uncondi_length)]
+ x = pad_sequence(x, batch_first=True)
+ if not x.shape[1] == pe.shape[2]:
+ pe = uncondi_pe
+ mask = mask_uncond
+ # print("double stream block", x.shape, pe.shape)
+ x = super().forward(x, vec, pe, mask, txt_length)
+ return x
+
+class SingleStreamBlockC(SingleStreamBlock):
+ """
+ A DiT block with parallel linear layers as described in
+ https://arxiv.org/abs/2302.05442 and adapted modulation interface.
+ """
+
+ def __init__(self, hidden_size: int,
+ num_heads: int,
+ mlp_ratio: float = 4.0,
+ qk_scale: float | None = None,
+ backend='pytorch',
+ abondon_cond = False):
+ super().__init__(hidden_size, num_heads, mlp_ratio,
+ qk_scale, backend)
+ self.abondon_cond = abondon_cond
+
+ def forward(self, x: Tensor, vec: Tensor, pe: Tensor, mask: Tensor = None,
+ uncondi_length = None, uncondi_pe = None, mask_uncond = None) -> Tensor:
+ if self.abondon_cond:
+ x = [ix[:u_l, :] for ix, u_l in zip(x, uncondi_length)]
+ x = pad_sequence(x, batch_first=True)
+ if not x.shape[1] == pe.shape[2]:
+ pe = uncondi_pe
+ mask = mask_uncond
+ # print("single stream block", x.shape, pe.shape)
+ x = super().forward(x, vec, pe, mask)
+ return x
+
+
+class DoubleStreamBlockD(DoubleStreamBlock):
+ """
+ A DiT block with parallel linear layers as described in
+ https://arxiv.org/abs/2302.05442 and adapted modulation interface.
+ """
+
+ def __init__(self, hidden_size: int, num_heads: int, mlp_ratio: float,
+ qkv_bias: bool = False, backend='pytorch'):
+ super().__init__(hidden_size, num_heads, mlp_ratio,
+ qkv_bias, backend)
+ mlp_hidden_dim = int(hidden_size * mlp_ratio)
+ self.edit_mod = Modulation(hidden_size, double=True)
+ self.edit_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.edit_attn = SelfAttention(dim=hidden_size, num_heads=num_heads, qkv_bias=qkv_bias)
+
+ self.edit_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.edit_mlp = nn.Sequential(
+ nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
+ nn.GELU(approximate="tanh"),
+ nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
+ )
+
+ def forward(self, x: Tensor, vec: Tensor,
+ pe: Tensor, mask: Tensor = None,
+ txt_length=None,
+ edit_length=None):
+ if edit_length is not None:
+ txt, edit, img = x[:, :txt_length], x[:, txt_length:txt_length + edit_length], x[:, txt_length + edit_length:]
+ else:
+ txt, img = x[:, :txt_length], x[:, txt_length:]
+ img_mod1, img_mod2 = self.img_mod(vec)
+ txt_mod1, txt_mod2 = self.txt_mod(vec)
+ # prepare image for attention
+ img_modulated = self.img_norm1(img)
+ img_modulated = (1 + img_mod1.scale) * img_modulated + img_mod1.shift
+ img_qkv = self.img_attn.qkv(img_modulated)
+ img_q, img_k, img_v = rearrange(img_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
+ img_q, img_k = self.img_attn.norm(img_q, img_k, img_v)
+ # prepare txt for attention
+ txt_modulated = self.txt_norm1(txt)
+ txt_modulated = (1 + txt_mod1.scale) * txt_modulated + txt_mod1.shift
+ txt_qkv = self.txt_attn.qkv(txt_modulated)
+ txt_q, txt_k, txt_v = rearrange(txt_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
+ txt_q, txt_k = self.txt_attn.norm(txt_q, txt_k, txt_v)
+
+ if edit_length is not None:
+ edit_mod1, edit_mod2 = self.edit_mod(vec)
+ # prepare edit for attention
+ edit_modulated = self.edit_norm1(edit)
+ edit_modulated = (1 + edit_mod1.scale) * edit_modulated + edit_mod1.shift
+ edit_qkv = self.edit_attn.qkv(edit_modulated)
+ edit_q, edit_k, edit_v = rearrange(edit_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
+ edit_q, edit_k = self.edit_attn.norm(edit_q, edit_k, edit_v)
+ else:
+ edit_q, edit_k, edit_v = None, None, None
+
+
+ # run actual attention
+ q = torch.cat((txt_q,) + ((edit_q,) if edit_q is not None else ()) + (img_q,), dim=2)
+ k = torch.cat((txt_k,) + ((edit_k,) if edit_k is not None else ()) + (img_k,), dim=2)
+ v = torch.cat((txt_v,) + ((edit_v,) if edit_v is not None else ()) + (img_v,), dim=2)
+ if mask is not None:
+ mask = repeat(mask, 'B L S-> B H L S', H=self.num_heads)
+ attn = attention(q, k, v, pe=pe, mask=mask, backend=self.backend)
+ if edit_length is not None:
+ txt_attn, edit_attn, img_attn = attn[:, : txt_length], attn[:, txt_length:txt_length + edit_length ], attn[:, txt_length + edit_length:]
+ else:
+ txt_attn, img_attn = attn[:, : txt_length], attn[:, txt_length:]
+
+ # calculate the img bloks
+ img = img + img_mod1.gate * self.img_attn.proj(img_attn)
+ img = img + img_mod2.gate * self.img_mlp((1 + img_mod2.scale) * self.img_norm2(img) + img_mod2.shift)
+
+ # calculate the txt bloks
+ txt = txt + txt_mod1.gate * self.txt_attn.proj(txt_attn)
+ txt = txt + txt_mod2.gate * self.txt_mlp((1 + txt_mod2.scale) * self.txt_norm2(txt) + txt_mod2.shift)
+
+ # calculate the img bloks
+ if edit_length is not None:
+ edit = edit + edit_mod1.gate * self.edit_attn.proj(edit_attn)
+ edit = edit + edit_mod2.gate * self.edit_mlp((1 + edit_mod2.scale) * self.edit_norm2(edit) + edit_mod2.shift)
+ x = torch.cat((txt, edit, img), 1)
+ else:
+ x = torch.cat((txt, img), 1)
+ return x
+
+
+class LastLayer(nn.Module):
+ def __init__(self, hidden_size: int, patch_size: int, out_channels: int):
+ super().__init__()
+ self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True)
+ self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True))
+
+ def forward(self, x: Tensor, vec: Tensor) -> Tensor:
+ shift, scale = self.adaLN_modulation(vec).chunk(2, dim=1)
+ x = (1 + scale[:, None, :]) * self.norm_final(x) + shift[:, None, :]
+ x = self.linear(x)
+ return x
+
+
+if __name__ == '__main__':
+ pe = EmbedND(dim=64, theta=10000, axes_dim=[16, 56, 56])
+
+ ix_id = torch.zeros(64 // 2, 64 // 2, 3)
+ ix_id[..., 1] = ix_id[..., 1] + torch.arange(64 // 2)[:, None]
+ ix_id[..., 2] = ix_id[..., 2] + torch.arange(64 // 2)[None, :]
+ ix_id = rearrange(ix_id, "h w c -> 1 (h w) c")
+ pos = torch.cat([ix_id, ix_id], dim = 1)
+ a = pe(pos)
+
+ b = torch.cat([pe(ix_id), pe(ix_id)], dim = 2)
+
+ print(a - b)
\ No newline at end of file
diff --git a/requirements.txt b/requirements.txt
new file mode 100644
index 0000000..dff0fa9
--- /dev/null
+++ b/requirements.txt
@@ -0,0 +1,2 @@
+scepter
+diffusers
\ No newline at end of file