diff --git a/README.md b/README.md index 15419b4..b99d8f8 100644 --- a/README.md +++ b/README.md @@ -206,6 +206,28 @@ export LOCAL_MODEL_PATH="ms://iic/ACE_Plus@local_editing/comfyui_local_lora16.sa python infer.py ``` +## 🚀 Train +We provide training code that allows users to train on their own data. Reference the data in 'data/train.csv' and 'data/eval.csv' to construct the training data and test data, respectively. We use '#;#' to separate fields. +The required fields include the following six, with their explanations as follows. +```angular2html +"edit_image": represents the input image for the editing task. If it is not an editing task but a reference generation, this field can be left empty. +"edit_mask": represents the input image mask for the editing task, used to specify the editing area. If it is not an editing task but rather for reference generation, this field can be left empty. +"ref_image": represents the input image for the reference image generation task; if it is a pure editing task, this field can be left empty. +"target_image": represents the generated target image and cannot be empty. +"prompt": represents the prompt for the generation task. +"data_type": represents the type of data, which can be 'portrait', 'subject', or 'local'. This field is not used in training phase. +``` + +All parameters related to training are stored in 'train_config/ace_plus_lora.yaml'. To run the training code, execute the following command. + +```bash +export FLUX_FILL_PATH="hf://black-forest-labs/FLUX.1-Fill-dev" +python run_train.py --cfg train_config/ace_plus_lora.yaml +``` + +The models trained by ACE++ can be found in ./examples/exp_example/xxxx/checkpoints/xxxx/0_SwiftLoRA/comfyui_model.safetensors. + + ## 💻 Demo We have built a GUI demo based on Gradio to help users better utilize the ACE++ model. Just execute the following command. ```bash diff --git a/__init__.py b/__init__.py index 422787a..4f6c3fc 100644 --- a/__init__.py +++ b/__init__.py @@ -1 +1 @@ -import modules \ No newline at end of file +from . import modules \ No newline at end of file diff --git a/data/eval.csv b/data/eval.csv new file mode 100644 index 0000000..1e249f6 --- /dev/null +++ b/data/eval.csv @@ -0,0 +1,6 @@ +#;##;#./assets/samples/portrait/human_1.jpg#;#./assets/samples/portrait/human_1_1.jpg#;#Maintain the facial features, A girl is wearing a neat police uniform and sporting a badge. She is smiling with a friendly and confident demeanor. The background is blurred, featuring a cartoon logo.#;#portrait +#;##;#./assets/samples/subject/subject_1.jpg#;#./assets/samples/subject/subject_1_1.jpg#;#Display the logo in a minimalist style printed in white on a matte black ceramic coffee mug, alongside a steaming cup of coffee on a cozy cafe table.#;#subject +./assets/samples/local/local_1.webp#;#./assets/samples/local/local_1_m.webp#;##;#./assets/samples/local/local_1_1.jpg#;#By referencing the mask, restore a partial image from the doodle {image} that aligns with the textual explanation: "1 white old owl".#;#local_editing +./assets/samples/application/photo_editing/1_1_edit.png#;#./assets/samples/application/photo_editing/1_1_m.png#;#./assets/samples/application/photo_editing/1_ref.png#;#./assets/samples/application/photo_editing/1_1_res.jpg#;#The item is put on the ground.#;#subject +./assets/samples/application/logo_paste/1_1_edit.png#;#./assets/samples/application/logo_paste/1_1_m.png#;#./assets/samples/application/logo_paste/1_ref.png#;#./assets/samples/application/logo_paste/1_1_res.webp#;#The logo is printed on the headphones.#;#subject +assets/samples/application/movie_poster/1_1_edit.png#;#assets/samples/application/movie_poster/1_1_m.png#;#assets/samples/application/movie_poster/1_ref.png#;#assets/samples/application/movie_poster/1_1_res.webp#;#The man is facing the camera and is smiling.#;#portrait diff --git a/data/train.csv b/data/train.csv new file mode 100644 index 0000000..1e249f6 --- /dev/null +++ b/data/train.csv @@ -0,0 +1,6 @@ +#;##;#./assets/samples/portrait/human_1.jpg#;#./assets/samples/portrait/human_1_1.jpg#;#Maintain the facial features, A girl is wearing a neat police uniform and sporting a badge. She is smiling with a friendly and confident demeanor. The background is blurred, featuring a cartoon logo.#;#portrait +#;##;#./assets/samples/subject/subject_1.jpg#;#./assets/samples/subject/subject_1_1.jpg#;#Display the logo in a minimalist style printed in white on a matte black ceramic coffee mug, alongside a steaming cup of coffee on a cozy cafe table.#;#subject +./assets/samples/local/local_1.webp#;#./assets/samples/local/local_1_m.webp#;##;#./assets/samples/local/local_1_1.jpg#;#By referencing the mask, restore a partial image from the doodle {image} that aligns with the textual explanation: "1 white old owl".#;#local_editing +./assets/samples/application/photo_editing/1_1_edit.png#;#./assets/samples/application/photo_editing/1_1_m.png#;#./assets/samples/application/photo_editing/1_ref.png#;#./assets/samples/application/photo_editing/1_1_res.jpg#;#The item is put on the ground.#;#subject +./assets/samples/application/logo_paste/1_1_edit.png#;#./assets/samples/application/logo_paste/1_1_m.png#;#./assets/samples/application/logo_paste/1_ref.png#;#./assets/samples/application/logo_paste/1_1_res.webp#;#The logo is printed on the headphones.#;#subject +assets/samples/application/movie_poster/1_1_edit.png#;#assets/samples/application/movie_poster/1_1_m.png#;#assets/samples/application/movie_poster/1_ref.png#;#assets/samples/application/movie_poster/1_1_res.webp#;#The man is facing the camera and is smiling.#;#portrait diff --git a/inference/ace_plus_diffusers.py b/inference/ace_plus_diffusers.py index 609c7e6..f3aca31 100644 --- a/inference/ace_plus_diffusers.py +++ b/inference/ace_plus_diffusers.py @@ -108,6 +108,8 @@ class ACEPlusDiffuserInference(): max_sequence_length=512, generator=generator ).images[0] + if lora_path is not None: + self.pipe.unload_lora_weights() return self.image_processor.postprocess(image, slice_w, out_w, out_h), seed diff --git a/modules/__init__.py b/modules/__init__.py index 7d1bd5d..d8bc0cb 100644 --- a/modules/__init__.py +++ b/modules/__init__.py @@ -1,2 +1,6 @@ -from .flux import Flux, ACEPlus -from .embedder import ACEHFEmbedder, T5ACEPlusClipFluxEmbedder \ No newline at end of file +from .flux import FluxMRACEPlus +from .ace_plus_dataset import ACEPlusDataset +from .ace_plus_ldm import LatentDiffusionACEPlus +from .ace_plus_solver import ACEPlusSolver +from .embedder import ACEHFEmbedder, T5ACEPlusClipFluxEmbedder +from .checkpoint import ACECheckpointHook, ACEBackwardHook \ No newline at end of file diff --git a/modules/ace_plus_dataset.py b/modules/ace_plus_dataset.py new file mode 100644 index 0000000..b67e3b6 --- /dev/null +++ b/modules/ace_plus_dataset.py @@ -0,0 +1,242 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import math +import re, io +import numpy as np +import random, torch +from PIL import Image +import torchvision.transforms as T +from collections import defaultdict +from scepter.modules.data.dataset.registry import DATASETS +from scepter.modules.data.dataset.base_dataset import BaseDataset +from scepter.modules.transform.io import pillow_convert +from scepter.modules.utils.directory import osp_path +from scepter.modules.utils.file_system import FS +from torchvision.transforms import InterpolationMode +def load_image(prefix, img_path, cvt_type=None): + if img_path is None or img_path == '': + return None + img_path = osp_path(prefix, img_path) + with FS.get_object(img_path) as image_bytes: + image = Image.open(io.BytesIO(image_bytes)) + if cvt_type is not None: + image = pillow_convert(image, cvt_type) + return image +def transform_image(image, std = 0.5, mean = 0.5): + return (image.permute(2, 0, 1)/255. - mean)/std +def transform_mask(mask): + return mask.unsqueeze(0)/255. +def ensure_src_align_target_h_mode(src_image, size, image_id, interpolation=InterpolationMode.BILINEAR): + # padding mode + H, W = size + ret_image = [] + for one_id in image_id: + edit_image = src_image[one_id] + _, eH, eW = edit_image.shape + scale = H/eH + tH, tW = H, int(eW * scale) + ret_image.append(T.Resize((tH, tW), interpolation=interpolation, antialias=True)(edit_image)) + return ret_image + +def ensure_limit_sequence(image, max_seq_len = 4096, d = 16, interpolation=InterpolationMode.BILINEAR): + # resize image for max_seq_len, while keep the aspect ratio + H, W = image.shape[-2:] + scale = min(1.0, math.sqrt(max_seq_len / ((H / d) * (W / d)))) + rH = int(H * scale) // d * d # ensure divisible by self.d + rW = int(W * scale) // d * d + # print(f"{H} {W} -> {rH} {rW}") + image = T.Resize((rH, rW), interpolation=interpolation, antialias=True)(image) + return image + +@DATASETS.register_class() +class ACEPlusDataset(BaseDataset): + para_dict = { + "DELIMITER": { + "value": "#;#", + "description": "The delimiter for records of data list." + }, + "FIELDS": { + "value": ["data_type", "edit_image", "edit_mask", "ref_image", "target_image", "prompt"], + "description": "The fields for every record." + }, + "PATH_PREFIX": { + "value": "", + "description": "The path prefix for every input image." + }, + "EDIT_TYPE_LIST": { + "value": [], + "description": "The edit type list to be trained for data list." + }, + "MAX_SEQ_LEN": { + "value": 4096, + "description": "The max sequence length for input image." + }, + "D": { + "value": 16, + "description": "Patch size for resized image." + } + } + para_dict.update(BaseDataset.para_dict) + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + delimiter = cfg.get("DELIMITER", "#;#") + fields = cfg.get("FIELDS", []) + prefix = cfg.get("PATH_PREFIX", "") + edit_type_list = cfg.get("EDIT_TYPE_LIST", []) + self.max_seq_len = cfg.get("MAX_SEQ_LEN", 4096) + self.repaiting_scale = cfg.get("REPAINTING_SCALE", 0.5) + self.d = cfg.get("D", 16) + prompt_file = cfg.DATA_LIST + self.items = self.read_data_list(delimiter, + fields, + prefix, + edit_type_list, + prompt_file) + random.shuffle(self.items) + use_num = int(cfg.get('USE_NUM', -1)) + if use_num > 0: + self.items = self.items[:use_num] + def read_data_list(self, delimiter, + fields, + prefix, + edit_type_list, + prompt_file): + with FS.get_object(prompt_file) as local_data: + rows = local_data.decode('utf-8').strip().split('\n') + items = list() + dtype_level_num = {} + for i, row in enumerate(rows): + item = {"prefix": prefix} + for key, val in zip(fields, row.split(delimiter)): + item[key] = val + edit_type = item["data_type"] + if len(edit_type_list) > 0: + for re_pattern in edit_type_list: + if re.match(re_pattern, edit_type): + items.append(item) + if edit_type not in dtype_level_num: + dtype_level_num[edit_type] = 0 + dtype_level_num[edit_type] += 1 + break + else: + items.append(item) + if edit_type not in dtype_level_num: + dtype_level_num[edit_type] = 0 + dtype_level_num[edit_type] += 1 + for edit_type in dtype_level_num: + self.logger.info(f"{edit_type} has {dtype_level_num[edit_type]} samples.") + return items + def __len__(self): + return len(self.items) + + def __getitem__(self, index): + item = self._get(index) + return self.pipeline(item) + + def _get(self, index): + # normalize + index = self.items[index%len(self)] + prefix = index.get("prefix", "") + edit_image = index.get("edit_image", "") + edit_mask = index.get("edit_mask", "") + ref_image = index.get("ref_image", "") + target_image = index.get("target_image", "") + prompt = index.get("prompt", "") + + edit_image = load_image(prefix, edit_image, cvt_type="RGB") if edit_image != "" else None + edit_mask = load_image(prefix, edit_mask, cvt_type="L") if edit_mask != "" else None + ref_image = load_image(prefix, ref_image, cvt_type="RGB") if ref_image != "" else None + target_image = load_image(prefix, target_image, cvt_type="RGB") if target_image != "" else None + assert target_image is not None + + edit_id, ref_id, src_image_list, src_mask_list = [], [], [], [] + # parse editing image + if edit_image is None: + edit_image = Image.new("RGB", target_image.size, 255) + edit_mask = Image.new("L", edit_image.size, 255) + elif edit_mask is None: + edit_mask = Image.new("L", edit_image.size, 255) + src_image_list.append(edit_image) + edit_id.append(0) + src_mask_list.append(edit_mask) + # parse reference image + if ref_image is not None: + src_image_list.append(ref_image) + ref_id.append(1) + src_mask_list.append(Image.new("L", ref_image.size, 255)) + + image = transform_image(torch.tensor(np.array(target_image).astype(np.float32))) + if edit_mask is not None: + image_mask = transform_mask(torch.tensor(np.array(edit_mask).astype(np.float32))) + else: + image_mask = Image.new("L", target_image.size, 255) + image_mask = transform_mask(torch.tensor(np.array(image_mask).astype(np.float32))) + + + src_image_list = [transform_image(torch.tensor(np.array(im).astype(np.float32))) for im in src_image_list] + src_mask_list = [transform_mask(torch.tensor(np.array(im).astype(np.float32))) for im in src_mask_list] + + # decide the repainting scale for the editing task + if len(ref_id) > 0: + repainting_scale = 1.0 + else: + repainting_scale = self.repaiting_scale + for e_i in edit_id: + src_image_list[e_i] = src_image_list[e_i] * (1 - repainting_scale * src_mask_list[e_i]) + + # use fill mode(cat img, not align) + # ensure the height of ref image is aligned with that of target image + size = image.shape[1:] + ref_image_list = ensure_src_align_target_h_mode(src_image_list, size, + image_id=ref_id, + interpolation=InterpolationMode.BILINEAR) + ref_mask_list = ensure_src_align_target_h_mode(src_mask_list, size, + image_id=ref_id, + interpolation=InterpolationMode.NEAREST_EXACT) + edit_image_list = ensure_src_align_target_h_mode(src_image_list, size, + image_id=edit_id, + interpolation=InterpolationMode.BILINEAR) + edit_mask_list = ensure_src_align_target_h_mode(src_mask_list, size, + image_id=edit_id, + interpolation=InterpolationMode.NEAREST_EXACT) + + src_image_list = [torch.cat(ref_image_list + edit_image_list, dim=-1)] + src_mask_list = [torch.cat(ref_mask_list + edit_mask_list, dim=-1)] + image = torch.cat(ref_image_list + [image], dim=-1) + image_mask = torch.cat(ref_mask_list + [image_mask], dim=-1) + + # limit max sequence length + image = ensure_limit_sequence(image, max_seq_len = self.max_seq_len, + d = self.d, interpolation=InterpolationMode.BILINEAR) + image_mask = ensure_limit_sequence(image_mask, max_seq_len = self.max_seq_len, + d = self.d, interpolation=InterpolationMode.NEAREST_EXACT) + src_image_list = [ensure_limit_sequence(i, max_seq_len = self.max_seq_len, + d = self.d, interpolation=InterpolationMode.BILINEAR) for i in src_image_list] + src_mask_list = [ensure_limit_sequence(i, max_seq_len = self.max_seq_len, + d = self.d, interpolation=InterpolationMode.NEAREST_EXACT) for i in src_mask_list] + # print(src_image_list[0].shape, src_mask_list[0].shape, image.shape, image_mask.shape) + item = { + "src_image_list": src_image_list, + "src_mask_list": src_mask_list, + "image": image, + "image_mask": image_mask, + "edit_id": edit_id, + "ref_id": ref_id, + "prompt": prompt, + "edit_key": index["edit_key"] if "edit_key" in index else "" + } + return item + + @staticmethod + def collate_fn(batch): + collect = defaultdict(list) + for sample in batch: + for k, v in sample.items(): + collect[k].append(v) + new_batch = dict() + for k, v in collect.items(): + if all([i is None for i in v]): + new_batch[k] = None + else: + new_batch[k] = v + return new_batch diff --git a/modules/ace_plus_ldm.py b/modules/ace_plus_ldm.py new file mode 100644 index 0000000..138ae09 --- /dev/null +++ b/modules/ace_plus_ldm.py @@ -0,0 +1,423 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import torch +import torch.nn.functional as F +import copy +import math +import random +from contextlib import nullcontext +from einops import rearrange +from scepter.modules.model.network.ldm import LatentDiffusion +from scepter.modules.model.registry import MODELS, DIFFUSIONS, BACKBONES, LOSSES, TOKENIZERS, EMBEDDERS +from scepter.modules.model.utils.basic_utils import check_list_of_list, to_device, pack_imagelist_into_tensor, \ + limit_batch_data, unpack_tensor_into_imagelist, count_params, disabled_train +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we + +@MODELS.register_class() +class LatentDiffusionACEPlus(LatentDiffusion): + para_dict = LatentDiffusion.para_dict + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.guide_scale = cfg.get('GUIDE_SCALE', 1.0) + + def init_params(self): + self.parameterization = self.cfg.get('PARAMETERIZATION', 'rf') + assert self.parameterization in [ + 'eps', 'x0', 'v', 'rf' + ], 'currently only supporting "eps" and "x0" and "v" and "rf"' + + diffusion_cfg = self.cfg.get("DIFFUSION", None) + assert diffusion_cfg is not None + if self.cfg.have("WORK_DIR"): + diffusion_cfg.WORK_DIR = self.cfg.WORK_DIR + self.diffusion = DIFFUSIONS.build(diffusion_cfg, logger=self.logger) + + self.pretrained_model = self.cfg.get('PRETRAINED_MODEL', None) + self.ignore_keys = self.cfg.get('IGNORE_KEYS', []) + + self.model_config = self.cfg.DIFFUSION_MODEL + self.first_stage_config = self.cfg.FIRST_STAGE_MODEL + self.cond_stage_config = self.cfg.COND_STAGE_MODEL + self.tokenizer_config = self.cfg.get('TOKENIZER', None) + self.loss_config = self.cfg.get('LOSS', None) + + self.scale_factor = self.cfg.get('SCALE_FACTOR', 0.18215) + self.size_factor = self.cfg.get('SIZE_FACTOR', 16) + self.default_n_prompt = self.cfg.get('DEFAULT_N_PROMPT', '') + self.default_n_prompt = '' if self.default_n_prompt is None else self.default_n_prompt + self.p_zero = self.cfg.get('P_ZERO', 0.0) + self.train_n_prompt = self.cfg.get('TRAIN_N_PROMPT', '') + if self.default_n_prompt is None: + self.default_n_prompt = '' + if self.train_n_prompt is None: + self.train_n_prompt = '' + self.use_ema = self.cfg.get('USE_EMA', False) + self.model_ema_config = self.cfg.get('DIFFUSION_MODEL_EMA', None) + + def construct_network(self): + # embedding_context = torch.device("meta") if self.model_config.get("PRETRAINED_MODEL", None) else nullcontext() + # with embedding_context: + self.model = BACKBONES.build(self.model_config, logger=self.logger).to(torch.bfloat16) + self.logger.info('all parameters:{}'.format(count_params(self.model))) + if self.use_ema: + if self.model_ema_config: + self.model_ema = BACKBONES.build(self.model_ema_config, + logger=self.logger) + else: + self.model_ema = copy.deepcopy(self.model) + self.model_ema = self.model_ema.eval() + for param in self.model_ema.parameters(): + param.requires_grad = False + if self.loss_config: + self.loss = LOSSES.build(self.loss_config, logger=self.logger) + if self.tokenizer_config is not None: + self.tokenizer = TOKENIZERS.build(self.tokenizer_config, + logger=self.logger) + if self.first_stage_config: + self.first_stage_model = MODELS.build(self.first_stage_config, + logger=self.logger) + self.first_stage_model = self.first_stage_model.eval() + self.first_stage_model.train = disabled_train + for param in self.first_stage_model.parameters(): + param.requires_grad = False + else: + self.first_stage_model = None + if self.tokenizer_config is not None: + self.cond_stage_config.KWARGS = { + 'vocab_size': self.tokenizer.vocab_size + } + if self.cond_stage_config == '__is_unconditional__': + print( + f'Training {self.__class__.__name__} as an unconditional model.' + ) + self.cond_stage_model = None + else: + model = EMBEDDERS.build(self.cond_stage_config, logger=self.logger) + self.cond_stage_model = model.eval().requires_grad_(False) + self.cond_stage_model.train = disabled_train + + @torch.no_grad() + def encode_first_stage(self, x, **kwargs): + def run_one_image(u): + zu = self.first_stage_model.encode(u) + if isinstance(zu, (tuple, list)): + zu = zu[0] + return zu + + z = [run_one_image(u.unsqueeze(0) if u.dim() == 3 else u) for u in x] + return z + + @torch.no_grad() + def decode_first_stage(self, z): + return [self.first_stage_model.decode(zu) for zu in z] + def noise_sample(self, num_samples, h, w, seed, dtype=torch.bfloat16): + noise = torch.randn( + num_samples, + 16, + # allow for packing + 2 * math.ceil(h / 16), + 2 * math.ceil(w / 16), + device=we.device_id, + dtype=dtype, + generator=torch.Generator(device=we.device_id).manual_seed(seed), + ) + return noise + def resize_func(self, x, size): + if x is None: return x + return F.interpolate(x.unsqueeze(0), size = size, mode='nearest-exact') + def parse_ref_and_edit(self, src_image, + src_image_mask, + text_embedding, + #text_mask, + edit_id): + edit_image = [] + edit_mask = [] + ref_image = [] + ref_mask = [] + ref_context = [] + ref_y = [] + ref_id = [] + txt = [] + txt_y = [] + for sample_id, (one_src, one_src_mask, + one_text_embedding, + one_text_y, + # one_text_mask, + one_edit_id) in enumerate(zip(src_image, + src_image_mask, + text_embedding["context"], + text_embedding["y"], + #text_mask, + edit_id) + ): + ref_id.append([i for i in range(len(one_src))]) + if hasattr(self, "ref_cond_stage_model") and self.ref_cond_stage_model: + ref_image.append(self.ref_cond_stage_model.encode_list([((i + 1.0) / 2.0 * 255).type(torch.uint8) for i in one_src])) + else: + ref_image.append(one_src) + ref_mask.append(one_src_mask) + # process edit image & edit image mask + current_edit_image = to_device([one_src[i] for i in one_edit_id], strict=False) + current_edit_image = [v.squeeze(0) for v in self.encode_first_stage(current_edit_image)] + current_edit_image_mask = to_device([one_src_mask[i] for i in one_edit_id], strict=False) + current_edit_image_mask = [self.reshape_func(m).squeeze(0) for m in current_edit_image_mask] + + edit_image.append(current_edit_image) + edit_mask.append(current_edit_image_mask) + ref_context.append(one_text_embedding[:len(ref_id[-1])]) + ref_y.append(one_text_y[:len(ref_id[-1])]) + if not sum(len(src_) for src_ in src_image) > 0: + ref_image = None + ref_context = None + ref_y = None + for sample_id, (one_text_embedding, one_text_y) in enumerate(zip(text_embedding["context"], + text_embedding["y"])): + txt.append(one_text_embedding[-1].squeeze(0)) + txt_y.append(one_text_y[-1]) + return { + "edit": edit_image, + "edit_mask": edit_mask, + "edit_id": edit_id, + "ref_context": ref_context, + "ref_y": ref_y, + "context": txt, + "y": txt_y, + "ref_x": ref_image, + "ref_mask": ref_mask, + "ref_id": ref_id + } + + + def reshape_func(self, mask): + mask = mask.to(torch.bfloat16) + mask = mask.view((-1, mask.shape[-2], mask.shape[-1])) + mask = rearrange( + mask, + "c (h ph) (w pw) -> c (ph pw) h w", + ph=8, + pw=8, + ) + return mask + + def forward_train(self, + src_image_list =[], + src_mask_list =[], + edit_id=[], + image=None, + image_mask=None, + noise=None, + prompt=[], + **kwargs): + ''' + Args: + src_image: list of list of src_image + src_image_mask: list of list of src_image_mask + image: target image + image_mask: target image mask + noise: default is None, generate automaticly + ref_prompt: list of list of text + prompt: list of text + **kwargs: + Returns: + ''' + assert check_list_of_list(src_image_list) and check_list_of_list(src_mask_list) + assert self.cond_stage_model is not None + + gc_seg = kwargs.pop("gc_seg", []) + gc_seg = int(gc_seg[0]) if len(gc_seg) > 0 else 0 + align = kwargs.pop("align", []) + prompt_ = [[pp] if isinstance(pp, str) else pp for pp in prompt] + if len(align) < 1: align = [0] * len(prompt_) + context = getattr(self.cond_stage_model, 'encode_list_of_list')(prompt_) + guide_scale = self.guide_scale + if guide_scale is not None: + guide_scale = torch.full((len(prompt_),), guide_scale, device=we.device_id) + else: + guide_scale = None + # image and image_mask + # print("is list of list", check_list_of_list(image)) + if check_list_of_list(image): + image = [to_device(ix) for ix in image] + x_start = [self.encode_first_stage(ix, **kwargs) for ix in image] + noise = [[torch.randn_like(ii) for ii in ix] for ix in x_start] + x_start = [torch.cat(ix, dim=-1) for ix in x_start] + noise = [torch.cat(ix, dim=-1) for ix in noise] + + noise, _ = pack_imagelist_into_tensor(noise) + + image_mask = [to_device(im, strict=False) for im in image_mask] + x_mask = [[self.reshape_func(i).squeeze(0) for i in im] if im is not None else [None] * len(ix) for ix, im in zip(image, image_mask)] + x_mask = [torch.cat(im, dim=-1) for im in x_mask] + else: + image = to_device(image) + x_start = self.encode_first_stage(image, **kwargs) + image_mask = to_device(image_mask, strict=False) + x_mask = [self.reshape_func(i).squeeze(0) for i in image_mask] if image_mask is not None else [None] * len( + image) + loss_mask, _ = pack_imagelist_into_tensor( + tuple(torch.ones_like(ix, dtype=torch.bool, device=ix.device) for ix in x_start)) + x_start, x_shapes = pack_imagelist_into_tensor(x_start) + context['x_shapes'] = x_shapes + context['align'] = align + # process image mask + + context['x_mask'] = x_mask + ref_edit_context = self.parse_ref_and_edit(src_image_list, src_mask_list, context, edit_id) + context.update(ref_edit_context) + + teacher_context = copy.deepcopy(context) + teacher_context["context"] = torch.cat(teacher_context["context"], dim=0) + teacher_context["y"] = torch.cat(teacher_context["y"], dim=0) + loss = self.diffusion.loss(x_0=x_start, + model=self.model, + model_kwargs={"cond": context, + "gc_seg": gc_seg, + "guidance": guide_scale}, + noise=noise, + reduction='none', + **kwargs) + loss = loss[loss_mask].mean() + ret = {'loss': loss, 'probe_data': {'prompt': prompt}} + return ret + + @torch.no_grad() + def forward_test(self, + src_image_list=[], + src_mask_list=[], + edit_id=[], + image=None, + image_mask=None, + prompt=[], + sampler='flow_euler', + sample_steps=20, + seed=2023, + guide_scale=3.5, + guide_rescale=0.0, + show_process=False, + log_num=-1, + **kwargs): + outputs = self.forward_editing( + src_image_list=src_image_list, + src_mask_list=src_mask_list, + edit_id=edit_id, + image=image, + image_mask=image_mask, + prompt=prompt, + sampler=sampler, + sample_steps=sample_steps, + seed=seed, + guide_scale=guide_scale, + guide_rescale=guide_rescale, + show_process=show_process, + log_num=log_num, + **kwargs + ) + return outputs + + @torch.no_grad() + def forward_editing(self, + src_image_list=[], + src_mask_list=[], + edit_id=[], + image=None, + image_mask=None, + prompt=[], + sampler='flow_euler', + sample_steps=20, + seed=2023, + guide_scale=3.5, + log_num=-1, + **kwargs + ): + # gc_seg is unused + prompt, image, image_mask, src_image, src_image_mask, edit_id = limit_batch_data( + [prompt, image, image_mask, src_image_list, src_mask_list, edit_id], log_num) + assert check_list_of_list(src_image) and check_list_of_list(src_image_mask) + assert self.cond_stage_model is not None + align = kwargs.pop("align", []) + prompt_ = [[pp] if isinstance(pp, str) else pp for pp in prompt] + if len(align) < 1: align = [0] * len(prompt_) + context = getattr(self.cond_stage_model, 'encode_list_of_list')(prompt_) + guide_scale = guide_scale or self.guide_scale + if guide_scale is not None: + guide_scale = torch.full((len(prompt),), guide_scale, device=we.device_id) + else: + guide_scale = None + # image and image_mask + seed = seed if seed >= 0 else random.randint(0, 2 ** 32 - 1) + if image is not None: + if check_list_of_list(image): + image = [torch.cat(ix, dim=-1) for ix in image] + image_mask = [torch.cat(im, dim=-1) for im in image_mask] + noise = [self.noise_sample(1, ix.shape[1], ix.shape[2], seed) for ix in image] + else: + height, width = kwargs.pop("height"), kwargs.pop("width") + noise = [self.noise_sample(1, height, width, seed) for _ in prompt] + noise, x_shapes = pack_imagelist_into_tensor(noise) + context['x_shapes'] = x_shapes + context['align'] = align + # process image mask + image_mask = to_device(image_mask, strict=False) + x_mask = [self.reshape_func(i).squeeze(0) for i in image_mask] + context['x_mask'] = x_mask + ref_edit_context = self.parse_ref_and_edit(src_image, src_image_mask, context, edit_id) + context.update(ref_edit_context) + # UNet use input n_prompt + # model = self.model_ema if self.use_ema and self.eval_ema else self.model + # import pdb;pdb.set_trace() + model = self.model + embedding_context = model.no_sync if isinstance(model, torch.distributed.fsdp.FullyShardedDataParallel) \ + else nullcontext + with embedding_context(): + samples = self.diffusion.sample( + noise=noise, + sampler=sampler, + model=self.model, + model_kwargs={"cond": context, "guidance": guide_scale, "gc_seg": -1 + }, + steps=sample_steps, + show_progress=True, + guide_scale=guide_scale, + return_intermediate=None, + **kwargs).float() + samples = unpack_tensor_into_imagelist(samples, x_shapes) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + x_samples = self.decode_first_stage(samples) + outputs = list() + for i in range(len(prompt)): + rec_img = torch.clamp((x_samples[i].float() + 1.0) / 2.0, min=0.0, max=1.0) + rec_img = rec_img.squeeze(0) + edit_imgs, edit_img_masks = [], [] + if src_image is not None and src_image[i] is not None: + if src_image_mask[i] is None: + src_image_mask[i] = [None] * len(src_image[i]) + for edit_img, edit_mask in zip(src_image[i], src_image_mask[i]): + edit_img = torch.clamp((edit_img.float() + 1.0) / 2.0, min=0.0, max=1.0) + edit_imgs.append(edit_img.squeeze(0)) + if edit_mask is None: + edit_mask = torch.ones_like(edit_img[[0], :, :]) + edit_img_masks.append(edit_mask) + one_tup = { + 'reconstruct_image': rec_img, + 'instruction': prompt[i], + 'edit_image': edit_imgs if len(edit_imgs) > 0 else None, + 'edit_mask': edit_img_masks if len(edit_imgs) > 0 else None + } + if image is not None: + if image_mask is None: + image_mask = [None] * len(image) + ori_img = torch.clamp((image[i] + 1.0) / 2.0, min=0.0, max=1.0) + one_tup['target_image'] = ori_img.squeeze(0) + one_tup['target_mask'] = image_mask[i] if image_mask[i] is not None else torch.ones_like( + ori_img[[0], :, :]) + outputs.append(one_tup) + return outputs + + @staticmethod + def get_config_template(): + return dict_to_yaml('MODEL', + __class__.__name__, + LatentDiffusionACEPlus.para_dict, + set_name=True) + diff --git a/modules/ace_plus_solver.py b/modules/ace_plus_solver.py new file mode 100644 index 0000000..90d3fc8 --- /dev/null +++ b/modules/ace_plus_solver.py @@ -0,0 +1,164 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import numpy as np +import torch +from scepter.modules.solver import LatentDiffusionSolver +from scepter.modules.solver.registry import SOLVERS +from scepter.modules.utils.data import transfer_data_to_cuda +from scepter.modules.utils.distribute import we +from scepter.modules.utils.probe import ProbeData +from tqdm import tqdm +@SOLVERS.register_class() +class ACEPlusSolver(LatentDiffusionSolver): + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.probe_prompt = cfg.get("PROBE_PROMPT", None) + self.probe_hw = cfg.get("PROBE_HW", []) + @torch.no_grad() + def run_eval(self): + self.eval_mode() + self.before_all_iter(self.hooks_dict[self._mode]) + all_results = [] + for batch_idx, batch_data in tqdm( + enumerate(self.datas[self._mode].dataloader)): + self.before_iter(self.hooks_dict[self._mode]) + if self.sample_args: + batch_data.update(self.sample_args.get_lowercase_dict()) + with torch.autocast(device_type='cuda', + enabled=self.use_amp, + dtype=self.dtype): + results = self.run_step_eval(transfer_data_to_cuda(batch_data), + batch_idx, + step=self.total_iter, + rank=we.rank) + all_results.extend(results) + self.after_iter(self.hooks_dict[self._mode]) + log_data, log_label = self.save_results(all_results) + self.register_probe({'eval_label': log_label}) + self.register_probe({ + 'eval_image': + ProbeData(log_data, + is_image=True, + build_html=True, + build_label=log_label) + }) + self.after_all_iter(self.hooks_dict[self._mode]) + + @torch.no_grad() + def run_test(self): + self.test_mode() + self.before_all_iter(self.hooks_dict[self._mode]) + all_results = [] + for batch_idx, batch_data in tqdm( + enumerate(self.datas[self._mode].dataloader)): + self.before_iter(self.hooks_dict[self._mode]) + if self.sample_args: + batch_data.update(self.sample_args.get_lowercase_dict()) + with torch.autocast(device_type='cuda', + enabled=self.use_amp, + dtype=self.dtype): + results = self.run_step_eval(transfer_data_to_cuda(batch_data), + batch_idx, + step=self.total_iter, + rank=we.rank) + all_results.extend(results) + self.after_iter(self.hooks_dict[self._mode]) + log_data, log_label = self.save_results(all_results) + self.register_probe({'test_label': log_label}) + self.register_probe({ + 'test_image': + ProbeData(log_data, + is_image=True, + build_html=True, + build_label=log_label) + }) + + self.after_all_iter(self.hooks_dict[self._mode]) + + def save_results(self, results): + log_data, log_label = [], [] + for result in results: + ret_images, ret_labels = [], [] + edit_image = result.get('edit_image', None) + edit_mask = result.get('edit_mask', None) + if edit_image is not None: + for i, edit_img in enumerate(result['edit_image']): + if edit_img is None: + continue + ret_images.append((edit_img.permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8)) + ret_labels.append(f'edit_image{i}; ') + if edit_mask is not None: + ret_images.append((edit_mask[i].permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8)) + ret_labels.append(f'edit_mask{i}; ') + + target_image = result.get('target_image', None) + target_mask = result.get('target_mask', None) + if target_image is not None: + ret_images.append((target_image.permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8)) + ret_labels.append(f'target_image; ') + if target_mask is not None: + ret_images.append((target_mask.permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8)) + ret_labels.append(f'target_mask; ') + teacher_image = result.get('image', None) + if teacher_image is not None: + ret_images.append((teacher_image.permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8)) + ret_labels.append(f"teacher_image") + reconstruct_image = result.get('reconstruct_image', None) + if reconstruct_image is not None: + ret_images.append((reconstruct_image.permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8)) + ret_labels.append(f"{result['instruction']}") + log_data.append(ret_images) + log_label.append(ret_labels) + return log_data, log_label + @property + def probe_data(self): + if not we.debug and self.mode == 'train': + batch_data = transfer_data_to_cuda(self.current_batch_data[self.mode]) + self.eval_mode() + with torch.autocast(device_type='cuda', + enabled=self.use_amp, + dtype=self.dtype): + batch_data['log_num'] = self.log_train_num + batch_data.update(self.sample_args.get_lowercase_dict()) + results = self.run_step_eval(batch_data) + self.train_mode() + log_data, log_label = self.save_results(results) + self.register_probe({ + 'train_image': + ProbeData(log_data, + is_image=True, + build_html=True, + build_label=log_label) + }) + self.register_probe({'train_label': log_label}) + if self.probe_prompt: + self.eval_mode() + all_results = [] + for prompt in self.probe_prompt: + with torch.autocast(device_type='cuda', + enabled=self.use_amp, + dtype=self.dtype): + batch_data = { + "prompt": [[prompt]], + "image": [torch.zeros(3, self.probe_hw[0], self.probe_hw[1])], + "image_mask": [torch.ones(1, self.probe_hw[0], self.probe_hw[1])], + "src_image_list": [[]], + "src_mask_list": [[]], + "edit_id": [[]], + "height": self.probe_hw[0], + "width": self.probe_hw[1] + } + batch_data.update(self.sample_args.get_lowercase_dict()) + results = self.run_step_eval(batch_data) + all_results.extend(results) + self.train_mode() + log_data, log_label = self.save_results(all_results) + self.register_probe({ + 'probe_image': + ProbeData(log_data, + is_image=True, + build_html=True, + build_label=log_label) + }) + + return super(LatentDiffusionSolver, self).probe_data diff --git a/modules/checkpoint.py b/modules/checkpoint.py new file mode 100644 index 0000000..8dcb4cf --- /dev/null +++ b/modules/checkpoint.py @@ -0,0 +1,135 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import os, torch +import os.path as osp +import warnings +from collections import OrderedDict +from safetensors.torch import save_file +from scepter.modules.solver.hooks import CheckpointHook, BackwardHook +from scepter.modules.solver.hooks.registry import HOOKS +from scepter.modules.utils.config import dict_to_yaml +from scepter.modules.utils.distribute import we +from scepter.modules.utils.file_system import FS + +_DEFAULT_CHECKPOINT_PRIORITY = 300 + +def convert_to_comfyui_lora(ori_sd, prefix = "lora_unet"): + new_ckpt = OrderedDict() + for k,v in ori_sd.items(): + new_k = k.replace(".lora_A.0_SwiftLoRA.", ".lora_down.").replace(".lora_B.0_SwiftLoRA.", ".lora_up.") + new_k = prefix + "_" + new_k.split(".lora")[0].replace("model.", "").replace(".", "_") + ".lora" + new_k.split(".lora")[1] + alpha_k = new_k.split(".lora")[0] + ".alpha" + new_ckpt[new_k] = v + if "lora_up" in new_k: + alpha = v.shape[-1] + elif "lora_down" in new_k: + alpha = v.shape[0] + new_ckpt[alpha_k] = torch.tensor(float(alpha)).to(v) + return new_ckpt + +@HOOKS.register_class() +class ACECheckpointHook(CheckpointHook): + """ Checkpoint resume or save hook. + Args: + interval (int): Save interval, by epoch. + save_best (bool): Save the best checkpoint by a metric key, default is False. + save_best_by (str): How to get the best the checkpoint by the metric key, default is ''. + + means the higher the best (default). + - means the lower the best. + E.g. +acc@1, -err@1, acc@5(same as +acc@5) + """ + + def __init__(self, cfg, logger=None): + super(ACECheckpointHook, self).__init__(cfg, logger=logger) + + def after_iter(self, solver): + super().after_iter(solver) + if solver.total_iter != 0 and ( + (solver.total_iter + 1) % self.interval == 0 + or solver.total_iter == solver.max_steps - 1): + from swift import SwiftModel + if isinstance(solver.model, SwiftModel) or ( + hasattr(solver.model, 'module') + and isinstance(solver.model.module, SwiftModel)): + save_path = osp.join( + solver.work_dir, + 'checkpoints/{}-{}'.format(self.save_name_prefix, + solver.total_iter + 1)) + if we.rank == 0: + tuner_model = os.path.join(save_path, '0_SwiftLoRA', 'adapter_model.bin') + save_model = os.path.join(save_path, '0_SwiftLoRA', 'comfyui_model.safetensors') + if FS.exists(tuner_model): + with FS.get_from(tuner_model) as local_file: + swift_lora_sd = torch.load(local_file, weights_only=True) + safetensor_lora_sd = convert_to_comfyui_lora(swift_lora_sd) + with FS.put_to(save_model) as local_file: + save_file(safetensor_lora_sd, local_file) + @staticmethod + def get_config_template(): + return dict_to_yaml('hook', + __class__.__name__, + ACECheckpointHook.para_dict, + set_name=True) + +@HOOKS.register_class() +class ACEBackwardHook(BackwardHook): + def grad_clip(self, optimizer): + for params_group in optimizer.param_groups: + train_params = [] + for param in params_group['params']: + if param.requires_grad: + train_params.append(param) + # print(len(train_params), self.gradient_clip) + torch.nn.utils.clip_grad_norm_(parameters=train_params, + max_norm=self.gradient_clip) + + def after_iter(self, solver): + if solver.optimizer is not None and solver.is_train_mode: + if solver.loss is None: + warnings.warn( + 'solver.loss should not be None in train mode, remember to call solver._reduce_scalar()!' + ) + return + if solver.scaler is not None: + solver.scaler.scale(solver.loss / + self.accumulate_step).backward() + self.current_step += 1 + # Suppose profiler run after backward, so we need to set backward_prev_step + # as the previous one step before the backward step + if self.current_step % self.accumulate_step == 0: + solver.scaler.unscale_(solver.optimizer) + if self.gradient_clip > 0: + self.grad_clip(solver.optimizer) + self.profile(solver) + solver.scaler.step(solver.optimizer) + solver.scaler.update() + solver.optimizer.zero_grad() + else: + (solver.loss / self.accumulate_step).backward() + self.current_step += 1 + # Suppose profiler run after backward, so we need to set backward_prev_step + # as the previous one step before the backward step + if self.current_step % self.accumulate_step == 0: + if self.gradient_clip > 0: + self.grad_clip(solver.optimizer) + self.profile(solver) + solver.optimizer.step() + solver.optimizer.zero_grad() + if solver.lr_scheduler: + if self.current_step % self.accumulate_step == 0: + solver.lr_scheduler.step() + if self.current_step % self.accumulate_step == 0: + setattr(solver, 'backward_step', True) + self.current_step = 0 + else: + setattr(solver, 'backward_step', False) + solver.loss = None + if self.empty_cache_step > 0 and solver.total_iter % self.empty_cache_step == 0: + torch.cuda.empty_cache() + + @staticmethod + def get_config_template(): + return dict_to_yaml('hook', + __class__.__name__, + ACEBackwardHook.para_dict, + set_name=True) diff --git a/modules/embedder.py b/modules/embedder.py index f1beece..f01346d 100644 --- a/modules/embedder.py +++ b/modules/embedder.py @@ -1,18 +1,17 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +# This file contains code that is adapted from +# https://github.com/black-forest-labs/flux.git 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) + basic_clean, canonicalize, 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: @@ -21,169 +20,6 @@ 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 = { diff --git a/modules/flux.py b/modules/flux.py index 6d097a3..7ac764c 100644 --- a/modules/flux.py +++ b/modules/flux.py @@ -1,6 +1,10 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. -import math, torch +# This file contains code that is adapted from +# https://github.com/black-forest-labs/flux.git +import math +import torch +from torch import Tensor, nn from collections import OrderedDict from functools import partial from einops import rearrange, repeat @@ -9,90 +13,82 @@ 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) - +from torch.nn.utils.rnn import pad_sequence +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." + 'IN_CHANNELS': { + 'value': 64, + 'description': "model's input channels." }, - "OUT_CHANNELS": { - "value": 64, - "description": "model's output channels." + 'OUT_CHANNELS': { + 'value': 64, + 'description': "model's output channels." }, - "HIDDEN_SIZE": { - "value": 1024, - "description": "model's hidden size." + 'HIDDEN_SIZE': { + 'value': 1024, + 'description': "model's hidden size." }, - "NUM_HEADS": { - "value": 16, - "description": "number of heads in the transformer." + '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." + 'AXES_DIM': { + 'value': [16, 56, 56], + 'description': 'dimensions of the axes of the positional encoding.' }, - "THETA": { - "value": 10_000, - "description": "theta for positional encoding." + 'THETA': { + 'value': 10_000, + 'description': 'theta for positional encoding.' }, - "VEC_IN_DIM": { - "value": 768, - "description": "dimension of the vector input." + 'VEC_IN_DIM': { + 'value': 768, + 'description': 'dimension of the vector input.' }, - "GUIDANCE_EMBED": { - "value": False, - "description": "whether to use guidance embedding." + 'GUIDANCE_EMBED': { + 'value': False, + 'description': 'whether to use guidance embedding.' }, - "CONTEXT_IN_DIM": { - "value": 4096, - "description": "dimension of the context input." + '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." + '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." + 'QKV_BIAS': { + 'value': True, + 'description': 'whether to use bias in qkv projection.' }, - "DEPTH": { - "value": 19, - "description": "number of transformer blocks." + 'DEPTH': { + 'value': 19, + 'description': 'number of transformer blocks.' }, - "DEPTH_SINGLE_BLOCKS": { - "value": 38, - "description": "number of transformer blocks in the single stream block." + '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'." + 'USE_GRAD_CHECKPOINT': { + 'value': False, + 'description': 'whether to use gradient checkpointing.' } } - def __init__( - self, - cfg, - logger = None - ): + + 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) + 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 @@ -117,16 +113,17 @@ class Flux(BaseModel): ) pe_dim = hidden_size // num_heads if sum(axes_dim) != pe_dim: - raise ValueError(f"Got {axes_dim} but expected positional 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.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.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( @@ -150,6 +147,28 @@ class Flux(BaseModel): ) self.final_layer = LastLayer(self.hidden_size, 1, self.out_channels) + + def prepare_input(self, x, context, y, x_shape=None): + # x.shape [6, 16, 16, 16] target is [6, 16, 768, 1360] + bs, c, h, w = x.shape + x = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=2, pw=2) + x_id = torch.zeros(h // 2, w // 2, 3) + x_id[..., 1] = x_id[..., 1] + torch.arange(h // 2)[:, None] + x_id[..., 2] = x_id[..., 2] + torch.arange(w // 2)[None, :] + x_ids = repeat(x_id, "h w c -> b (h w) c", b=bs) + txt_ids = torch.zeros(bs, context.shape[1], 3) + return x, x_ids.to(x), context.to(x), txt_ids.to(x), y.to(x), h, w + + def unpack(self, x: Tensor, height: int, width: int) -> Tensor: + return rearrange( + x, + "b (h w) (c ph pw) -> b c (h ph) (w pw)", + h=math.ceil(height/2), + w=math.ceil(width/2), + ph=2, + pw=2, + ) + def merge_diffuser_lora(self, ori_sd, lora_sd, scale=1.0): key_map = { "single_blocks.{}.linear1.weight": {"key_list": [ @@ -261,6 +280,7 @@ class Flux(BaseModel): 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() @@ -329,9 +349,6 @@ class Flux(BaseModel): 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" @@ -438,6 +455,71 @@ class Flux(BaseModel): if len(unexpected) > 0: self.logger.info(f'\nUnexpected Keys:\n {unexpected}') + def forward( + self, + x: Tensor, + t: Tensor, + cond: dict = {}, + guidance: Tensor | None = None, + gc_seg: int = 0 + ) -> Tensor: + x, x_ids, txt, txt_ids, y, h, w = self.prepare_input(x, cond["context"], cond["y"]) + # running on sequences img + x = self.img_in(x) + vec = self.time_in(timestep_embedding(t, 256)) + if self.guidance_embed: + 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) + txt = self.txt_in(txt) + ids = torch.cat((txt_ids, x_ids), dim=1) + pe = self.pe_embedder(ids) + kwargs = dict( + vec=vec, + pe=pe, + 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, + ) + + 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, h, w) + return x + + @staticmethod + def get_config_template(): + return dict_to_yaml('BACKBONE', + __class__.__name__, + Flux.para_dict, + set_name=True) +@BACKBONES.register_class() +class FluxMR(Flux): 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) @@ -562,10 +644,10 @@ class Flux(BaseModel): def get_config_template(): return dict_to_yaml('MODEL', __class__.__name__, - Flux.para_dict, + FluxMR.para_dict, set_name=True) @BACKBONES.register_class() -class ACEPlus(Flux): +class FluxMRACEPlus(FluxMR): def __init__(self, cfg, logger = None): super().__init__(cfg, logger) def prepare_input(self, x, cond): @@ -577,8 +659,10 @@ class ACEPlus(Flux): 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) + ie = [iie.squeeze(0) for iie in ie] + ie_mask = [torch.ones((ix.shape[0] * 4, ix.shape[1], ix.shape[2])) if iime is None else iime.squeeze(0) for iime in ie_mask] + ie = torch.cat(ie, dim=-1) + ie_mask = torch.cat(ie_mask, dim=-1) else: ie, ie_mask = torch.zeros_like(ix).to(x), torch.ones_like(imask).to(x) ix = torch.cat([ix, ie, ie_mask], dim=0) @@ -606,7 +690,6 @@ class ACEPlus(Flux): 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)): @@ -628,5 +711,5 @@ class ACEPlus(Flux): def get_config_template(): return dict_to_yaml('MODEL', __class__.__name__, - ACEPlus.para_dict, - set_name=True) \ No newline at end of file + FluxMRACEPlus.para_dict, + set_name=True) diff --git a/modules/layers.py b/modules/layers.py index 6c5dcce..348e754 100644 --- a/modules/layers.py +++ b/modules/layers.py @@ -1,5 +1,7 @@ # -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. +# This file contains code that is adapted from +# https://github.com/black-forest-labs/flux.git from __future__ import annotations import math diff --git a/requirements.txt b/requirements.txt index e9ce089..105fe57 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,3 +1,7 @@ -scepter +huggingface_hub diffusers -gradio>=4.44.1 \ No newline at end of file +transformers +torch>=2.4.1 +xformers>=0.0.27.post2 +gradio>=4.44.1 +scepter \ No newline at end of file diff --git a/run_train.py b/run_train.py new file mode 100644 index 0000000..6d5b64b --- /dev/null +++ b/run_train.py @@ -0,0 +1,70 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import argparse +import importlib +import os +import sys +from datetime import datetime +sys.dont_write_bytecode = True +from scepter.modules.solver.registry import SOLVERS +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 + +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) + +def run_task(cfg): + std_logger = get_logger(name='scepter') + solver = SOLVERS.build(cfg.SOLVER, logger=std_logger) + solver.set_up_pre() + solver.set_up() + if we.rank == 0: + FS.put_object_from_local_file(cfg.args.cfg_file, os.path.join(solver.work_dir, "train.yaml")) + if cfg.args.stage == "train": + solver.solve() + elif cfg.args.stage == "eval": + solver.run_eval() + + +def update_config(cfg): + if hasattr(cfg.args, 'learning_rate') and cfg.args.learning_rate: + print( + f'learning_rate change from {cfg.SOLVER.OPTIMIZER.LEARNING_RATE} to {cfg.args.learning_rate}' + ) + cfg.SOLVER.OPTIMIZER.LEARNING_RATE = float(cfg.args.learning_rate) + if hasattr(cfg.args, 'max_steps') and cfg.args.max_steps: + print( + f'max_steps change from {cfg.SOLVER.MAX_STEPS} to {cfg.args.max_steps}' + ) + cfg.SOLVER.MAX_STEPS = int(cfg.args.max_steps) + cfg.SOLVER.WORK_DIR = os.path.join(cfg.SOLVER.WORK_DIR, "{0:%Y%m%d%H%M%S}".format(datetime.now())) + return cfg + + +if __name__ == '__main__': + parser = argparse.ArgumentParser(description='Argparser for Scepter:\n') + parser.add_argument( + "--stage", + dest="stage", + help="Running stage!", + default="train", + choices=["train", "eval"] + ) + parser.add_argument('--learning_rate', + dest='learning_rate', + help='The learning rate for our network!', + default=None) + parser.add_argument('--max_steps', + dest='max_steps', + help='The max steps for training!', + default=None) + + cfg = Config(load=True, parser_ins=parser) + cfg = update_config(cfg) + we.init_env(cfg, logger=None, fn=run_task) diff --git a/train_config/ace_plus_lora.yaml b/train_config/ace_plus_lora.yaml new file mode 100644 index 0000000..8080f94 --- /dev/null +++ b/train_config/ace_plus_lora.yaml @@ -0,0 +1,279 @@ +ENV: + BACKEND: nccl + SEED: 1999 +SOLVER: + # NAME DESCRIPTION: TYPE: default: 'LatentUfitSolver' + NAME: ACEPlusSolver + # MAX_STEPS DESCRIPTION: The total steps for training. TYPE: int default: 100000 + MAX_STEPS: 100000 + # USE_AMP DESCRIPTION: Use amp to surpport mix precision or not, default is False. TYPE: bool default: False + USE_AMP: True + # DTYPE DESCRIPTION: The precision for training. TYPE: str default: 'float32' + DTYPE: bfloat16 + ENABLE_GRADSCALER: False + # USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False + USE_FAIRSCALE: False + USE_ORIG_PARAMS: True + USE_FSDP: True # lora use ddp(USE_FSDP=False), else use fsdp(USE_FSDP=True) + # LOAD_MODEL_ONLY DESCRIPTION: Only load the model rather than the optimizer and schedule, default is False. TYPE: bool default: False + LOAD_MODEL_ONLY: False + # RESUME_FROM DESCRIPTION: Resume from some state of training! TYPE: str default: '' + RESUME_FROM: + # WORK_DIR DESCRIPTION: Save dir of the training log or model. TYPE: str default: '' + WORK_DIR: ./examples/exp_example/ + # LOG_FILE DESCRIPTION: Save log path. TYPE: str default: '' + LOG_FILE: std_log.txt + # LOG_TRAIN_NUM DESCRIPTION: The number samples used to log in training phase. TYPE: int default: -1 + LOG_TRAIN_NUM: 16 + # FSDP_REDUCE_DTYPE DESCRIPTION: The dtype of reduce in FSDP. TYPE: str default: 'float16' + FSDP_REDUCE_DTYPE: float32 + # FSDP_BUFFER_DTYPE DESCRIPTION: The dtype of buffer in FSDP. TYPE: str default: 'float16' + FSDP_BUFFER_DTYPE: float32 + # FSDP_SHARD_MODULES DESCRIPTION: The modules to be sharded in FSDP. TYPE: list default: ['model'] + FSDP_SHARD_MODULES: + - MODULE: 'model.model' + FSDP_GROUP: [ 'single_blocks', 'double_blocks'] + - MODULE: 'cond_stage_model.t5_model.hf_module.encoder' + FSDP_GROUP: [ 'block' ] # + SAVE_MODULES: [ 'model'] # + TRAIN_MODULES: ['model'] + + # + FILE_SYSTEM: + - NAME: AliyunOssFs + ENDPOINT: http://oss-cn-wulanchabu-internal.aliyuncs.com + BUCKET: visionai-wlcb + OSS_AK: ${OSS_AK} + OSS_SK: ${OSS_SK} + TEMP_DIR: ${TEMP_DIR} + # + MODEL: + NAME: LatentDiffusionACEPlus + PARAMETERIZATION: rf + TIMESTEPS: 1000 + GUIDE_SCALE: 1.0 + PRETRAINED_MODEL: + IGNORE_KEYS: [ ] + USE_EMA: False + EVAL_EMA: False + SIZE_FACTOR: 8 + DIFFUSION: + NAME: DiffusionFluxRF + PREDICTION_TYPE: raw + NOISE_NORM: True + # NOISE_SCHEDULER DESCRIPTION: TYPE: default: '' + NOISE_SCHEDULER: + NAME: FlowMatchFluxShiftScheduler + SHIFT: False + PRE_T_SAMPLE: True + PRE_T_SAMPLE_FOLD: 1 + SIGMOID_SCALE: 1 + BASE_SHIFT: 0.5 + MAX_SHIFT: 1.15 + SAMPLER_SCHEDULER: + NAME: FlowMatchFluxShiftScheduler + SHIFT: True + PRE_T_SAMPLE: False + SIGMOID_SCALE: 1 + BASE_SHIFT: 0.5 + MAX_SHIFT: 1.15 + + # + DIFFUSION_MODEL: + # NAME DESCRIPTION: TYPE: default: 'Flux' + NAME: FluxMRACEPlus + PRETRAINED_MODEL: ${FLUX_FILL_PATH}/flux1-fill-dev.safetensors + # IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64 + IN_CHANNELS: 384 + # OUT_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64 + OUT_CHANNELS: 64 + # HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024 + HIDDEN_SIZE: 3072 + REDUX_DIM: 1152 + # NUM_HEADS DESCRIPTION: number of heads in the transformer. TYPE: int default: 16 + NUM_HEADS: 24 + # AXES_DIM DESCRIPTION: dimensions of the axes of the positional encoding. TYPE: list default: [16, 56, 56] + AXES_DIM: [ 16, 56, 56 ] + # THETA DESCRIPTION: theta for positional encoding. TYPE: int default: 10000 + THETA: 10000 + # VEC_IN_DIM DESCRIPTION: dimension of the vector input. TYPE: int default: 768 + VEC_IN_DIM: 768 + # GUIDANCE_EMBED DESCRIPTION: whether to use guidance embedding. TYPE: bool default: False + GUIDANCE_EMBED: True + # CONTEXT_IN_DIM DESCRIPTION: dimension of the context input. TYPE: int default: 4096 + CONTEXT_IN_DIM: 4096 + # MLP_RATIO DESCRIPTION: ratio of mlp hidden size to hidden size. TYPE: float default: 4.0 + MLP_RATIO: 4.0 + # QKV_BIAS DESCRIPTION: whether to use bias in qkv projection. TYPE: bool default: True + QKV_BIAS: True + # DEPTH DESCRIPTION: number of transformer blocks. TYPE: int default: 19 + DEPTH: 19 + # DEPTH_SINGLE_BLOCKS DESCRIPTION: number of transformer blocks in the single stream block. TYPE: int default: 38 + DEPTH_SINGLE_BLOCKS: 38 + ATTN_BACKEND: flash_attn + + # + FIRST_STAGE_MODEL: + NAME: AutoencoderKLFlux + EMBED_DIM: 16 + PRETRAINED_MODEL: ${FLUX_FILL_PATH}/ae.safetensors + IGNORE_KEYS: [ ] + BATCH_SIZE: 8 + USE_CONV: False + SCALE_FACTOR: 0.3611 + SHIFT_FACTOR: 0.1159 + # + ENCODER: + NAME: Encoder + CH: 128 + OUT_CH: 3 + NUM_RES_BLOCKS: 2 + IN_CHANNELS: 3 + ATTN_RESOLUTIONS: [ ] + CH_MULT: [ 1, 2, 4, 4 ] + Z_CHANNELS: 16 + DOUBLE_Z: True + DROPOUT: 0.0 + RESAMP_WITH_CONV: True + # + DECODER: + NAME: Decoder + CH: 128 + OUT_CH: 3 + NUM_RES_BLOCKS: 2 + IN_CHANNELS: 3 + ATTN_RESOLUTIONS: [ ] + CH_MULT: [ 1, 2, 4, 4 ] + Z_CHANNELS: 16 + DROPOUT: 0.0 + RESAMP_WITH_CONV: True + GIVE_PRE_END: False + TANH_OUT: False + # + COND_STAGE_MODEL: + # NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder' + NAME: T5ACEPlusClipFluxEmbedder + # T5_MODEL DESCRIPTION: TYPE: default: '' + T5_MODEL: + # NAME DESCRIPTION: TYPE: default: 'HFEmbedder' + NAME: ACEHFEmbedder + # HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None + HF_MODEL_CLS: T5EncoderModel + # MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None + MODEL_PATH: ${FLUX_FILL_PATH}/text_encoder_2/ + # HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None + HF_TOKENIZER_CLS: T5Tokenizer + # TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None + TOKENIZER_PATH: ${FLUX_FILL_PATH}/tokenizer_2/ + ADDED_IDENTIFIER: [ '','{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ] + # MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77 + MAX_LENGTH: 512 + # OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state' + OUTPUT_KEY: last_hidden_state + # D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16' + D_TYPE: bfloat16 + # BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False + BATCH_INFER: False + CLEAN: whitespace + # CLIP_MODEL DESCRIPTION: TYPE: default: '' + CLIP_MODEL: + # NAME DESCRIPTION: TYPE: default: 'HFEmbedder' + NAME: ACEHFEmbedder + # HF_MODEL_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None + HF_MODEL_CLS: CLIPTextModel + # MODEL_PATH DESCRIPTION: model folder path TYPE: NoneType default: None + MODEL_PATH: ${FLUX_FILL_PATH}/text_encoder/ + # HF_TOKENIZER_CLS DESCRIPTION: huggingface cls in transfomer TYPE: NoneType default: None + HF_TOKENIZER_CLS: CLIPTokenizer + # TOKENIZER_PATH DESCRIPTION: tokenizer folder path TYPE: NoneType default: None + TOKENIZER_PATH: ${FLUX_FILL_PATH}/tokenizer/ + # MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77 + MAX_LENGTH: 77 + # OUTPUT_KEY DESCRIPTION: output key TYPE: str default: 'last_hidden_state' + OUTPUT_KEY: pooler_output + # D_TYPE DESCRIPTION: dtype TYPE: str default: 'bfloat16' + D_TYPE: bfloat16 + # BATCH_INFER DESCRIPTION: batch infer TYPE: bool default: False + BATCH_INFER: True + CLEAN: whitespace + TUNER: + # THE LORA PARAMETERS + - NAME: SwiftLoRA + R: 64 + LORA_ALPHA: 64 + LORA_DROPOUT: 0.0 + BIAS: "none" + TARGET_MODULES: "(model.double_blocks.*(.qkv|.img_mlp.0|.img_mlp.2|.txt_mlp.0|.txt_mlp.2|.proj|.img_mod.lin|.txt_mod.lin))|(model.single_blocks.*(.linear1|.linear2|.modulation.lin))$" + # + SAMPLE_ARGS: + SAMPLE_STEPS: 28 + SAMPLER: flow_euler + SEED: 42 + IMAGE_SIZE: [ 1024, 1024 ] + GUIDE_SCALE: 50 + + LR_SCHEDULER: + NAME: StepAnnealingLR + WARMUP_STEPS: 0 + TOTAL_STEPS: 100000 + DECAY_MODE: 'cosine' + # + OPTIMIZER: + NAME: AdamW + LEARNING_RATE: 1e-3 + BETAS: [ 0.9, 0.999 ] + EPS: 1e-6 + WEIGHT_DECAY: 1e-2 + AMSGRAD: False + # + TRAIN_DATA: + NAME: ACEPlusDataset + MODE: train + DATA_LIST: data/train.csv + DELIMITER: "#;#" + # input_image, input_mask, input_reference_image, target_image, instruction, task_type + FIELDS: ["edit_image", "edit_mask", "ref_image", "target_image", "prompt", "data_type"] + PATH_PREFIX: "" + EDIT_TYPE_LIST: [] + MAX_SEQ_LEN: 2048 + D: 16 + PIN_MEMORY: True + BATCH_SIZE: 1 + NUM_WORKERS: 4 + SAMPLER: + NAME: LoopSampler + + EVAL_DATA: + NAME: ACEPlusDataset + MODE: eval + DATA_LIST: data/train.csv + DELIMITER: "#;#" + # input_image, input_mask, input_reference_image, target_image, instruction, task_type + FIELDS: [ "edit_image", "edit_mask", "ref_image", "target_image", "prompt", "data_type" ] + PATH_PREFIX: "" + EDIT_TYPE_LIST: [ ] + MAX_SEQ_LEN: 2048 + D: 16 + PIN_MEMORY: True + BATCH_SIZE: 1 + NUM_WORKERS: 4 + + TRAIN_HOOKS: + - NAME: ACEBackwardHook + GRADIENT_CLIP: 1.0 + PRIORITY: 10 + - NAME: LogHook + LOG_INTERVAL: 20 + - NAME: ACECheckpointHook + INTERVAL: 250 + PRIORITY: 200 + DISABLE_SNAPSHOT: True + - NAME: ProbeDataHook + PROB_INTERVAL: 50 + PRIORITY: 0 + - NAME: TensorboardLogHook + LOG_INTERVAL: 50 + EVAL_HOOKS: + - NAME: ProbeDataHook + PROB_INTERVAL: 50 + PRIORITY: 0 \ No newline at end of file