surpport the training for ace++.

This commit is contained in:
皓童
2025-01-16 18:06:28 +08:00
parent b4d446933d
commit 7a371e31ef
16 changed files with 1526 additions and 248 deletions
+22
View File
@@ -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
+1 -1
View File
@@ -1 +1 @@
import modules
from . import modules
+6
View File
@@ -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
Can't render this file because it contains an unexpected character in line 3 and column 236.
+6
View File
@@ -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
Can't render this file because it contains an unexpected character in line 3 and column 236.
+2
View File
@@ -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
+6 -2
View File
@@ -1,2 +1,6 @@
from .flux import Flux, ACEPlus
from .embedder import ACEHFEmbedder, T5ACEPlusClipFluxEmbedder
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
+242
View File
@@ -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
+423
View File
@@ -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)
+164
View File
@@ -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
+135
View File
@@ -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)
+3 -167
View File
@@ -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 = {
+159 -76
View File
@@ -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)
FluxMRACEPlus.para_dict,
set_name=True)
+2
View File
@@ -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
+6 -2
View File
@@ -1,3 +1,7 @@
scepter
huggingface_hub
diffusers
gradio>=4.44.1
transformers
torch>=2.4.1
xformers>=0.0.27.post2
gradio>=4.44.1
scepter
+70
View File
@@ -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)
+279
View File
@@ -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: [ '<img>','{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