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

+ +

+ Paper PDF + Project Page + + + + +
+ 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