Compare commits

..
17 Commits
Author SHA1 Message Date
jiangzeyinzi adda36e39d Merge pull request #66 from yaosheng216/patch-7
Update model_node.py
2024-11-26 14:58:52 +08:00
jiangzeyinzi 1da1864993 Merge pull request #65 from yaosheng216/patch-6
Update parameter_node.py
2024-11-26 14:58:36 +08:00
Great 5eac362325 Update model_node.py 2024-11-26 14:57:20 +08:00
Great 9ae3bca43d Update parameter_node.py 2024-11-26 14:55:29 +08:00
jiangzeyinzi ca034ef765 Merge pull request #64 from modelscope/v1.3.0_dev
V1.3.0 dev
2024-11-26 12:13:08 +08:00
皓童 2a29446d45 modify ace inference and ace yaml 2024-11-25 14:24:54 +08:00
maochaojie cd33b4ab15 Merge branch 'v1.3.0_dev' of https://github.com/modelscope/scepter into v1.3.0_dev 2024-11-21 15:42:05 +08:00
maochaojie 7d7943fed3 modify yaml and workflow 2024-11-21 15:41:45 +08:00
jiangzeyinzi d48b2f110f Merge pull request #63 from yaosheng216/patch-5
Update model_node.py
2024-11-20 13:26:33 +08:00
Great 82486adf38 Update model_node.py 2024-11-20 13:24:28 +08:00
maochaojie a683061c6f upgrade from 1.2.0 to 1.3.0 2024-11-19 19:20:02 +08:00
mcj 4ec0492897 Merge pull request #62 from modelscope/v1.2.0_dev
update chatbot example
2024-11-07 20:30:51 +08:00
mcj 02c0ba9757 Merge pull request #61 from modelscope/v1.2.0_dev
update instr
2024-11-04 17:14:27 +08:00
jiangzeyinzi 82132ff3a1 Merge pull request #60 from modelscope/v1.2.0_dev
add instruction
2024-11-04 14:38:03 +08:00
mcj 73984c4f9e Merge pull request #59 from modelscope/v1.2.0_dev
update readme
2024-11-02 06:48:03 +08:00
jiangzeyinzi edb46a615c Merge pull request #58 from modelscope/v1.2.0_dev
V1.2.0 dev
2024-11-01 21:30:20 +08:00
jiangzeyinzi d9b207cf5b Merge pull request #55 from modelscope/v1.2.0_dev
V1.2.0 dev
2024-11-01 16:53:16 +08:00
89 changed files with 10627 additions and 1266 deletions
+30 -13
View File
@@ -18,7 +18,11 @@ SCEPTER offers 3 core components:
## 🎉 News
- [🔥🔥🔥2024.10]: We are pleased to announce the release of the code for [ACE](https://arxiv.org/abs/2410.00086), supporting Customized Training / Comfy UI Workflow / gradio-based ChatBot Interface. The detailed documents can be found at [ACE repo](https://github.com/ali-vilab/ACE.git).
- [🔥🔥🔥2024.11]: We're excited to announce the upcoming release of the [ACE-0.6b-1024px](https://huggingface.co/scepter-studio/ACE-0.6B-1024px) model,
which significantly enhances image generation quality compared with [ACE-0.6b-512px](https://huggingface.co/scepter-studio/ACE-0.6B-512px). The detailed documents can be found at [ACE repo](https://github.com/ali-vilab/ACE.git).
At the same time, based on the editing results of ACE, combined with the powerful text-to-image capabilities of the [FLUX-dev](https://huggingface.co/black-forest-labs/FLUX.1-dev) model through SDEdit as an image quality refiner, the quality of image editing can be further enhanced.
- [🔥2024.11]: Supports video files, video annotation, caption translation in data management, and inference & training of the [CogVideoX](https://arxiv.org/abs/2408.06072).
- [2024.10]: We are pleased to announce the release of the code for [ACE](https://arxiv.org/abs/2410.00086), supporting Customized Training / Comfy UI Workflow / gradio-based ChatBot Interface.
- [2024.10]: Support for inference and tuning with [FLUX](https://huggingface.co/black-forest-labs/FLUX.1-dev), as well as for building [ComfyUI](https://github.com/comfyanonymous/ComfyUI) workflows using this framework.
- [2024.09]: We introduce **ACE**, an **A**ll-round **C**reator and **E**ditor adept at executing a diverse array of image editing tasks tailored to your specifications. Built upon the cutting-edge Diffusion Transformer architecture, ACE has been extensively trained on a comprehensive dataset to seamlessly interpret and execute any natural language instruction. For further information, please consult the [project page](https://ali-vilab.github.io/ace-page/).
- [2024.07]: Support the inference and training of open-source generative models based on the [DiT](https://arxiv.org/abs/2212.09748) architecture, such as [SD3](https://arxiv.org/pdf/2403.03206) and [PixArt](https://arxiv.org/abs/2310.00426).
@@ -32,19 +36,25 @@ SCEPTER offers 3 core components:
- [2023.12]: We release [🪄SCEPTER](https://github.com/modelscope/scepter/) library.
## 🖼 Gallery for Recent Works
### ACE
## 🪄ACE
ACE is a unified foundational model framework that supports a wide range of visual generation tasks. By defining CU for unifying multi-modal inputs across different tasks and incorporating long-context CU, we introduce historical contextual information into visual generation tasks, paving the way for ChatGPT-like dialog systems in visual generation.
[![Watch the demo](https://ali-vilab.github.io/ace-page/static/images/tasks.png)](https://ali-vilab.github.io/ace-page/)
#### ACE Training
### ACE Models
| **Model** | **Status** |
|:----------------:|:---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------:|
| ACE-0.6B-512px | [![Demo link](https://img.shields.io/badge/Demo-ACE_Chat-purple)](https://huggingface.co/spaces/scepter-studio/ACE-Chat)<br>[![ModelScope link](https://img.shields.io/badge/ModelScope-Model-blue)](https://www.modelscope.cn/models/iic/ACE-0.6B-512px) [![HuggingFace link](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Model-yellow)](https://huggingface.co/scepter-studio/ACE-0.6B-512px) |
| ACE-0.6B-1024px | [![Demo link](https://img.shields.io/badge/Demo-ACE_Refiner_Chat-purple)](https://huggingface.co/spaces/scepter-studio/ACE-Refiner-Chat)<br>[![ModelScope link](https://img.shields.io/badge/ModelScope-Model-blue)](https://www.modelscope.cn/models/iic/ACE-0.6B-1024px) [![HuggingFace link](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Model-yellow)](https://huggingface.co/scepter-studio/ACE-0.6B-1024px) | |
| ACE-12B-FLUX-dev | Coming Soon |
### ACE Training
We offer a demonstration training YAML that enables the end-to-end training of ACE using a toy dataset. For a comprehensive overview of the hyperparameter configurations, please consult `scepter/methods/edit/dit_ace_0.6b_512.yaml`.
##### Prepare datasets
#### Prepare datasets
Please find the dataset class located in `scepter/modules/data/dataset/ms_dataset.py`,
designed to facilitate end-to-end training using an open-source toy dataset.
@@ -52,7 +62,7 @@ Download a dataset zip file from [modelscope](https://www.modelscope.cn/models/i
Should you wish to prepare your own datasets, we recommend consulting `scepter/modules/data/dataset/ms_dataset.py` for detailed guidance on the required data format.
##### Prepare initial weight
#### Prepare initial weight
The ACE checkpoint has been uploaded to both ModelScope and HuggingFace platforms:
* [ModelScope](https://www.modelscope.cn/models/iic/ACE-0.6B-512px)
* [HuggingFace](https://huggingface.co/scepter-studio/ACE-0.6B-512px)
@@ -60,22 +70,25 @@ The ACE checkpoint has been uploaded to both ModelScope and HuggingFace platform
In the provided training YAML configuration, we have designated the Modelscope URL as the default checkpoint URL. Should you wish to transition to Hugging Face, you can effortlessly achieve this by modifying the PRETRAINED_MODEL value within the YAML file (replace the prefix "ms://iic" to "hf://scepter-studio").
##### Start training
#### Start training
You can easily start training procedure by executing the following command:
```bash
# ACE-0.6B-512px
PYTHONPATH=. python scepter/tools/run_train.py --cfg scepter/methods/edit/dit_ace_0.6b_512.yaml
# ACE-0.6B-1024px
PYTHONPATH=. python scepter/tools/run_train.py --cfg scepter/methods/edit/dit_ace_0.6b_1024.yaml
```
#### ACE Chat Bot
### ACE Chat Bot
We have developed a chatbot interface utilizing Gradio, designed to convert user input in natural language into visually captivating images that align semantically with the specified instructions. You can easily access this functionality by launching Scepter Studio with the following command:
```bash
PYTHONPATH=. python scepter/tools/webui.py --cfg scepter/methods/studio/scepter_ui.yaml --language zh
PYTHONPATH=. python scepter/tools/webui.py --cfg scepter/methods/studio/scepter_ui.yaml --language zh --tab chatbot
```
Upon starting, you will find a "ChatBot" tab within the Gradio application, which serves as a chat-based interface to handle any requests related to image editing or generation.
#### ACE ComfyUI Workflow
### ACE ComfyUI Workflow
![Workflow](https://github.com/ali-vilab/ace-page/raw/main/assets/comfyui/ace_example.jpg)
@@ -108,6 +121,8 @@ Upon starting, you will find a "ChatBot" tab within the Gradio application, whic
</tbody>
</table>
## 🖼 Gallery for Recent Works
### FLUX Tuners
<table><tbody>
@@ -258,18 +273,20 @@ We deploy a work studio on Modelscope that includes only the inference tab, plea
## ⚙️️ ComfyUI Workflow
### Launch
We support the use of all models in the ComfyUI Workflow through the following methods:
Manually install by moving custom_nodes to ComfyUI.
1) Automatic installation directly via the ComfyUI Manager by searching for the **ComfyUI-Scepter** node.
2) Manually install by moving custom_nodes from Scepter to ComfyUI.
```shell
git clone https://github.com/modelscope/scepter.git
cd path/to/scepter
pip install -e .
cp -r path/to/scepter/workflow/ path/to/ComfyUI/custom_nodes/ComfyUI-Scepter
cd path/to/ComfyUI
python main.py
```
In addition, we also support installation and usage through the ComfyUI Manager.
**Note**: You can use the nodes by dragging the sample images into ComfyUI. Additionally, our nodes can automatically pull models from ModelScope or HuggingFace by selecting the *model_source* field, or you can place the already downloaded models in a local path.
## 🔍 Learn More
+3 -2
View File
@@ -2,7 +2,7 @@ albumentations
beautifulsoup4
bezier
einops
modelscope
modelscope[framework]
ms-swift
numpy
open_clip_torch
@@ -12,6 +12,7 @@ oss2>=2.15.0
pycocotools
pyyaml>=5.3.1
scikit-image
scikit-learn
sentencepiece
torchsde
transformers
scikit-learn
+4 -3
View File
@@ -1,4 +1,5 @@
git+https://github.com/cocodataset/panopticapi.git
torch==2.0.1
torchvision==0.15.2
xformers==0.0.21
torch==2.4.1
torchvision==.19.1
flash-attn==2.5.8
xformers==0.0.28
+1 -1
View File
@@ -1,5 +1,5 @@
bitsandbytes
gradio==4.44.1
gradio
gradio_imageslider
imagehash
psutil
+161
View File
@@ -0,0 +1,161 @@
ENV:
BACKEND: nccl
SEED: 2024
#
SOLVER:
NAME: ACESolver
RESUME_FROM:
LOAD_MODEL_ONLY: True
USE_FSDP: False
SHARDING_STRATEGY:
USE_AMP: True
DTYPE: float16
CHANNELS_LAST: True
MAX_STEPS: 500
MAX_EPOCHS: -1
NUM_FOLDS: 1
ACCU_STEP: 1
EVAL_INTERVAL: 50
RESCALE_LR: False
#
WORK_DIR: ./cache/save_data/ace_0.6b_1024
LOG_FILE: std_log.txt
#
FILE_SYSTEM:
- NAME: "HuggingfaceFs"
TEMP_DIR: ./cache/cache_data
- NAME: "LocalFs"
TEMP_DIR: ./cache/cache_data
- NAME: "ModelscopeFs"
TEMP_DIR: ./cache/cache_data
#
MODEL:
NAME: LatentDiffusionACE
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
DECODER_BIAS: 0.5
DEFAULT_N_PROMPT:
USE_EMA: True
EVAL_EMA: False
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
USE_TEXT_POS_EMBEDDINGS: True
#
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: eps
MIN_SNR_GAMMA:
NOISE_SCHEDULER:
NAME: LinearScheduler
NUM_TIMESTEPS: 1000
BETA_MIN: 0.0001
BETA_MAX: 0.02
#
DIFFUSION_MODEL:
NAME: ACE
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/dit/ace_0.6b_1024px.pth
IGNORE_KEYS: [ ]
PATCH_SIZE: 2
IN_CHANNELS: 4
HIDDEN_SIZE: 1152
DEPTH: 28
NUM_HEADS: 16
MLP_RATIO: 4.0
PRED_SIGMA: True
DROP_PATH: 0.0
WINDOW_DIZE: 0
Y_CHANNELS: 4096
MAX_SEQ_LEN: 4096
QK_NORM: True
USE_GRAD_CHECKPOINT: True
ATTENTION_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/vae/vae.bin
IGNORE_KEYS: []
#
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: 4
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: 4
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/text_encoder/t5-v1_1-xxl/
TOKENIZER_PATH: ms://iic/ACE-0.6B-1024px@models/tokenizer/t5-v1_1-xxl
LENGTH: 120
T5_DTYPE: bfloat16
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
CLEAN: whitespace
USE_GRAD: False
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 20
GUIDE_SCALE: 4.5
GUIDE_RESCALE: 0.5
#
OPTIMIZER:
NAME: AdamW
LEARNING_RATE: 1e-7
EPS: 1e-10
WEIGHT_DECAY: 5e-4
#
TRAIN_DATA:
NAME: ImageTextPairMSDatasetForACE
MODE: train
MS_DATASET_NAME: cache/datasets/hed_pair
MS_DATASET_NAMESPACE: ""
MS_DATASET_SPLIT: "train"
MS_DATASET_SUBNAME: ""
PROMPT_PREFIX: ""
REPLACE_STYLE: False
MAX_SEQ_LEN: 4096
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 1
SAMPLER:
NAME: LoopSampler
#
TRAIN_HOOKS:
-
NAME: BackwardHook
PRIORITY: 0
-
NAME: LogHook
LOG_INTERVAL: 50
-
NAME: CheckpointHook
INTERVAL: 100
-
NAME: ProbeDataHook
PROB_INTERVAL: 100
@@ -0,0 +1,235 @@
ENV:
BACKEND: nccl
SEED: 42
TENSOR_PARALLEL_SIZE: 1
PIPELINE_PARALLEL_SIZE: 1
SYS_ENVS:
TORCH_CUDNN_V8_API_ENABLED: '1'
TOKENIZERS_PARALLELISM: 'false'
TF_CPP_MIN_LOG_LEVEL: '3'
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
#
SOLVER:
NAME: LatentDiffusionVideoSolver
MAX_STEPS: 2000
USE_AMP: True
DTYPE: bfloat16
USE_FAIRSCALE: False
USE_FSDP: True
LOAD_MODEL_ONLY: False
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_cogvideox_2b_lora
LOG_FILE: std_log.txt
EVAL_INTERVAL: 100
LOG_TRAIN_NUM: 4
ENABLE_GRADSCALER: False
USE_SCALER: False
FPS: 8
SHARDING_STRATEGY: full_shard
FSDP_REDUCE_DTYPE: float32
FSDP_BUFFER_DTYPE: float32
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
TUNER:
- NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
#
MODEL:
NAME: LatentDiffusionCogVideoX
PRETRAINED_MODEL:
PARAMETERIZATION: v
TIMESTEPS: 1000
MIN_SNR_GAMMA: 3.0
ZERO_TERMINAL_SNR: True
SCALE_FACTOR_SPATIAL: 8
SCALE_FACTOR_TEMPORAL: 4
SCALING_FACTOR_IMAGE: 1.15258426
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: v
NOISE_SCHEDULER:
NAME: ScaledLinearScheduler
BETA_MIN: 0.00085
BETA_MAX: 0.012
SNR_SHIFT_SCALE: 3.0
RESCALE_BETAS_ZERO_SNR: True
DIFFUSION_SAMPLERS:
NAME: DDIMSampler
DISCRETIZATION_TYPE: trailing
ETA: 0.0
#
DIFFUSION_MODEL:
NAME: CogVideoXTransformer3DModel
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@transformer/diffusion_pytorch_model.safetensors
NUM_ATTENTION_HEADS: 30
ATTENTION_HEAD_DIM: 64
IN_CHANNELS: 16
OUT_CHANNELS: 16
FLIP_SIN_TO_COS: True
FREQ_SHIFT: 0
TIME_EMBED_DIM: 512
TEXT_EMBED_DIM: 4096
NUM_LAYERS: 30
DROPOUT: 0.0
ATTENTION_BIAS: True
SAMPLE_WIDTH: 90
SAMPLE_HEIGHT: 60
SAMPLE_FRAMES: 49
PATCH_SIZE: 2
TEMPORAL_COMPRESSION_RATIO: 4
MAX_TEXT_SEQ_LENGTH: 226
ACTIVATION_FN: "gelu-approximate"
TIMESTEP_ACTIVATION_FN: "silu"
NORM_ELEMENTWISE_AFFINE: True
NORM_EPS: 1e-5
SPATIAL_INTERPOLATION_SCALE: 1.875
TEMPORAL_INTERPOLATION_SCALE: 1.0
USE_ROTARY_POSITIONAL_EMBEDDINGS: False
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
GRADIENT_CHECKPOINTING: False
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@vae/diffusion_pytorch_model.safetensors
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: False
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: False
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
LENGTH: 226
CLEAN:
USE_GRAD: False
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 42
GUIDE_SCALE: 6.0
GUIDE_RESCALE: 0.0
NUM_FRAMES: 49
#
OPTIMIZER:
NAME: Adam
LEARNING_RATE: 1e-3
BETAS: [ 0.9, 0.95 ]
EPS: 1e-8
WEIGHT_DECAY: 0.0
AMSGRAD: False
#
# LR_SCHEDULER:
# NAME: StepAnnealingLR
# WARMUP_STEPS: 200
# TOTAL_STEPS: 2000
# DECAY_MODE: 'cosine'
#
TRAIN_DATA:
NAME: VideoGenDataset
MODE: train
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
PROMPT_PREFIX: 'DISNEY '
SAMPLER:
NAME: MixtureOfSamplers
SUB_SAMPLERS:
- NAME: MultiLevelBatchSampler
PROB: 1.0
FIELDS: [ "video_path", "prompt" ]
DELIMITER: '#;#'
PATH_PREFIX: cache/datasets/Disney-VideoGeneration-Dataset/
INDEX_FILE: cache/datasets/Disney-VideoGeneration-Dataset/index.jsonl
TRANSFORMS:
- NAME: Select
KEYS: [ 'video', "prompt" ]
META_KEYS: [ ]
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
PROMPT_FILE:
PROMPT_DATA: [ "A girl riding a bike.", "A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance." ]
IMAGE_SIZE: [ 480, 720 ]
FIELDS: [ "prompt" ]
DELIMITER: '#;#'
PROMPT_PREFIX: 'DISNEY '
PIN_MEMORY: True
BATCH_SIZE: 1
USE_NUM: 8
NUM_WORKERS: 4
TRANSFORMS:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
PRIORITY: 20
- NAME: CheckpointHook
INTERVAL: 1000
PRIORITY: 40
#
EVAL_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
@@ -0,0 +1,266 @@
ENV:
BACKEND: nccl
SEED: 42
TENSOR_PARALLEL_SIZE: 1
PIPELINE_PARALLEL_SIZE: 1
SYS_ENVS:
TORCH_CUDNN_V8_API_ENABLED: '1'
TOKENIZERS_PARALLELISM: 'false'
TF_CPP_MIN_LOG_LEVEL: '3'
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
#
SOLVER:
NAME: LatentDiffusionVideoSolver
MAX_STEPS: 2000
USE_AMP: True
DTYPE: bfloat16
USE_FAIRSCALE: False
USE_FSDP: True
LOAD_MODEL_ONLY: False
ENABLE_GRADSCALER: False
USE_SCALER: False
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_cogvideox_5b_i2v_lora
LOG_FILE: std_log.txt
EVAL_INTERVAL: 100
LOG_TRAIN_NUM: 4
FPS: 8
SHARDING_STRATEGY: full_shard
FSDP_REDUCE_DTYPE: float32
FSDP_BUFFER_DTYPE: float32
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
TUNER:
- NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
#
MODEL:
NAME: LatentDiffusionCogVideoX
PRETRAINED_MODEL:
PARAMETERIZATION: v
TIMESTEPS: 1000
MIN_SNR_GAMMA: 3.0
ZERO_TERMINAL_SNR: True
SCALE_FACTOR_SPATIAL: 8
SCALE_FACTOR_TEMPORAL: 4
SCALING_FACTOR_IMAGE: 0.7 # 5b diff
NOISED_IMAGE_DROPOUT: 0.05
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: v
NOISE_SCHEDULER:
NAME: ScaledLinearScheduler
BETA_MIN: 0.00085
BETA_MAX: 0.012
SNR_SHIFT_SCALE: 1.0 # 5b diff
RESCALE_BETAS_ZERO_SNR: True
DIFFUSION_SAMPLERS:
NAME: DDIMSampler
DISCRETIZATION_TYPE: trailing
ETA: 0.0
#
DIFFUSION_MODEL:
NAME: CogVideoXTransformer3DModel
DTYPE: bfloat16
PRETRAINED_MODEL: # 5b-I2V diff
- ms://AI-ModelScope/CogVideoX-5b-I2V@transformer/diffusion_pytorch_model-00001-of-00003.safetensors
- ms://AI-ModelScope/CogVideoX-5b-I2V@transformer/diffusion_pytorch_model-00002-of-00003.safetensors
- ms://AI-ModelScope/CogVideoX-5b-I2V@transformer/diffusion_pytorch_model-00003-of-00003.safetensors
NUM_ATTENTION_HEADS: 48 # 5b diff
ATTENTION_HEAD_DIM: 64
IN_CHANNELS: 32 # 5b-I2V diff
LATENT_CHANNELS: 16
OUT_CHANNELS: 16
FLIP_SIN_TO_COS: True
FREQ_SHIFT: 0
TIME_EMBED_DIM: 512
TEXT_EMBED_DIM: 4096
NUM_LAYERS: 42 # 5b diff
DROPOUT: 0.0
ATTENTION_BIAS: True
SAMPLE_WIDTH: 90
SAMPLE_HEIGHT: 60
SAMPLE_FRAMES: 49
PATCH_SIZE: 2
TEMPORAL_COMPRESSION_RATIO: 4
MAX_TEXT_SEQ_LENGTH: 226
ACTIVATION_FN: "gelu-approximate"
TIMESTEP_ACTIVATION_FN: "silu"
NORM_ELEMENTWISE_AFFINE: True
NORM_EPS: 1e-5
SPATIAL_INTERPOLATION_SCALE: 1.875
TEMPORAL_INTERPOLATION_SCALE: 1.0
USE_ROTARY_POSITIONAL_EMBEDDINGS: True # 5b diff
USE_LEARNED_POSITIONAL_EMBEDDINGS: True # 5b-I2V diff
GRADIENT_CHECKPOINTING: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b-I2V@vae/diffusion_pytorch_model.safetensors # 5b diff
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: True
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
LENGTH: 226
CLEAN:
USE_GRAD: False
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 42
GUIDE_SCALE: 6.0
GUIDE_RESCALE: 0.0
NUM_FRAMES: 49
#
OPTIMIZER:
NAME: Adam
LEARNING_RATE: 1e-3
BETAS: [ 0.9, 0.95 ]
EPS: 1e-8
WEIGHT_DECAY: 0.0
AMSGRAD: False
#
# LR_SCHEDULER:
# NAME: StepAnnealingLR
# WARMUP_STEPS: 200
# TOTAL_STEPS: 2000
# DECAY_MODE: 'cosine'
#
TRAIN_DATA:
NAME: VideoGenDataset
MODE: train
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 0
PROMPT_PREFIX: 'DISNEY '
DATA_TYPE: 'i2v'
SAMPLER:
NAME: MixtureOfSamplers
SUB_SAMPLERS:
- NAME: MultiLevelBatchSampler
PROB: 1.0
FIELDS: [ "video_path", "prompt" ]
DELIMITER: '#;#'
PATH_PREFIX: cache/datasets/Disney-VideoGeneration-Dataset/
INDEX_FILE: cache/datasets/Disney-VideoGeneration-Dataset/index.jsonl
TRANSFORMS:
- NAME: Select
KEYS: [ "video", "image", "prompt" ]
META_KEYS: [ ]
#
# EVAL_DATA:
# NAME: Text2ImageDataset
# MODE: eval
# PROMPT_FILE:
# PROMPT_DATA: [ "A cat running.#;#asset/images/edit_tuner/cat_512.jpg" ]
# FIELDS: [ "prompt", "img_path" ]
# DELIMITER: '#;#'
# PROMPT_PREFIX: ''
# PIN_MEMORY: True
# BATCH_SIZE: 1
# USE_NUM: 8
# NUM_WORKERS: 0
# IMAGE_SIZE: [ 480, 720 ]
# TRANSFORMS:
# - NAME: LoadImageFromFileList
# FILE_KEYS: [ 'img_path' ]
# RGB_ORDER: RGB
# BACKEND: pillow
# - NAME: FlexibleResize
# INTERPOLATION: bilinear
# SIZE: [ 480, 720 ]
# INPUT_KEY: [ 'img' ]
# OUTPUT_KEY: [ 'img' ]
# BACKEND: pillow
# - NAME: FlexibleCenterCrop
# SIZE: [ 480, 720 ]
# INPUT_KEY: [ 'img' ]
# OUTPUT_KEY: [ 'img' ]
# BACKEND: pillow
# - NAME: ImageToTensor
# INPUT_KEY: [ 'img' ]
# OUTPUT_KEY: [ 'img' ]
# BACKEND: pillow
# - NAME: Normalize
# MEAN: [ 0.5, 0.5, 0.5 ]
# STD: [ 0.5, 0.5, 0.5 ]
# INPUT_KEY: [ 'img' ]
# OUTPUT_KEY: [ 'image' ]
# BACKEND: torchvision
# - NAME: Select
# KEYS: [ 'image', 'prompt' ]
# META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
PRIORITY: 20
- NAME: CheckpointHook
INTERVAL: 1000
PRIORITY: 40
#
# EVAL_HOOKS:
# - NAME: ProbeDataHook
# PROB_INTERVAL: 100
# PRIORITY: 0
@@ -0,0 +1,273 @@
ENV:
BACKEND: nccl
SEED: 42
TENSOR_PARALLEL_SIZE: 1
PIPELINE_PARALLEL_SIZE: 1
SYS_ENVS:
TORCH_CUDNN_V8_API_ENABLED: '1'
TOKENIZERS_PARALLELISM: 'false'
TF_CPP_MIN_LOG_LEVEL: '3'
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
#
SOLVER:
NAME: LatentDiffusionVideoSolver
MAX_STEPS: 2000
USE_AMP: True
DTYPE: bfloat16
USE_FAIRSCALE: False
USE_FSDP: True
LOAD_MODEL_ONLY: False
ENABLE_GRADSCALER: False
USE_SCALER: False
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_cogvideox_5b_lora
LOG_FILE: std_log.txt
EVAL_INTERVAL: 100
LOG_TRAIN_NUM: 4
FPS: 8
SHARDING_STRATEGY: full_shard
FSDP_REDUCE_DTYPE: float32
FSDP_BUFFER_DTYPE: float32
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
TUNER:
- NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
#
MODEL:
NAME: LatentDiffusionCogVideoX
PRETRAINED_MODEL:
PARAMETERIZATION: v
TIMESTEPS: 1000
MIN_SNR_GAMMA: 3.0
ZERO_TERMINAL_SNR: True
SCALE_FACTOR_SPATIAL: 8
SCALE_FACTOR_TEMPORAL: 4
SCALING_FACTOR_IMAGE: 0.7 # 5b diff
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: v
NOISE_SCHEDULER:
NAME: ScaledLinearScheduler
BETA_MIN: 0.00085
BETA_MAX: 0.012
SNR_SHIFT_SCALE: 1.0 # 5b diff
RESCALE_BETAS_ZERO_SNR: True
DIFFUSION_SAMPLERS:
NAME: DDIMSampler
DISCRETIZATION_TYPE: trailing
ETA: 0.0
#
DIFFUSION_MODEL:
NAME: CogVideoXTransformer3DModel
DTYPE: bfloat16
PRETRAINED_MODEL: # 5b diff
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00001-of-00002.safetensors
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00002-of-00002.safetensors
NUM_ATTENTION_HEADS: 48 # 5b diff
ATTENTION_HEAD_DIM: 64
IN_CHANNELS: 16
OUT_CHANNELS: 16
FLIP_SIN_TO_COS: True
FREQ_SHIFT: 0
TIME_EMBED_DIM: 512
TEXT_EMBED_DIM: 4096
NUM_LAYERS: 42 # 5b diff
DROPOUT: 0.0
ATTENTION_BIAS: True
SAMPLE_WIDTH: 90
SAMPLE_HEIGHT: 60
SAMPLE_FRAMES: 49
PATCH_SIZE: 2
TEMPORAL_COMPRESSION_RATIO: 4
MAX_TEXT_SEQ_LENGTH: 226
ACTIVATION_FN: "gelu-approximate"
TIMESTEP_ACTIVATION_FN: "silu"
NORM_ELEMENTWISE_AFFINE: True
NORM_EPS: 1e-5
SPATIAL_INTERPOLATION_SCALE: 1.875
TEMPORAL_INTERPOLATION_SCALE: 1.0
USE_ROTARY_POSITIONAL_EMBEDDINGS: True # 5b diff
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
GRADIENT_CHECKPOINTING: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b@vae/diffusion_pytorch_model.safetensors # 5b diff
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: True
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
LENGTH: 226
CLEAN:
USE_GRAD: False
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 42
GUIDE_SCALE: 6.0
GUIDE_RESCALE: 0.0
NUM_FRAMES: 49
#
OPTIMIZER:
NAME: Adam
LEARNING_RATE: 1e-3
BETAS: [ 0.9, 0.95 ]
EPS: 1e-8
WEIGHT_DECAY: 0.0
AMSGRAD: False
#
# LR_SCHEDULER:
# NAME: StepAnnealingLR
# WARMUP_STEPS: 200
# TOTAL_STEPS: 2000
# DECAY_MODE: 'cosine'
#
TRAIN_DATA:
NAME: VideoGenDatasetOTF
MODE: train
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
PROMPT_PREFIX: 'DISNEY '
DELIMITER: '#;#'
FIELDS: [ 'video_path', 'prompt' ]
PATH_PREFIX: cache/datasets/Disney-VideoGeneration-Dataset/
DATA_FILE: cache/datasets/Disney-VideoGeneration-Dataset/index.jsonl
SAMPLER:
NAME: LoopSampler
TRANSFORMS:
- NAME: Select
KEYS: [ 'video', 'video_latent', "prompt" ]
META_KEYS: [ ]
MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b@vae/diffusion_pytorch_model.safetensors
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: True
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
PROMPT_FILE:
PROMPT_DATA: [ "A girl riding a bike.", "A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance." ]
IMAGE_SIZE: [ 480, 720 ]
FIELDS: [ "prompt" ]
DELIMITER: '#;#'
PROMPT_PREFIX: 'DISNEY '
PIN_MEMORY: True
BATCH_SIZE: 1
USE_NUM: 8
NUM_WORKERS: 4
TRANSFORMS:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
PRIORITY: 20
- NAME: CheckpointHook
INTERVAL: 1000
PRIORITY: 40
#
EVAL_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
@@ -2,33 +2,22 @@ ENV:
BACKEND: nccl
SEED: 166666
SOLVER:
# NAME DESCRIPTION: TYPE: default: 'LatentUfitSolver'
NAME: LatentDiffusionSolver
# 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
# USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
USE_FAIRSCALE: False
# USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
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: ''
ENABLE_GRADSCALER: False
USE_SCALER: False
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_flux_dev_1024_lora
LOG_FILE: std_log.txt
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
EVAL_INTERVAL: 100
# 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: [ 'model', 'cond_stage_model.t5_model' ] #
SAVE_MODULES: [ 'model'] #
TRAIN_MODULES: ['model']
@@ -58,61 +47,36 @@ SOLVER:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
NOISE_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
NAME: FlowMatchSigmaScheduler
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
# LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
LOGIT_MEAN: 0.0
# LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
LOGIT_STD: 1.0
# MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
MODE_SCALE: 1.29
SAMPLER_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
NAME: FlowMatchFluxShiftScheduler
# SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
SHIFT: False
# SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
SIGMOID_SCALE: 1
# BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
BASE_SHIFT: 0.5
# MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
MAX_SHIFT: 1.15
#
DIFFUSION_MODEL:
# NAME DESCRIPTION: TYPE: default: 'Flux'
NAME: Flux
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
# IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
IN_CHANNELS: 64
# HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
HIDDEN_SIZE: 3072
# 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
USE_GRAD_CHECKPOINT: True
@@ -157,55 +121,34 @@ SOLVER:
TANH_OUT: False
#
COND_STAGE_MODEL:
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
NAME: T5PlusClipFluxEmbedder
# T5_MODEL DESCRIPTION: TYPE: default: ''
T5_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# 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: ms://AI-ModelScope/FLUX.1-dev@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: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
# 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: HFEmbedder
# 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: ms://AI-ModelScope/FLUX.1-dev@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: ms://AI-ModelScope/FLUX.1-dev@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
USE_GRAD_CHECKPOINT: True
#
SAMPLE_ARGS:
SAMPLE_STEPS: 50
SAMPLER: flow_eluer
SAMPLER: flow_euler
SEED: 2024
IMAGE_SIZE: [ 1024, 1024 ]
GUIDE_SCALE: 3.5
@@ -2,35 +2,24 @@ ENV:
BACKEND: nccl
SEED: 166666
SOLVER:
# NAME DESCRIPTION: TYPE: default: 'LatentUfitSolver'
NAME: LatentDiffusionSolver
# 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
# USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
USE_FAIRSCALE: False
# USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
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: ''
ENABLE_GRADSCALER: False
USE_SCALER: False
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_flux_schnell_1024_lora
LOG_FILE: std_log.txt
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
EVAL_INTERVAL: 100
# 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: [ 'model', 'cond_stage_model.t5_model' ] #
SAVE_MODULES: [ 'model'] #
SAVE_MODULES: [ 'model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
@@ -58,12 +47,9 @@ SOLVER:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
NOISE_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
NAME: FlowMatchSigmaScheduler
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
WEIGHTING_SCHEME: logit_normal
@@ -157,54 +143,33 @@ SOLVER:
TANH_OUT: False
#
COND_STAGE_MODEL:
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
NAME: T5PlusClipFluxEmbedder
# T5_MODEL DESCRIPTION: TYPE: default: ''
T5_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# 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: ms://AI-ModelScope/FLUX.1-schnell@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: ms://AI-ModelScope/FLUX.1-schnell@tokenizer_2/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 256
# 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: HFEmbedder
# 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: ms://AI-ModelScope/FLUX.1-schnell@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: ms://AI-ModelScope/FLUX.1-schnell@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
#
SAMPLE_ARGS:
SAMPLE_STEPS: 4
SAMPLER: flow_eluer
SAMPLER: flow_euler
SEED: 2024
IMAGE_SIZE: [ 1024, 1024 ]
GUIDE_SCALE: 3.5
+1 -1
View File
@@ -8,11 +8,11 @@ FILE_SYSTEM:
TEMP_DIR: ./cache/cache_data
#
ENABLE_I2V: False
SKIP_EXAMPLES: True
#
MODEL:
EDIT_MODEL:
MODEL_CFG_DIR: scepter/methods/studio/chatbot/models/
DEFAULT: ace_0.6b_512
I2V:
MODEL_NAME: CogVideoX-5b-I2V
MODEL_DIR: ms://ZhipuAI/CogVideoX-5b-I2V/
@@ -0,0 +1,128 @@
NAME: ACE_0.6B_1024
IS_DEFAULT: False
USE_DYNAMIC_MODEL: True
DEFAULT_PARAS:
PARAS:
#
INPUT:
INPUT_IMAGE:
INPUT_MASK:
TASK:
PROMPT: ""
NEGATIVE_PROMPT: ""
OUTPUT_HEIGHT: 1024
OUTPUT_WIDTH: 1024
SAMPLER: ddim
SAMPLE_STEPS: 50
GUIDE_SCALE: 4.5
GUIDE_RESCALE: 0.5
SEED: -1
TAR_INDEX: 0
OUTPUT:
LATENT:
IMAGES:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: float16
INPUT: ["IMAGE"]
- NAME: decode
DTYPE: float16
INPUT: ["LATENT"]
#
DIFFUSION_MODEL:
FUNCTION:
- NAME: forward
DTYPE: float16
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE"]
#
COND_STAGE_MODEL:
FUNCTION:
- NAME: encode_list_of_list
DTYPE: bfloat16
INPUT: ["PROMPT"]
#
MODEL:
NAME: LatentDiffusionACE
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
DECODER_BIAS: 0.5
DEFAULT_N_PROMPT: ""
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
USE_TEXT_POS_EMBEDDINGS: True
#
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: eps
MIN_SNR_GAMMA:
NOISE_SCHEDULER:
NAME: LinearScheduler
NUM_TIMESTEPS: 1000
BETA_MIN: 0.0001
BETA_MAX: 0.02
#
DIFFUSION_MODEL:
NAME: ACE
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/dit/ace_0.6b_1024px.pth
IGNORE_KEYS: [ ]
PATCH_SIZE: 2
IN_CHANNELS: 4
HIDDEN_SIZE: 1152
DEPTH: 28
NUM_HEADS: 16
MLP_RATIO: 4.0
PRED_SIGMA: True
DROP_PATH: 0.0
WINDOW_DIZE: 0
Y_CHANNELS: 4096
MAX_SEQ_LEN: 4096
QK_NORM: True
USE_GRAD_CHECKPOINT: True
ATTENTION_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/vae/vae.bin
IGNORE_KEYS: []
#
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: 4
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: 4
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/text_encoder/t5-v1_1-xxl/
TOKENIZER_PATH: ms://iic/ACE-0.6B-1024px@models/tokenizer/t5-v1_1-xxl
LENGTH: 120
T5_DTYPE: bfloat16
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
CLEAN: whitespace
USE_GRAD: False
@@ -0,0 +1,284 @@
NAME: ACE_0.6B_1024_REFINER
IS_DEFAULT: False
USE_DYNAMIC_MODEL: True
DEFAULT_PARAS:
PARAS:
#
INPUT:
INPUT_IMAGE:
INPUT_MASK:
TASK:
PROMPT: ""
NEGATIVE_PROMPT: ""
OUTPUT_HEIGHT: 1024
OUTPUT_WIDTH: 1024
SAMPLER: ddim
SAMPLE_STEPS: 50
GUIDE_SCALE: 4.5
GUIDE_RESCALE: 0.5
SEED: -1
TAR_INDEX: 0
REFINER_SCALE: 0.2
USE_ACE: True
#REFINER_PROMPT: "High Resolution, Sharpness, Clarity, Detail Enhancement, Noise Reduction, HD, 4k, Image Restoration, HDR"
REFINER_PROMPT: "High Resolution, Sharpness, Clarity, Detail Enhancement, Noise Reduction, HD, 4k, Image Restoration, HDR"
OUTPUT:
LATENT:
IMAGES:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: float16
INPUT: ["IMAGE"]
- NAME: decode
DTYPE: float16
INPUT: ["LATENT"]
#
DIFFUSION_MODEL:
FUNCTION:
- NAME: forward
DTYPE: float16
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE"]
#
COND_STAGE_MODEL:
FUNCTION:
- NAME: encode_list_of_list
DTYPE: bfloat16
INPUT: ["PROMPT"]
#
MODEL:
NAME: LatentDiffusionACE
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
DECODER_BIAS: 0.5
DEFAULT_N_PROMPT: ""
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
USE_TEXT_POS_EMBEDDINGS: True
#
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: eps
MIN_SNR_GAMMA:
NOISE_SCHEDULER:
NAME: LinearScheduler
NUM_TIMESTEPS: 1000
BETA_MIN: 0.0001
BETA_MAX: 0.02
#
DIFFUSION_MODEL:
NAME: ACE
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/dit/ace_0.6b_1024px.pth
IGNORE_KEYS: [ ]
PATCH_SIZE: 2
IN_CHANNELS: 4
HIDDEN_SIZE: 1152
DEPTH: 28
NUM_HEADS: 16
MLP_RATIO: 4.0
PRED_SIGMA: True
DROP_PATH: 0.0
WINDOW_DIZE: 0
Y_CHANNELS: 4096
MAX_SEQ_LEN: 4096
QK_NORM: True
USE_GRAD_CHECKPOINT: True
ATTENTION_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/vae/vae.bin
IGNORE_KEYS: []
#
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: 4
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: 4
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/text_encoder/t5-v1_1-xxl/
TOKENIZER_PATH: ms://iic/ACE-0.6B-1024px@models/tokenizer/t5-v1_1-xxl
LENGTH: 120
T5_DTYPE: bfloat16
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
CLEAN: whitespace
USE_GRAD: False
ACE_PROMPT: [
"A cute cartoon rabbit holding a whiteboard that says 'ACE Refiner', standing in a sunny meadow filled with flowers, with a big smile and bright colors.",
"A beautiful young woman with long flowing hair, wearing a summer dress, holding a whiteboard that reads 'ACE Refiner' while sitting on a park bench surrounded by cherry blossoms.",
"An adorable cartoon cat wearing oversized glasses, holding a whiteboard that says 'ACE Refiner', perched on a stack of colorful books in a cozy library setting.",
"A charming girl with pigtails, wearing a cute school uniform, enthusiastically holding a whiteboard that has 'ACE Refiner' written on it, in a bright and cheerful classroom full of educational posters.",
"A friendly cartoon dog with floppy ears, sitting in front of a doghouse, proudly holding a whiteboard that says 'ACE Refiner', with a playful expression and a blue sky in the background.",
"A cute anime girl with big expressive eyes, dressed in a colorful outfit, holding a whiteboard that reads 'ACE Refiner' in a fantastical landscape filled with mythical creatures.",
"A vibrant cartoon fox holding a whiteboard that says 'ACE Refiner', standing on a rock by a sparkling stream, surrounded by lush greenery and butterflies.",
"A stylish young woman in a business outfit, smiling as she holds a whiteboard written with 'ACE Refiner', in a modern office filled with plants and natural light.",
"A cute cartoon unicorn holding a sparkling whiteboard that says 'ACE Refiner', frolicking in a magical forest, with rainbows and stars in the background.",
"A happy family, consisting of a cute little girl and her playful puppy, holding a whiteboard that says 'ACE Refiner', together in their backyard on a sunny day."
]
REFINER_MODEL:
NAME: ""
IS_DEFAULT: False
DEFAULT_PARAS:
PARAS:
RESOLUTIONS: [ [ 1024, 1024 ] ]
INPUT:
INPUT_IMAGE:
INPUT_MASK:
TASK:
PROMPT: ""
NEGATIVE_PROMPT: ""
OUTPUT_HEIGHT: 1024
OUTPUT_WIDTH: 1024
SAMPLER: flow_euler
SAMPLE_STEPS: 30
GUIDE_SCALE: 3.5
GUIDE_RESCALE:
OUTPUT:
LATENT:
IMAGES:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: bfloat16
INPUT: [ "IMAGE" ]
- NAME: decode
DTYPE: bfloat16
INPUT: [ "LATENT" ]
PARAS:
SCALE_FACTOR: 1.5305
SHIFT_FACTOR: 0.0609
SIZE_FACTOR: 8
DIFFUSION_MODEL:
FUNCTION:
- NAME: forward
DTYPE: bfloat16
INPUT: [ "SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE" ]
COND_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: bfloat16
INPUT: [ "PROMPT" ]
MODEL:
DIFFUSION:
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
NOISE_SCHEDULER:
NAME: FlowMatchSigmaScheduler
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
LOGIT_MEAN: 0.0
LOGIT_STD: 1.0
MODE_SCALE: 1.29
DIFFUSION_MODEL:
NAME: FluxMR
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
IN_CHANNELS: 64
OUT_CHANNELS: 64
HIDDEN_SIZE: 3072
NUM_HEADS: 24
AXES_DIM: [ 16, 56, 56 ]
THETA: 10000
VEC_IN_DIM: 768
GUIDANCE_EMBED: True
CONTEXT_IN_DIM: 4096
MLP_RATIO: 4.0
QKV_BIAS: True
DEPTH: 19
DEPTH_SINGLE_BLOCKS: 38
USE_GRAD_CHECKPOINT: True
ATTN_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLFlux
EMBED_DIM: 16
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@ae.safetensors
IGNORE_KEYS: [ ]
BATCH_SIZE: 8
USE_CONV: False
SCALE_FACTOR: 0.3611
SHIFT_FACTOR: 0.1159
#
ENCODER:
NAME: Encoder
USE_CHECKPOINT: False
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
USE_CHECKPOINT: False
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: T5PlusClipFluxEmbedder
T5_MODEL:
NAME: HFEmbedder
HF_MODEL_CLS: T5EncoderModel
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder_2/
HF_TOKENIZER_CLS: T5Tokenizer
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
MAX_LENGTH: 512
OUTPUT_KEY: last_hidden_state
D_TYPE: bfloat16
BATCH_INFER: False
CLEAN: whitespace
CLIP_MODEL:
NAME: HFEmbedder
HF_MODEL_CLS: CLIPTextModel
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder/
HF_TOKENIZER_CLS: CLIPTokenizer
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer/
MAX_LENGTH: 77
OUTPUT_KEY: pooler_output
D_TYPE: bfloat16
BATCH_INFER: True
CLEAN: whitespace
@@ -1,5 +1,6 @@
NAME: ACE_0.6B_512
IS_DEFAULT: False
IS_DEFAULT: True
USE_DYNAMIC_MODEL: True
DEFAULT_PARAS:
PARAS:
#
@@ -39,7 +40,7 @@ DEFAULT_PARAS:
#
COND_STAGE_MODEL:
FUNCTION:
- NAME: encode_list
- NAME: encode_list_of_list
DTYPE: bfloat16
INPUT: ["PROMPT"]
#
@@ -0,0 +1,151 @@
NAME: COGVIDEOX_2B
IS_DEFAULT: False
DEFAULT_PARAS:
PARAS:
RESOLUTIONS: [[480, 720]]
INPUT:
IMAGE:
ORIGINAL_SIZE_AS_TUPLE: [480, 720]
TARGET_SIZE_AS_TUPLE: [480, 720]
PROMPT: ""
NEGATIVE_PROMPT: ""
PROMPT_PREFIX: ""
SAMPLE: ddim
SAMPLE_STEPS: 50
GUIDE_SCALE: 6.0
GUIDE_RESCALE: 0.0
DISCRETIZATION: trailing
NUM_FRAMES:
DEFAULT: 49
VISIBLE: True
FPS:
DEFAULT: 8
VISIBLE: True
OUTPUT:
VIDEOS:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
-
NAME: decode
DTYPE: bfloat16
INPUT: ["LATENT"]
PARAS:
SCALING_FACTOR_IMAGE: 1.15258426
DIFFUSION_MODEL:
FUNCTION:
-
NAME: forward
DTYPE: bfloat16
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE", "GUIDE_RESCALE", "DISCRETIZATION", "NUM_FRAMES", "FPS"]
PARAS:
USE_ROTARY_POSITIONAL_EMBEDDINGS: False
PATCH_SIZE: 2
LATENT_CHANNELS: 16
SCALE_FACTOR_SPATIAL: 8
SCALE_FACTOR_TEMPORAL: 4
ATTENTION_HEAD_DIM: 64
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
COND_STAGE_MODEL:
FUNCTION:
-
NAME: encode
DTYPE: bfloat16
INPUT: ["PROMPT"]
#
MODEL:
PRETRAINED_MODEL:
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: v
NOISE_SCHEDULER:
NAME: ScaledLinearScheduler
BETA_MIN: 0.00085
BETA_MAX: 0.012
SNR_SHIFT_SCALE: 3.0
RESCALE_BETAS_ZERO_SNR: True
DIFFUSION_SAMPLERS:
NAME: DDIMSampler
DISCRETIZATION_TYPE: trailing
ETA: 0.0
#
DIFFUSION_MODEL:
NAME: CogVideoXTransformer3DModel
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@transformer/diffusion_pytorch_model.safetensors
NUM_ATTENTION_HEADS: 30
ATTENTION_HEAD_DIM: 64
IN_CHANNELS: 16
OUT_CHANNELS: 16
FLIP_SIN_TO_COS: True
FREQ_SHIFT: 0
TIME_EMBED_DIM: 512
TEXT_EMBED_DIM: 4096
NUM_LAYERS: 30
DROPOUT: 0.0
ATTENTION_BIAS: True
SAMPLE_WIDTH: 90
SAMPLE_HEIGHT: 60
SAMPLE_FRAMES: 49
PATCH_SIZE: 2
TEMPORAL_COMPRESSION_RATIO: 4
MAX_TEXT_SEQ_LENGTH: 226
ACTIVATION_FN: "gelu-approximate"
TIMESTEP_ACTIVATION_FN: "silu"
NORM_ELEMENTWISE_AFFINE: True
NORM_EPS: 1e-5
SPATIAL_INTERPOLATION_SCALE: 1.875
TEMPORAL_INTERPOLATION_SCALE: 1.0
USE_ROTARY_POSITIONAL_EMBEDDINGS: False
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
GRADIENT_CHECKPOINTING: False
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@vae/diffusion_pytorch_model.safetensors
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: False
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: False
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
LENGTH: 226
CLEAN:
USE_GRAD: False
@@ -0,0 +1,153 @@
NAME: COGVIDEOX_5B
IS_DEFAULT: False
DEFAULT_PARAS:
PARAS:
RESOLUTIONS: [[480, 720]]
INPUT:
IMAGE:
ORIGINAL_SIZE_AS_TUPLE: [480, 720]
TARGET_SIZE_AS_TUPLE: [480, 720]
PROMPT: ""
NEGATIVE_PROMPT: ""
PROMPT_PREFIX: ""
SAMPLE: ddim
SAMPLE_STEPS: 50
GUIDE_SCALE: 6.0
GUIDE_RESCALE: 0.0
DISCRETIZATION: trailing
NUM_FRAMES:
DEFAULT: 49
VISIBLE: True
FPS:
DEFAULT: 8
VISIBLE: True
OUTPUT:
VIDEOS:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
-
NAME: decode
DTYPE: bfloat16
INPUT: ["LATENT"]
PARAS:
SCALING_FACTOR_IMAGE: 0.7 # 5b diff
DIFFUSION_MODEL:
FUNCTION:
-
NAME: forward
DTYPE: bfloat16
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE", "GUIDE_RESCALE", "DISCRETIZATION", "NUM_FRAMES", "FPS"]
PARAS:
USE_ROTARY_POSITIONAL_EMBEDDINGS: True
PATCH_SIZE: 2
LATENT_CHANNELS: 16
SCALE_FACTOR_SPATIAL: 8
SCALE_FACTOR_TEMPORAL: 4
ATTENTION_HEAD_DIM: 64
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
COND_STAGE_MODEL:
FUNCTION:
-
NAME: encode
DTYPE: bfloat16
INPUT: ["PROMPT"]
#
MODEL:
PRETRAINED_MODEL:
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: v
NOISE_SCHEDULER:
NAME: ScaledLinearScheduler
BETA_MIN: 0.00085
BETA_MAX: 0.012
SNR_SHIFT_SCALE: 1.0 # 5b diff
RESCALE_BETAS_ZERO_SNR: True
DIFFUSION_SAMPLERS:
NAME: DDIMSampler
DISCRETIZATION_TYPE: trailing
ETA: 0.0
#
DIFFUSION_MODEL:
NAME: CogVideoXTransformer3DModel
DTYPE: bfloat16
PRETRAINED_MODEL: # 5b diff
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00001-of-00002.safetensors
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00002-of-00002.safetensors
NUM_ATTENTION_HEADS: 48 # 5b diff
ATTENTION_HEAD_DIM: 64
IN_CHANNELS: 16
OUT_CHANNELS: 16
FLIP_SIN_TO_COS: True
FREQ_SHIFT: 0
TIME_EMBED_DIM: 512
TEXT_EMBED_DIM: 4096
NUM_LAYERS: 42 # 5b diff
DROPOUT: 0.0
ATTENTION_BIAS: True
SAMPLE_WIDTH: 90
SAMPLE_HEIGHT: 60
SAMPLE_FRAMES: 49
PATCH_SIZE: 2
TEMPORAL_COMPRESSION_RATIO: 4
MAX_TEXT_SEQ_LENGTH: 226
ACTIVATION_FN: "gelu-approximate"
TIMESTEP_ACTIVATION_FN: "silu"
NORM_ELEMENTWISE_AFFINE: True
NORM_EPS: 1e-5
SPATIAL_INTERPOLATION_SCALE: 1.875
TEMPORAL_INTERPOLATION_SCALE: 1.0
USE_ROTARY_POSITIONAL_EMBEDDINGS: True # 5b diff
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
GRADIENT_CHECKPOINTING: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b@vae/diffusion_pytorch_model.safetensors # 5b diff
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: True
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
LENGTH: 226
CLEAN:
USE_GRAD: False
@@ -13,8 +13,8 @@ DEFAULT_PARAS:
VISIBLE: False
PROMPT_PREFIX: ""
SAMPLE:
VALUES: ["flow_eluer"]
DEFAULT: "flow_eluer"
VALUES: ["flow_euler"]
DEFAULT: "flow_euler"
SAMPLE_STEPS: 50
GUIDE_SCALE: 3.5
GUIDE_RESCALE:
@@ -13,8 +13,8 @@ DEFAULT_PARAS:
VISIBLE: False
PROMPT_PREFIX: ""
SAMPLE:
VALUES: ["flow_eluer"]
DEFAULT: "flow_eluer"
VALUES: ["flow_euler"]
DEFAULT: "flow_euler"
SAMPLE_STEPS: 4
GUIDE_SCALE: 3.5
GUIDE_RESCALE:
@@ -18,6 +18,16 @@ DIFFUSION_PARAS:
MAX: 4
DEFAULT: 1
VISIBLE: True
NUM_FRAMES:
MIN: 1
MAX: 100
DEFAULT: 49
VISIBLE: False
FPS:
MIN: 1
MAX: 50
DEFAULT: 8
VISIBLE: False
SAMPLE_STEPS:
MIN: 1
MAX: 100
@@ -93,7 +103,8 @@ DIFFUSION_PARAS:
[1664, 576], [1728, 576],
[2048, 2048], [2048, 1920], [1920, 2048],
[1536, 2560], [2560, 1536], [2560, 1440],
[2560, 1440]
[2560, 1440],
[480, 720], [720, 480]
]
DEFAULT: [1024, 1024]
VISIBLE: True
@@ -450,3 +450,28 @@ PROCESSORS:
SRC_IMAGE_TOOL: sketch
SRC_IMAGE_INTERACTIVE: True
CAPTION_INTERACTIVE: False
VIDEO_PROCESSORS:
- NAME: CogVLM2Llama3Caption
TYPE: caption
MODEL_PATH: ms://ZhipuAI/cogvlm2-llama3-caption
DEVICE: "gpu"
MEMORY: 20000
PROMPT: Please describe this video in detail.
TEMPERATURE: 0.1
MAX_NEW_TOKENS: 2048
PAD_TOKEN_ID: 128002
TOP_K: 1
TOP_P: 0.1
TRANSLATION_PROCESSORS:
- NAME: OpusMtZhEn
TYPE: caption
MODEL_PATH: ms://cubeai/trans-opus-mt-zh-en
DEVICE: "gpu"
MEMORY: 5000
- NAME: OpusMtEnZh
TYPE: caption
MODEL_PATH: ms://cubeai/trans-opus-mt-en-zh
DEVICE: "gpu"
MEMORY: 5000
+1 -1
View File
@@ -89,5 +89,5 @@ INTERFACE:
CONFIG: scepter/methods/studio/inference/inference.yaml
- NAME: 对话式编辑
NAME_EN: ChatBot
IFID: ChatBot
IFID: chatbot
CONFIG: scepter/methods/studio/chatbot/chatbot.yaml
@@ -0,0 +1,315 @@
ENV:
BACKEND: nccl
SEED: 42
TENSOR_PARALLEL_SIZE: 1
PIPELINE_PARALLEL_SIZE: 1
SYS_ENVS:
TORCH_CUDNN_V8_API_ENABLED: '1'
TOKENIZERS_PARALLELISM: 'false'
TF_CPP_MIN_LOG_LEVEL: '3'
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
META:
VERSION: 'COGVIDEOX_2B'
DESCRIPTION: "cogvideox 2b"
IS_DEFAULT: False
IS_SHARE: True
INFERENCE_PARAS:
INFERENCE_BATCH_SIZE: 1
INFERENCE_PREFIX: ""
DEFAULT_SAMPLER: "ddim"
DEFAULT_SAMPLE_STEPS: 50
INFERENCE_N_PROMPT: ""
RESOLUTION: [ 480, 720 ]
PARAS:
- TRAIN_BATCH_SIZE: 1
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: [ 480, 720 ]
MEMORY: 89000
EPOCHS: 50
SAVE_INTERVAL: 25
EPSEC: 0.818
LEARNING_RATE: 4e-4
IS_DEFAULT: False
TUNER: FULL
- TRAIN_BATCH_SIZE: 1
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: [ 480, 720 ]
MEMORY: 89000
EPOCHS: 50
SAVE_INTERVAL: 25
EPSEC: 0.818
LEARNING_RATE: 4e-4
IS_DEFAULT: True
TUNER: LORA
#
TUNERS:
LORA:
- NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
#
SOLVER:
NAME: LatentDiffusionVideoSolver
MAX_STEPS: 2000
USE_AMP: True
DTYPE: bfloat16
USE_FAIRSCALE: False
USE_FSDP: True
LOAD_MODEL_ONLY: False
ENABLE_GRADSCALER: False
USE_SCALER: False
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_cogvideox_2b_lora
LOG_FILE: std_log.txt
EVAL_INTERVAL: 100
LOG_TRAIN_NUM: 4
FPS: 8
SHARDING_STRATEGY: full_shard
FSDP_REDUCE_DTYPE: float32
FSDP_BUFFER_DTYPE: float32
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
TUNER:
#
MODEL:
NAME: LatentDiffusionCogVideoX
PRETRAINED_MODEL:
PARAMETERIZATION: v
TIMESTEPS: 1000
MIN_SNR_GAMMA: 3.0
ZERO_TERMINAL_SNR: True
SCALE_FACTOR_SPATIAL: 8
SCALE_FACTOR_TEMPORAL: 4
SCALING_FACTOR_IMAGE: 1.15258426
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: v
NOISE_SCHEDULER:
NAME: ScaledLinearScheduler
BETA_MIN: 0.00085
BETA_MAX: 0.012
SNR_SHIFT_SCALE: 3.0
RESCALE_BETAS_ZERO_SNR: True
DIFFUSION_SAMPLERS:
NAME: DDIMSampler
DISCRETIZATION_TYPE: trailing
ETA: 0.0
#
DIFFUSION_MODEL:
NAME: CogVideoXTransformer3DModel
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@transformer/diffusion_pytorch_model.safetensors
NUM_ATTENTION_HEADS: 30
ATTENTION_HEAD_DIM: 64
IN_CHANNELS: 16
OUT_CHANNELS: 16
FLIP_SIN_TO_COS: True
FREQ_SHIFT: 0
TIME_EMBED_DIM: 512
TEXT_EMBED_DIM: 4096
NUM_LAYERS: 30
DROPOUT: 0.0
ATTENTION_BIAS: True
SAMPLE_WIDTH: 90
SAMPLE_HEIGHT: 60
SAMPLE_FRAMES: 49
PATCH_SIZE: 2
TEMPORAL_COMPRESSION_RATIO: 4
MAX_TEXT_SEQ_LENGTH: 226
ACTIVATION_FN: "gelu-approximate"
TIMESTEP_ACTIVATION_FN: "silu"
NORM_ELEMENTWISE_AFFINE: True
NORM_EPS: 1e-5
SPATIAL_INTERPOLATION_SCALE: 1.875
TEMPORAL_INTERPOLATION_SCALE: 1.0
USE_ROTARY_POSITIONAL_EMBEDDINGS: False
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
GRADIENT_CHECKPOINTING: False
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@vae/diffusion_pytorch_model.safetensors
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: False
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: False
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
LENGTH: 226
CLEAN:
USE_GRAD: False
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 42
GUIDE_SCALE: 6.0
GUIDE_RESCALE: 0.0
NUM_FRAMES: 49
#
OPTIMIZER:
NAME: Adam
LEARNING_RATE: 1e-3
BETAS: [ 0.9, 0.95 ]
EPS: 1e-8
WEIGHT_DECAY: 0.0
AMSGRAD: False
#
# LR_SCHEDULER:
# NAME: StepAnnealingLR
# WARMUP_STEPS: 200
# TOTAL_STEPS: 2000
# DECAY_MODE: 'cosine'
#
TRAIN_DATA:
NAME: VideoGenDatasetOTF
MODE: train
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
PROMPT_PREFIX: ''
DELIMITER: '#;#'
FIELDS: [ 'video_path', 'width', 'height', 'prompt' ]
PATH_PREFIX:
DATA_FILE:
SAMPLER:
NAME: LoopSampler
TRANSFORMS:
- NAME: Select
KEYS: [ 'video', 'video_latent', "prompt" ]
META_KEYS: [ ]
MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-2b@vae/diffusion_pytorch_model.safetensors
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: True
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
PROMPT_FILE:
PROMPT_DATA: [ "A girl riding a bike.", "A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance." ]
IMAGE_SIZE: [ 480, 720 ]
FIELDS: [ "prompt" ]
DELIMITER: '#;#'
PROMPT_PREFIX: ''
PIN_MEMORY: True
BATCH_SIZE: 1
# USE_NUM: 8
NUM_WORKERS: 4
TRANSFORMS:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
PRIORITY: 20
- NAME: CheckpointHook
INTERVAL: 1000
PRIORITY: 40
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
DISABLE_SNAPSHOT: True
#
EVAL_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
SAVE_PROBE_PREFIX: 'image'
@@ -0,0 +1,317 @@
ENV:
BACKEND: nccl
SEED: 42
TENSOR_PARALLEL_SIZE: 1
PIPELINE_PARALLEL_SIZE: 1
SYS_ENVS:
TORCH_CUDNN_V8_API_ENABLED: '1'
TOKENIZERS_PARALLELISM: 'false'
TF_CPP_MIN_LOG_LEVEL: '3'
PYTORCH_CUDA_ALLOC_CONF: 'expandable_segments:True'
META:
VERSION: 'COGVIDEOX_5B'
DESCRIPTION: "cogvideox 5b"
IS_DEFAULT: False
IS_SHARE: True
INFERENCE_PARAS:
INFERENCE_BATCH_SIZE: 1
INFERENCE_PREFIX: ""
DEFAULT_SAMPLER: "ddim"
DEFAULT_SAMPLE_STEPS: 50
INFERENCE_N_PROMPT: ""
RESOLUTION: [ 480, 720 ]
PARAS:
- TRAIN_BATCH_SIZE: 1
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: [ 480, 720 ]
MEMORY: 89000
EPOCHS: 50
SAVE_INTERVAL: 25
EPSEC: 0.818
LEARNING_RATE: 4e-4
IS_DEFAULT: False
TUNER: FULL
- TRAIN_BATCH_SIZE: 1
TRAIN_PREFIX: ""
TRAIN_N_PROMPT: ""
RESOLUTION: [ 480, 720 ]
MEMORY: 89000
EPOCHS: 50
SAVE_INTERVAL: 25
EPSEC: 0.818
LEARNING_RATE: 4e-4
IS_DEFAULT: True
TUNER: LORA
#
TUNERS:
LORA:
- NAME: SwiftLoRA
R: 64
LORA_ALPHA: 64
LORA_DROPOUT: 0.0
BIAS: "none"
TARGET_MODULES: "model.*(.to_k|.to_q|.to_v|.to_out.0)$"
#
SOLVER:
NAME: LatentDiffusionVideoSolver
MAX_STEPS: 2000
USE_AMP: True
DTYPE: bfloat16
USE_FAIRSCALE: False
USE_FSDP: True
LOAD_MODEL_ONLY: False
ENABLE_GRADSCALER: False
USE_SCALER: False
RESUME_FROM:
WORK_DIR: ./cache/save_data/dit_cogvideox_5b_lora
LOG_FILE: std_log.txt
EVAL_INTERVAL: 100
LOG_TRAIN_NUM: 4
FPS: 8
SHARDING_STRATEGY: full_shard
FSDP_REDUCE_DTYPE: float32
FSDP_BUFFER_DTYPE: float32
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.model']
SAVE_MODULES: [ 'model', 'cond_stage_model.model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
TUNER:
#
MODEL:
NAME: LatentDiffusionCogVideoX
PRETRAINED_MODEL:
PARAMETERIZATION: v
TIMESTEPS: 1000
MIN_SNR_GAMMA: 3.0
ZERO_TERMINAL_SNR: True
SCALE_FACTOR_SPATIAL: 8
SCALE_FACTOR_TEMPORAL: 4
SCALING_FACTOR_IMAGE: 0.7 # 5b diff
IGNORE_KEYS: [ ]
DEFAULT_N_PROMPT:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: v
NOISE_SCHEDULER:
NAME: ScaledLinearScheduler
BETA_MIN: 0.00085
BETA_MAX: 0.012
SNR_SHIFT_SCALE: 1.0 # 5b diff
RESCALE_BETAS_ZERO_SNR: True
DIFFUSION_SAMPLERS:
NAME: DDIMSampler
DISCRETIZATION_TYPE: trailing
ETA: 0.0
#
DIFFUSION_MODEL:
NAME: CogVideoXTransformer3DModel
DTYPE: bfloat16
PRETRAINED_MODEL: # 5b diff
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00001-of-00002.safetensors
- ms://AI-ModelScope/CogVideoX-5b@transformer/diffusion_pytorch_model-00002-of-00002.safetensors
NUM_ATTENTION_HEADS: 48 # 5b diff
ATTENTION_HEAD_DIM: 64
IN_CHANNELS: 16
OUT_CHANNELS: 16
FLIP_SIN_TO_COS: True
FREQ_SHIFT: 0
TIME_EMBED_DIM: 512
TEXT_EMBED_DIM: 4096
NUM_LAYERS: 42 # 5b diff
DROPOUT: 0.0
ATTENTION_BIAS: True
SAMPLE_WIDTH: 90
SAMPLE_HEIGHT: 60
SAMPLE_FRAMES: 49
PATCH_SIZE: 2
TEMPORAL_COMPRESSION_RATIO: 4
MAX_TEXT_SEQ_LENGTH: 226
ACTIVATION_FN: "gelu-approximate"
TIMESTEP_ACTIVATION_FN: "silu"
NORM_ELEMENTWISE_AFFINE: True
NORM_EPS: 1e-5
SPATIAL_INTERPOLATION_SCALE: 1.875
TEMPORAL_INTERPOLATION_SCALE: 1.0
USE_ROTARY_POSITIONAL_EMBEDDINGS: True # 5b diff
USE_LEARNED_POSITIONAL_EMBEDDINGS: False
GRADIENT_CHECKPOINTING: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b@vae/diffusion_pytorch_model.safetensors # 5b diff
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: True
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://AI-ModelScope/t5-v1_1-xxl
TOKENIZER_PATH: ms://AI-ModelScope/t5-v1_1-xxl
LENGTH: 226
CLEAN:
USE_GRAD: False
#
LOSS:
NAME: ReconstructLoss
LOSS_TYPE: l2
#
SAMPLE_ARGS:
SAMPLER: ddim
SAMPLE_STEPS: 50
SEED: 42
GUIDE_SCALE: 6.0
GUIDE_RESCALE: 0.0
NUM_FRAMES: 49
#
OPTIMIZER:
NAME: Adam
LEARNING_RATE: 1e-3
BETAS: [ 0.9, 0.95 ]
EPS: 1e-8
WEIGHT_DECAY: 0.0
AMSGRAD: False
#
# LR_SCHEDULER:
# NAME: StepAnnealingLR
# WARMUP_STEPS: 200
# TOTAL_STEPS: 2000
# DECAY_MODE: 'cosine'
#
TRAIN_DATA:
NAME: VideoGenDatasetOTF
MODE: train
PIN_MEMORY: True
BATCH_SIZE: 1
NUM_WORKERS: 4
PROMPT_PREFIX: ''
DELIMITER: '#;#'
FIELDS: [ 'video_path', 'width', 'height', 'prompt' ]
PATH_PREFIX:
DATA_FILE:
SAMPLER:
NAME: LoopSampler
TRANSFORMS:
- NAME: Select
KEYS: [ 'video', 'video_latent', "prompt" ]
META_KEYS: [ ]
MODEL:
NAME: AutoencoderKLCogVideoX
DTYPE: bfloat16
PRETRAINED_MODEL: ms://AI-ModelScope/CogVideoX-5b@vae/diffusion_pytorch_model.safetensors
SAMPLE_HEIGHT: 480
SAMPLE_WIDTH: 720
USE_QUANT_CONV: False
USE_POST_QUANT_CONV: False
USE_SLICING: True
USE_TILING: True
GRADIENT_CHECKPOINTING: True
ENCODER:
NAME: CogVideoXEncoder3D
IN_CHANNELS: 3
OUT_CHANNELS: 16
UP_BLOCK_TYPES: [ "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D", "CogVideoXDownBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
DECODER:
NAME: CogVideoXDecoder3D
IN_CHANNELS: 16
OUT_CHANNELS: 3
UP_BLOCK_TYPES: [ "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D", "CogVideoXUpBlock3D" ]
BLOCK_OUT_CHANNELS: [ 128, 256, 256, 512 ]
LAYERS_PER_BLOCK: 3
ACT_FN: "silu"
NORM_EPS: 1e-6
NORM_NUM_GROUPS: 32
DROPOUT: 0.0
PAD_MODE: "first"
TEMPORAL_COMPRESSION_RATIO: 4
GRADIENT_CHECKPOINTING: True
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
PROMPT_FILE:
PROMPT_DATA: [ "A girl riding a bike.", "A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance." ]
IMAGE_SIZE: [ 480, 720 ]
FIELDS: [ "prompt" ]
DELIMITER: '#;#'
PROMPT_PREFIX: ''
PIN_MEMORY: True
BATCH_SIZE: 1
# USE_NUM: 8
NUM_WORKERS: 4
TRANSFORMS:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
PRIORITY: 20
- NAME: CheckpointHook
INTERVAL: 1000
PRIORITY: 40
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
DISABLE_SNAPSHOT: True
#
EVAL_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
SAVE_PROBE_PREFIX: 'image'
@@ -3,7 +3,7 @@ ENV:
META:
VERSION: 'FLUX1.0_DEV'
DESCRIPTION: "flux 1.0 dev"
IS_DEFAULT: False
IS_DEFAULT: True
IS_SHARE: True
INFERENCE_PARAS:
INFERENCE_BATCH_SIZE: 1
@@ -50,43 +50,33 @@ META:
#
SOLVER:
NAME: LatentDiffusionSolver
# 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
# USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
USE_FAIRSCALE: False
# USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
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: ./cache/save_data/dit_flux_dev_1024_lora
LOG_FILE: std_log.txt
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
EVAL_INTERVAL: 100
# 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'
ENABLE_GRADSCALER: False
USE_SCALER: False
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: [ 'model', 'cond_stage_model.t5_model' ] #
SAVE_MODULES: [ 'model'] #
SAVE_MODULES: [ 'model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
FREEZE:
#
TUNER:
#
MODEL:
NAME: LatentDiffusionFlux
PARAMETERIZATION: rf
@@ -99,65 +89,39 @@ SOLVER:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
NOISE_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
NAME: FlowMatchSigmaScheduler
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
# LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
LOGIT_MEAN: 0.0
# LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
LOGIT_STD: 1.0
# MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
MODE_SCALE: 1.29
SAMPLER_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
NAME: FlowMatchFluxShiftScheduler
# SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
SHIFT: False
# SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
SIGMOID_SCALE: 1
# BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
BASE_SHIFT: 0.5
# MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
MAX_SHIFT: 1.15
#
DIFFUSION_MODEL:
# NAME DESCRIPTION: TYPE: default: 'Flux'
NAME: Flux
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
# IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
IN_CHANNELS: 64
# HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
HIDDEN_SIZE: 3072
# 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: False
# 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
USE_GRAD_CHECKPOINT: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLFlux
EMBED_DIM: 16
@@ -167,7 +131,7 @@ SOLVER:
USE_CONV: False
SCALE_FACTOR: 0.3611
SHIFT_FACTOR: 0.1159
#
ENCODER:
NAME: Encoder
USE_CHECKPOINT: True
@@ -181,7 +145,7 @@ SOLVER:
DOUBLE_Z: True
DROPOUT: 0.0
RESAMP_WITH_CONV: True
#
DECODER:
NAME: Decoder
USE_CHECKPOINT: True
@@ -196,61 +160,40 @@ SOLVER:
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
NAME: T5PlusClipFluxEmbedder
# T5_MODEL DESCRIPTION: TYPE: default: ''
T5_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# 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: ms://AI-ModelScope/FLUX.1-dev@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: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
# 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: HFEmbedder
# 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: ms://AI-ModelScope/FLUX.1-dev@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: ms://AI-ModelScope/FLUX.1-dev@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
#
SAMPLE_ARGS:
SAMPLE_STEPS: 50
SAMPLER: flow_eluer
SAMPLER: flow_euler
SEED: 2024
IMAGE_SIZE: [ 1024, 1024 ]
SHIFT: True
GUIDE_SCALE: 3.5
#
OPTIMIZER:
NAME: AdamW
LEARNING_RATE: 4e-4
@@ -258,7 +201,7 @@ SOLVER:
EPS: 1e-8
WEIGHT_DECAY: 1e-2
AMSGRAD: False
#
TRAIN_DATA:
NAME: ImageTextPairMSDataset
MODE: train
@@ -302,7 +245,7 @@ SOLVER:
- NAME: Select
KEYS: [ 'image', 'prompt' ]
META_KEYS: [ 'data_key' ]
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
@@ -319,13 +262,12 @@ SOLVER:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
# GRADIENT_CLIP: 1.0
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
@@ -50,43 +50,33 @@ META:
#
SOLVER:
NAME: LatentDiffusionSolver
# 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
# USE_FAIRSCALE DESCRIPTION: Use fairscale as the backend of ddp, default False. TYPE: bool default: False
USE_FAIRSCALE: False
# USE_FSDP DESCRIPTION: Use fsdp as the backend of ddp, default False. TYPE: bool default: False
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: ./cache/save_data/dit_flux_schnell_1024_lora
LOG_FILE: std_log.txt
# EVAL_INTERVAL DESCRIPTION: Eval the model interval. TYPE: int default: 1
EVAL_INTERVAL: 100
# 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'
ENABLE_GRADSCALER: False
USE_SCALER: False
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: [ 'model', 'cond_stage_model.t5_model' ] #
SAVE_MODULES: [ 'model'] #
FSDP_SHARD_MODULES: [ 'model', 'cond_stage_model.t5_model' ]
SAVE_MODULES: [ 'model']
TRAIN_MODULES: ['model']
#
FILE_SYSTEM:
NAME: "ModelscopeFs"
TEMP_DIR: "./cache/cache_data"
#
FREEZE:
#
TUNER:
#
MODEL:
NAME: LatentDiffusionFlux
PARAMETERIZATION: rf
@@ -99,65 +89,39 @@ SOLVER:
USE_EMA: False
EVAL_EMA: False
DIFFUSION:
# NAME DESCRIPTION: TYPE: default: 'DiffusionFluxRF'
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
# NOISE_SCHEDULER DESCRIPTION: TYPE: default: ''
NOISE_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchSigmaScheduler'
NAME: FlowMatchSigmaScheduler
# WEIGHTING_SCHEME DESCRIPTION: The weighting scheme for sampling timesteps, choose from ['sigma_sqrt', 'logit_normal', 'mode', 'cosmap', 'none']. TYPE: str default: 'logit_normal'
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
# LOGIT_MEAN DESCRIPTION: The mean of the logit distribution for sampling timesteps. TYPE: float default: 0.0
LOGIT_MEAN: 0.0
# LOGIT_STD DESCRIPTION: The standard deviation of the logit distribution for sampling timesteps. TYPE: float default: 1.0
LOGIT_STD: 1.0
# MODE_SCALE DESCRIPTION: The scale factor for the mode of the logit distribution for sampling timesteps. TYPE: float default: 1.29
MODE_SCALE: 1.29
SAMPLER_SCHEDULER:
# NAME DESCRIPTION: TYPE: default: 'FlowMatchFluxShiftScheduler'
NAME: FlowMatchFluxShiftScheduler
# SHIFT DESCRIPTION: Use timestamp shift or not, default is True. TYPE: bool default: True
SHIFT: False
# SIGMOID_SCALE DESCRIPTION: The scale of sigmoid function for sampling timesteps. TYPE: int default: 1
SIGMOID_SCALE: 1
# BASE_SHIFT DESCRIPTION: The base shift factor for the timestamp. TYPE: float default: 0.5
BASE_SHIFT: 0.5
# MAX_SHIFT DESCRIPTION: The max shift factor for the timestamp. TYPE: float default: 1.15
MAX_SHIFT: 1.15
#
DIFFUSION_MODEL:
# NAME DESCRIPTION: TYPE: default: 'Flux'
NAME: Flux
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-schnell@flux1-schnell.safetensors
# IN_CHANNELS DESCRIPTION: model's input channels. TYPE: int default: 64
IN_CHANNELS: 64
# HIDDEN_SIZE DESCRIPTION: model's hidden size. TYPE: int default: 1024
HIDDEN_SIZE: 3072
# 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: False
# 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
USE_GRAD_CHECKPOINT: True
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLFlux
EMBED_DIM: 16
@@ -167,7 +131,7 @@ SOLVER:
USE_CONV: False
SCALE_FACTOR: 0.3611
SHIFT_FACTOR: 0.1159
#
ENCODER:
NAME: Encoder
USE_CHECKPOINT: True
@@ -181,7 +145,7 @@ SOLVER:
DOUBLE_Z: True
DROPOUT: 0.0
RESAMP_WITH_CONV: True
#
DECODER:
NAME: Decoder
USE_CHECKPOINT: True
@@ -196,60 +160,39 @@ SOLVER:
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
# NAME DESCRIPTION: TYPE: default: 'T5PlusClipFluxEmbedder'
NAME: T5PlusClipFluxEmbedder
# T5_MODEL DESCRIPTION: TYPE: default: ''
T5_MODEL:
# NAME DESCRIPTION: TYPE: default: 'HFEmbedder'
NAME: HFEmbedder
# 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: ms://AI-ModelScope/FLUX.1-schnell@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: ms://AI-ModelScope/FLUX.1-schnell@tokenizer_2/
# MAX_LENGTH DESCRIPTION: max length of input TYPE: int default: 77
MAX_LENGTH: 256
# 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: HFEmbedder
# 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: ms://AI-ModelScope/FLUX.1-schnell@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: ms://AI-ModelScope/FLUX.1-schnell@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
#
SAMPLE_ARGS:
SAMPLE_STEPS: 4
SAMPLER: flow_eluer
SAMPLER: flow_euler
SEED: 2024
IMAGE_SIZE: [ 1024, 1024 ]
GUIDE_SCALE: 3.5
#
OPTIMIZER:
NAME: AdamW
LEARNING_RATE: 4e-4
@@ -257,7 +200,7 @@ SOLVER:
EPS: 1e-8
WEIGHT_DECAY: 1e-2
AMSGRAD: False
#
TRAIN_DATA:
NAME: ImageTextPairMSDataset
MODE: train
@@ -301,7 +244,7 @@ SOLVER:
- NAME: Select
KEYS: [ 'image', 'prompt' ]
META_KEYS: [ 'data_key' ]
#
EVAL_DATA:
NAME: Text2ImageDataset
MODE: eval
@@ -318,13 +261,12 @@ SOLVER:
- NAME: Select
KEYS: [ 'index', 'prompt' ]
META_KEYS: [ 'image_size' ]
#
TRAIN_HOOKS:
- NAME: ProbeDataHook
PROB_INTERVAL: 100
PRIORITY: 0
- NAME: BackwardHook
# GRADIENT_CLIP: 1.0
PRIORITY: 10
- NAME: LogHook
LOG_INTERVAL: 10
@@ -339,4 +281,4 @@ SOLVER:
PROB_INTERVAL: 100
SAVE_LAST: True
SAVE_NAME_PREFIX: 'step'
SAVE_PROBE_PREFIX: 'image'
SAVE_PROBE_PREFIX: 'image'
@@ -35,7 +35,7 @@ META:
SAVE_INTERVAL: 25
EPSEC: 0.818
LEARNING_RATE: 0.0001
IS_DEFAULT: False
IS_DEFAULT: True
TUNER: LORA
#
TUNERS:
@@ -13,8 +13,10 @@ TRAIN_PARAS:
VALUES: [[256, 256], [320, 180], [180, 320],
[512, 512], [640, 360], [360, 640],
[768, 768], [960, 540], [540, 960],
[1024, 1024], [1280, 720], [720, 1280]]
[1024, 1024], [1280, 720], [720, 1280],
[720, 480], [480, 720]]
DEFAULT: [1024, 1024]
EVAL_PROMPTS:
- a boy wearing a jacket
- a dog running on the lawn
SAVE_FILE_LOCAL_PATH: "cache/scepter_ui/datasets/train_data_from_list"
+2 -2
View File
@@ -7,6 +7,6 @@ from scepter.modules.data.dataset.dataset import (Image2ImageDataset,
ImageTextPairDataset,
Text2ImageDataset)
from scepter.modules.data.dataset.ms_dataset import (
ImageTextPairFolderDataset, ImageTextPairMSDataset,
ImageTextPairMSDatasetForACE)
ImageTextPairFolderDataset, ImageTextPairMSDataset)
from scepter.modules.data.dataset.registry import DATASETS
from scepter.modules.data.dataset.video_gen_dataset import VideoGenDataset
+8 -1
View File
@@ -242,6 +242,8 @@ class Text2ImageDataset(BaseDataset):
prompt_prefix = cfg.get('PROMPT_PREFIX', '')
path_prefix = cfg.get('PATH_PREFIX', '')
use_num = cfg.get('USE_NUM', -1)
meta_cfg = cfg.get('META_CFG', None)
meta_cfg = meta_cfg.get_lowercase_dict() if meta_cfg is not None else None
image_size = cfg.get('IMAGE_SIZE', 1024)
if isinstance(image_size, numbers.Number):
@@ -264,7 +266,12 @@ class Text2ImageDataset(BaseDataset):
self.items = list()
for i, row in enumerate(rows):
item = {'index': i, 'meta': {'image_size': image_size}}
if meta_cfg is not None:
meta_cfg_copy = copy.deepcopy(meta_cfg)
meta_cfg_copy['image_size'] = image_size
item = {'index': i, 'meta': meta_cfg_copy}
else:
item = {'index': i, 'meta': {'image_size': image_size}}
for key, value in zip(fields, row):
if key in ['prompt', 'caption', 'text']:
item['ori_prompt'] = value
@@ -0,0 +1,184 @@
import io
import random
import sys
import os
import warnings
import torch
import numpy as np
from tqdm import tqdm
from scepter.modules.utils.distribute import we
from scepter.modules.data.dataset import DATASETS, BaseDataset
from scepter.modules.utils.file_system import FS
try:
import decord
decord.bridge.set_bridge("torch")
except ImportError:
warnings.warn(
"The `decord` package is required for loading the video dataset. Install with `pip install decord`"
)
@DATASETS.register_class()
class VideoGenDataset(BaseDataset):
def __init__(self, cfg, logger = None):
super().__init__(cfg, logger=logger)
self.prompt_prefix = cfg.get('PROMPT_PREFIX', '')
self.path_prefix = cfg.get('PATH_PREFIX', '')
self.p_zero = cfg.get('P_ZERO', 0.0)
self.max_num_frames = cfg.get("NUM_FRAMES", 49)
self.fps = cfg.get("FPS", 8)
self.height = cfg.get("HEIGHT", 480)
self.width = cfg.get("WIDTH", 720)
self.skip_frames_start = cfg.get("SKIP_FRAMES_START", 0)
self.skip_frames_end = cfg.get("SKIP_FRAMES_END", 0)
self.data_type = cfg.get('DATA_TYPE', 't2v')
def worker_init_fn(self, worker_id, num_workers=1):
super().worker_init_fn(worker_id, num_workers=num_workers)
randseed = np.random.randint(0, 2 ** 32 - num_workers - 1)
workerseed = randseed + worker_id
random.seed(workerseed)
np.random.seed(workerseed)
def _preprocess_video_data(self, video_path):
with FS.get_object(video_path) as video_data:
video_reader = decord.VideoReader(io.BytesIO(video_data), width=self.width, height=self.height)
video_num_frames = len(video_reader)
start_frame = min(self.skip_frames_start, video_num_frames)
end_frame = max(0, video_num_frames - self.skip_frames_end)
if end_frame <= start_frame:
frames = video_reader.get_batch([start_frame])
elif end_frame - start_frame <= self.max_num_frames:
frames = video_reader.get_batch(list(range(start_frame, end_frame)))
else:
indices = list(range(start_frame, end_frame, (end_frame - start_frame) // self.max_num_frames))
frames = video_reader.get_batch(indices)
# Ensure that we don't go over the limit
frames = frames[: self.max_num_frames]
selected_num_frames = frames.shape[0]
# Choose first (4k + 1) frames as this is how many is required by the VAE
remainder = (3 + (selected_num_frames % 4)) % 4
if remainder != 0:
frames = frames[:-remainder]
selected_num_frames = frames.shape[0]
assert (selected_num_frames - 1) % 4 == 0
# Training transforms
frames = frames.float().div_(127.5).sub_(1.)
frames = frames.permute(3, 0, 1, 2).contiguous() # [C, F, H, W]
return frames
def _parse_index(self, index):
meta = dict()
for key, value in zip(index[-1], index[:-1]):
if key in ['oss_key', 'path', 'video_path']:
meta['video_path'] = value
elif key in ['prompt', 'caption', 'text']:
meta['prompt'] = value
elif key in ['width', 'height']:
meta[key] = int(value)
else:
meta[key] = value
return meta
def _get(self, index):
meta = self._parse_index(index)
video_path = os.path.join(self.path_prefix, meta.get('video_path', ''))
video = self._preprocess_video_data(video_path)
prompt = self.prompt_prefix + meta.get('prompt', '')
if self.mode == 'train' and np.random.uniform() < self.p_zero:
prompt = ''
item = {
'video': video,
'prompt': prompt,
'meta': meta,
}
if self.data_type == 'i2v':
item['image'] = item['video'][:, :1, :, :]
return item
def __len__(self):
return sys.maxsize
@staticmethod
def collate_fn(batch):
collect = {}
for sample in batch:
for k, v in sample.items():
if k not in collect:
collect[k] = []
collect[k].append(v)
return collect
@DATASETS.register_class()
class VideoGenDatasetOTF(VideoGenDataset):
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger)
self.data_file = cfg.DATA_FILE
self.delimiter = cfg.get('DELIMITER', '#;#')
self.fields = cfg.get('FIELDS', ['video_path', 'prompt'])
self.use_num = cfg.get('USE_NUM', -1)
from scepter.modules.model.registry import MODELS
model_cfg = cfg.get('MODEL', None)
if model_cfg is not None:
self.model = MODELS.build(cfg.MODEL, logger=logger).eval().requires_grad_(False).to(we.device_id)
self.items = self.parse_data(self.data_file, self.delimiter, self.fields)
if self.use_num and self.use_num > 0:
self.items = self.items[:self.use_num]
self.data = self.encode(self.items)
self.real_number = len(self.data)
if model_cfg is not None:
self.model.to('cpu')
del self.model
torch.cuda.empty_cache()
def parse_data(self, data_file, delimiter, fields):
items = list()
with FS.get_object(data_file) as local_data:
rows = [
i.split(delimiter,
len(fields) - 1)
for i in local_data.decode('utf-8').strip().split('\n')
]
for i, row in enumerate(rows):
item = {}
for key, value in zip(self.fields, row):
if key in ['oss_key', 'path', 'video_path']:
item['video_path'] = value
elif key in ['prompt', 'caption', 'text']:
item['prompt'] = value
elif key in ['width', 'height']:
item[key] = int(value)
else:
item[key] = value
items.append(item)
return items
def encode(self, items):
self.logger.info("Start to encode video data [{}]!".format(len(items)))
for item in tqdm(items):
video_path = os.path.join(self.path_prefix, item.get('video_path', ''))
video = self._preprocess_video_data(video_path)
latent = self.model.encode_first_stage(video.unsqueeze(0).to(we.device_id)).squeeze(0)
item['video_latent'] = latent.detach().cpu()
item['video'] = video
if self.data_type == 'i2v':
item['image'] = item['video'][:, :1, :, :]
return items
def _get(self, index):
return self.data[index % self.real_number]
+281 -106
View File
@@ -10,7 +10,7 @@ import torch.nn as nn
import torch.nn.functional as F
import torchvision.transforms.functional as TF
from PIL import Image
import torchvision.transforms as T
from scepter.modules.model.registry import DIFFUSIONS
from scepter.modules.model.utils.basic_utils import check_list_of_list
from scepter.modules.model.utils.basic_utils import \
@@ -85,6 +85,138 @@ class TextEmbedding(nn.Module):
super().__init__()
self.pos = nn.Parameter(data=torch.zeros(embedding_shape))
class RefinerInference(DiffusionInference):
def init_from_cfg(self, cfg):
self.use_dynamic_model = cfg.get('USE_DYNAMIC_MODEL', True)
super().init_from_cfg(cfg)
self.diffusion = DIFFUSIONS.build(cfg.MODEL.DIFFUSION, logger=self.logger) \
if cfg.MODEL.have('DIFFUSION') else None
self.max_seq_length = cfg.MODEL.get("MAX_SEQ_LENGTH", 4096)
assert self.diffusion is not None
if not self.use_dynamic_model:
self.dynamic_load(self.first_stage_model, 'first_stage_model')
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
self.dynamic_load(self.diffusion_model, 'diffusion_model')
@torch.no_grad()
def encode_first_stage(self, x, **kwargs):
_, dtype = self.get_function_info(self.first_stage_model, 'encode')
with torch.autocast('cuda',
enabled=dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
def run_one_image(u):
zu = get_model(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
def upscale_resize(self, image, interpolation=T.InterpolationMode.BILINEAR):
c, H, W = image.shape
scale = max(1.0, math.sqrt(self.max_seq_length / ((H / 16) * (W / 16))))
rH = int(H * scale) // 16 * 16 # ensure divisible by self.d
rW = int(W * scale) // 16 * 16
image = T.Resize((rH, rW), interpolation=interpolation, antialias=True)(image)
return image
@torch.no_grad()
def decode_first_stage(self, z):
_, dtype = self.get_function_info(self.first_stage_model, 'decode')
with torch.autocast('cuda',
enabled=dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
return [get_model(self.first_stage_model).decode(zu) for zu in z]
def noise_sample(self, num_samples, h, w, seed, device = None, dtype = torch.bfloat16):
noise = torch.randn(
num_samples,
16,
# allow for packing
2 * math.ceil(h / 16),
2 * math.ceil(w / 16),
device=device,
dtype=dtype,
generator=torch.Generator(device=device).manual_seed(seed),
)
return noise
def refine(self,
x_samples=None,
prompt=None,
reverse_scale=-1.,
seed = 2024,
**kwargs
):
print(prompt)
value_input = copy.deepcopy(self.input)
x_samples = [self.upscale_resize(x) for x in x_samples]
noise = []
for i, x in enumerate(x_samples):
noise_ = self.noise_sample(1, x.shape[1],
x.shape[2], seed,
device = x.device)
noise.append(noise_)
noise, x_shapes = pack_imagelist_into_tensor(noise)
if reverse_scale > 0:
self.dynamic_load(self.first_stage_model, 'first_stage_model')
x_samples = [x.unsqueeze(0) for x in x_samples]
x_start = self.encode_first_stage(x_samples, **kwargs)
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=not self.use_dynamic_model)
x_start, _ = pack_imagelist_into_tensor(x_start)
else:
x_start = None
# cond stage
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
function_name, dtype = self.get_function_info(self.cond_stage_model)
with torch.autocast('cuda',
enabled=dtype == 'float16',
dtype=getattr(torch, dtype)):
ctx = getattr(get_model(self.cond_stage_model),
function_name)(prompt)
ctx["x_shapes"] = x_shapes
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=not self.use_dynamic_model)
self.dynamic_load(self.diffusion_model, 'diffusion_model')
# UNet use input n_prompt
function_name, dtype = self.get_function_info(
self.diffusion_model)
with torch.autocast('cuda',
enabled=dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
solver_sample = value_input.get('sample', 'flow_euler')
sample_steps = value_input.get('sample_steps', 20)
guide_scale = value_input.get('guide_scale', 3.5)
if guide_scale is not None:
guide_scale = torch.full((noise.shape[0],), guide_scale, device=noise.device,
dtype=noise.dtype)
else:
guide_scale = None
latent = self.diffusion.sample(
noise=noise,
sampler=solver_sample,
model=get_model(self.diffusion_model),
model_kwargs={"cond": ctx, "guidance": guide_scale},
steps=sample_steps,
show_progress=True,
guide_scale=guide_scale,
return_intermediate=None,
reverse_scale=reverse_scale,
x=x_start,
**kwargs).float()
latent = unpack_tensor_into_imagelist(latent, x_shapes)
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=not self.use_dynamic_model)
self.dynamic_load(self.first_stage_model, 'first_stage_model')
x_samples = self.decode_first_stage(latent)
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=not self.use_dynamic_model)
return x_samples
class ACEInference(DiffusionInference):
def __init__(self, logger=None):
@@ -99,6 +231,7 @@ class ACEInference(DiffusionInference):
def init_from_cfg(self, cfg):
self.name = cfg.NAME
self.is_default = cfg.get('IS_DEFAULT', False)
self.use_dynamic_model = cfg.get('USE_DYNAMIC_MODEL', True)
module_paras = self.load_default(cfg.get('DEFAULT_PARAS', None))
assert cfg.have('MODEL')
@@ -116,9 +249,22 @@ class ACEInference(DiffusionInference):
module_paras.get(
'COND_STAGE_MODEL',
None)) if cfg.MODEL.have('COND_STAGE_MODEL') else None
self.refiner_model_cfg = cfg.get('REFINER_MODEL', None)
# self.refiner_scale = cfg.get('REFINER_SCALE', 0.)
# self.refiner_prompt = cfg.get('REFINER_PROMPT', "")
self.ace_prompt = cfg.get("ACE_PROMPT", [])
if self.refiner_model_cfg:
self.refiner_model_cfg.USE_DYNAMIC_MODEL = self.use_dynamic_model
self.refiner_module = RefinerInference(self.logger)
self.refiner_module.init_from_cfg(self.refiner_model_cfg)
else:
self.refiner_module = None
self.diffusion = DIFFUSIONS.build(cfg.MODEL.DIFFUSION,
logger=self.logger)
self.interpolate_func = lambda x: (F.interpolate(
x.unsqueeze(0),
scale_factor=1 / self.size_factor,
@@ -137,6 +283,10 @@ class ACEInference(DiffusionInference):
self.size_factor = cfg.get('SIZE_FACTOR', 8)
self.decoder_bias = cfg.get('DECODER_BIAS', 0)
self.default_n_prompt = cfg.get('DEFAULT_N_PROMPT', '')
if not self.use_dynamic_model:
self.dynamic_load(self.first_stage_model, 'first_stage_model')
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
self.dynamic_load(self.diffusion_model, 'diffusion_model')
@torch.no_grad()
def encode_first_stage(self, x, **kwargs):
@@ -163,6 +313,8 @@ class ACEInference(DiffusionInference):
]
return x
@torch.no_grad()
def __call__(self,
image=None,
@@ -184,7 +336,6 @@ class ACEInference(DiffusionInference):
g = torch.Generator(device=we.device_id)
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
g.manual_seed(int(seed))
if input_image is not None:
# assert isinstance(input_image, list) and isinstance(input_mask, list)
if task is None:
@@ -237,118 +388,142 @@ class ACEInference(DiffusionInference):
assert isinstance(nn_p, list)
n_prompt[nn_p_id][-1] = negative_prompt
ctx, null_ctx = {}, {}
# Get Noise Shape
self.dynamic_load(self.first_stage_model, 'first_stage_model')
is_txt_image = sum([len(e_i) for e_i in edit_image]) < 1
image = to_device(image)
x = self.encode_first_stage(image)
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=True)
noise = [
torch.empty(*i.shape, device=we.device_id).normal_(generator=g)
for i in x
]
noise, x_shapes = pack_imagelist_into_tensor(noise)
ctx['x_shapes'] = null_ctx['x_shapes'] = x_shapes
image_mask = to_device(image_mask, strict=False)
cond_mask = [self.interpolate_func(i) for i in image_mask
] if image_mask is not None else [None] * len(image)
ctx['x_mask'] = null_ctx['x_mask'] = cond_mask
refiner_scale = kwargs.pop("refiner_scale", 0.0)
refiner_prompt = kwargs.pop("refiner_prompt", "")
use_ace = kwargs.pop("use_ace", True)
# <= 0 use ace as the txt2img generator.
if use_ace and (not is_txt_image or refiner_scale <= 0):
ctx, null_ctx = {}, {}
# Get Noise Shape
self.dynamic_load(self.first_stage_model, 'first_stage_model')
x = self.encode_first_stage(image)
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=not self.use_dynamic_model)
noise = [
torch.empty(*i.shape, device=we.device_id).normal_(generator=g)
for i in x
]
noise, x_shapes = pack_imagelist_into_tensor(noise)
ctx['x_shapes'] = null_ctx['x_shapes'] = x_shapes
# Encode Prompt
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
function_name, dtype = self.get_function_info(self.cond_stage_model)
cont, cont_mask = getattr(get_model(self.cond_stage_model),
function_name)(prompt)
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont,
cont_mask)
null_cont, null_cont_mask = getattr(get_model(self.cond_stage_model),
function_name)(n_prompt)
null_cont, null_cont_mask = self.cond_stage_embeddings(
prompt, edit_image, null_cont, null_cont_mask)
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=False)
ctx['crossattn'] = cont
null_ctx['crossattn'] = null_cont
image_mask = to_device(image_mask, strict=False)
cond_mask = [self.interpolate_func(i) for i in image_mask
] if image_mask is not None else [None] * len(image)
ctx['x_mask'] = null_ctx['x_mask'] = cond_mask
# Encode Edit Images
self.dynamic_load(self.first_stage_model, 'first_stage_model')
edit_image = [to_device(i, strict=False) for i in edit_image]
edit_image_mask = [to_device(i, strict=False) for i in edit_image_mask]
e_img, e_mask = [], []
for u, m in zip(edit_image, edit_image_mask):
if u is None:
continue
if m is None:
m = [None] * len(u)
e_img.append(self.encode_first_stage(u, **kwargs))
e_mask.append([self.interpolate_func(i) for i in m])
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=True)
null_ctx['edit'] = ctx['edit'] = e_img
null_ctx['edit_mask'] = ctx['edit_mask'] = e_mask
# Encode Prompt
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
function_name, dtype = self.get_function_info(self.cond_stage_model)
cont, cont_mask = getattr(get_model(self.cond_stage_model),
function_name)(prompt)
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont,
cont_mask)
null_cont, null_cont_mask = getattr(get_model(self.cond_stage_model),
function_name)(n_prompt)
null_cont, null_cont_mask = self.cond_stage_embeddings(
prompt, edit_image, null_cont, null_cont_mask)
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=not self.use_dynamic_model)
ctx['crossattn'] = cont
null_ctx['crossattn'] = null_cont
# Diffusion Process
self.dynamic_load(self.diffusion_model, 'diffusion_model')
function_name, dtype = self.get_function_info(self.diffusion_model)
with torch.autocast('cuda',
enabled=dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
latent = self.diffusion.sample(
noise=noise,
sampler=sampler,
model=get_model(self.diffusion_model),
model_kwargs=[{
'cond':
ctx,
'mask':
cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
}, {
'cond':
null_ctx,
'mask':
null_cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
}] if guide_scale is not None and guide_scale > 1 else {
'cond':
null_ctx,
'mask':
cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
},
steps=sample_steps,
show_progress=True,
seed=seed,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
return_intermediate=None,
**kwargs)
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=False)
# Encode Edit Images
self.dynamic_load(self.first_stage_model, 'first_stage_model')
edit_image = [to_device(i, strict=False) for i in edit_image]
edit_image_mask = [to_device(i, strict=False) for i in edit_image_mask]
e_img, e_mask = [], []
for u, m in zip(edit_image, edit_image_mask):
if u is None:
continue
if m is None:
m = [None] * len(u)
e_img.append(self.encode_first_stage(u, **kwargs))
e_mask.append([self.interpolate_func(i) for i in m])
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=not self.use_dynamic_model)
null_ctx['edit'] = ctx['edit'] = e_img
null_ctx['edit_mask'] = ctx['edit_mask'] = e_mask
# Decode to Pixel Space
self.dynamic_load(self.first_stage_model, 'first_stage_model')
samples = unpack_tensor_into_imagelist(latent, x_shapes)
x_samples = self.decode_first_stage(samples)
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=False)
# Diffusion Process
self.dynamic_load(self.diffusion_model, 'diffusion_model')
function_name, dtype = self.get_function_info(self.diffusion_model)
with torch.autocast('cuda',
enabled=dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
latent = self.diffusion.sample(
noise=noise,
sampler=sampler,
model=get_model(self.diffusion_model),
model_kwargs=[{
'cond':
ctx,
'mask':
cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
}, {
'cond':
null_ctx,
'mask':
null_cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
}] if guide_scale is not None and guide_scale > 1 else {
'cond':
null_ctx,
'mask':
cont_mask,
'text_position_embeddings':
self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
},
steps=sample_steps,
show_progress=True,
seed=seed,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
return_intermediate=None,
**kwargs)
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=not self.use_dynamic_model)
# Decode to Pixel Space
self.dynamic_load(self.first_stage_model, 'first_stage_model')
samples = unpack_tensor_into_imagelist(latent, x_shapes)
x_samples = self.decode_first_stage(samples)
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=not self.use_dynamic_model)
x_samples = [x.squeeze(0) for x in x_samples]
else:
x_samples = image
if self.refiner_module and refiner_scale > 0:
if is_txt_image:
random.shuffle(self.ace_prompt)
input_refine_prompt = [self.ace_prompt[0] + refiner_prompt if p[0] == "" else p[0] for p in prompt]
input_refine_scale = -1.
else:
input_refine_prompt = [p[0].replace("{image}", "") + " " + refiner_prompt for p in prompt]
input_refine_scale = refiner_scale
print(input_refine_prompt)
x_samples = self.refiner_module.refine(x_samples,
reverse_scale = input_refine_scale,
prompt= input_refine_prompt,
seed=seed,
use_dynamic_model=self.use_dynamic_model)
imgs = [
torch.clamp((x_i + 1.0) / 2.0 + self.decoder_bias / 255,
torch.clamp((x_i.float() + 1.0) / 2.0 + self.decoder_bias / 255,
min=0.0,
max=1.0).squeeze(0).permute(1, 2, 0).cpu().numpy()
for x_i in x_samples
@@ -0,0 +1,181 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import numpy as np
from typing import Tuple
import random
import torch
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.distribute import we
from scepter.modules.model.backbone.cogvideox.utils import get_3d_rotary_pos_embed, get_resize_crop_region_for_grid
from .diffusion_inference import DiffusionInference, get_model
from .tuner_inference import TunerInference
class CogVideoXInference(DiffusionInference):
def __init__(self, logger=None):
self.logger = logger
self.is_redefine_paras = False
self.loaded_model = {}
self.loaded_model_name = [
'diffusion_model', 'first_stage_model', 'cond_stage_model'
]
self.tuner_infer = TunerInference(self.logger)
@torch.no_grad()
def decode_first_stage(self, latents):
latents = latents.permute(0, 2, 1, 3, 4)
latents = 1 / self.first_stage_model['paras']['scaling_factor_image'] * latents
frames = get_model(self.first_stage_model).decode(latents)
return frames
def _prepare_rotary_positional_embeddings(
self,
height: int,
width: int,
num_frames: int,
device: torch.device,
) -> Tuple[torch.Tensor, torch.Tensor]:
grid_height = height // (self.diffusion_model['paras']['scale_factor_spatial'] * self.diffusion_model['paras']['patch_size'])
grid_width = width // (self.diffusion_model['paras']['scale_factor_spatial'] * self.diffusion_model['paras']['patch_size'])
base_size_width = self.diffusion_model['paras']['sample_width'] // (self.diffusion_model['paras']['scale_factor_spatial'] * self.diffusion_model['paras']['patch_size'])
base_size_height = self.diffusion_model['paras']['sample_height'] // (self.diffusion_model['paras']['scale_factor_spatial'] * self.diffusion_model['paras']['patch_size'])
grid_crops_coords = get_resize_crop_region_for_grid(
(grid_height, grid_width), base_size_width, base_size_height
)
freqs_cos, freqs_sin = get_3d_rotary_pos_embed(
embed_dim=self.diffusion_model['paras']['attention_head_dim'],
crops_coords=grid_crops_coords,
grid_size=(grid_height, grid_width),
temporal_size=num_frames,
)
freqs_cos = freqs_cos.to(device=device)
freqs_sin = freqs_sin.to(device=device)
return freqs_cos, freqs_sin
@torch.no_grad()
def __call__(self,
input,
num_samples=1,
cat_uc=True,
tuner_model=None,
**kwargs):
value_input = copy.deepcopy(self.input)
value_input.update(input)
print(value_input)
height, width = value_input['target_size_as_tuple']
value_output = copy.deepcopy(self.output)
# register tuner
if tuner_model is not None and tuner_model != '' and len(
tuner_model) > 0:
if not isinstance(tuner_model, list):
tuner_model = [tuner_model]
self.dynamic_load(self.diffusion_model, 'diffusion_model')
self.tuner_infer.register_tuner(tuner_model, self.diffusion_model,
cond_stage_model=None)
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=True)
# cond stage
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
function_name, dtype = self.get_function_info(self.cond_stage_model)
with torch.autocast(device_type='cuda', enabled=True, dtype=torch.bfloat16):
cont = getattr(get_model(self.cond_stage_model),
function_name)(value_input['prompt'], return_mask=False, use_mask=False)
null_cont = getattr(get_model(self.cond_stage_model),
function_name)(value_input['negative_prompt'] * num_samples, return_mask=False, use_mask=False)
self.dynamic_unload(self.cond_stage_model,
'cond_stage_model',
skip_loaded=True)
# get noise
seed = kwargs.pop('seed', -1)
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
generator = torch.Generator().manual_seed(seed)
if 'seed' in value_output:
value_output['seed'] = seed
for sample_id in range(num_samples):
if self.diffusion_model is not None:
noise_shape = (1,
(value_input['num_frames'] - 1) // self.diffusion_model['paras']['scale_factor_temporal'] + 1,
self.diffusion_model['paras']['latent_channels'],
height // self.diffusion_model['paras']['scale_factor_spatial'],
width // self.diffusion_model['paras']['scale_factor_spatial']
)
noise = torch.randn(noise_shape, generator=generator, dtype=getattr(torch, dtype), device='cpu').to(we.device_id)
self.dynamic_load(self.diffusion_model, 'diffusion_model')
image_rotary_emb = (
self._prepare_rotary_positional_embeddings(height, width, noise.size(1), we.device_id)
if self.diffusion_model['paras']['use_rotary_positional_embeddings']
else None
)
function_name, dtype = self.get_function_info(
self.diffusion_model)
with torch.autocast('cuda',
enabled=dtype=='bfloat16',
dtype=getattr(torch, dtype)):
solver_sample = value_input.get('sample', 'ddim')
sample_steps = value_input.get('sample_steps', 50)
guide_scale = value_input.get('guide_scale', 7.5)
guide_rescale = value_input.get('guide_rescale', 0.5)
latent = self.diffusion.sample(noise=noise,
sampler=solver_sample,
model=get_model(self.diffusion_model),
model_kwargs=[{
'cond': cont,
'image_latent': None,
'image_rotary_emb': image_rotary_emb,
}, {
'cond': null_cont,
'image_latent': None,
'image_rotary_emb': image_rotary_emb,
}],
steps=sample_steps,
show_progress=True,
use_dynamic_cfg=True,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
return_intermediate=None,
**kwargs).float()
self.dynamic_unload(self.diffusion_model,
'diffusion_model',
skip_loaded=True)
self.dynamic_load(self.first_stage_model, 'first_stage_model')
x_samples = self.decode_first_stage(latent).float() # [B, C, F, H, W]
self.dynamic_unload(self.first_stage_model,
'first_stage_model',
skip_loaded=True)
x_frames = torch.clamp(x_samples / 2 + 0.5, min=0.0, max=1.0)
if 'videos' in value_output:
if value_output['videos'] is None or (
isinstance(value_output['videos'], list)
and len(value_output['videos']) < 1):
value_output['videos'] = []
value_output['videos'].append(x_frames)
for k, v in value_output.items():
if isinstance(v, list):
value_output[k] = torch.cat(v, dim=0)
if isinstance(v, torch.Tensor):
value_output[k] = v.cpu()
# unregister tuner
if tuner_model is not None and tuner_model != '' and len(
tuner_model) > 0:
self.tuner_infer.unregister_tuner(tuner_model,
self.diffusion_model,
cond_stage_model=None)
return value_output
@@ -14,6 +14,7 @@ from scepter.modules.model.registry import (BACKBONES, EMBEDDERS, MODELS,
TOKENIZERS, DIFFUSIONS)
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.config import Config
from scepter.studio.utils.env import get_available_memory
from .control_inference import ControlInference
@@ -316,7 +317,8 @@ class DiffusionInference():
module_paras = {}
if cfg is not None:
self.paras = cfg.PARAS
self.input = {k.lower(): dict(v).get('DEFAULT', None) if isinstance(v, (dict, OrderedDict)) else v for k, v in cfg.INPUT.items()}
self.input_cfg = {k.lower(): v for k, v in cfg.INPUT.items()}
self.input = {k.lower(): dict(v).get('DEFAULT', None) if isinstance(v, (dict, OrderedDict, Config)) else v for k, v in cfg.INPUT.items()}
self.output = {k.lower(): v for k, v in cfg.OUTPUT.items()}
module_paras = cfg.MODULES_PARAS
return module_paras
+1 -1
View File
@@ -151,7 +151,7 @@ class FluxInference(DiffusionInference):
with torch.autocast('cuda',
enabled= dtype in ('float16', 'bfloat16'),
dtype=getattr(torch, dtype)):
solver_sample = value_input.get('sample', 'flow_eluer')
solver_sample = value_input.get('sample', 'flow_euler')
sample_steps = value_input.get('sample_steps', 20)
guide_scale = value_input.get('guide_scale', 3.5)
if guide_scale is not None:
+2 -2
View File
@@ -29,11 +29,11 @@ class TunerInference():
warnings.warn(f'Import swift error, please deal with this problem: {e}')
self.logger.info('Unloading tuner model')
if isinstance(diffusion_model['model'], SwiftModel):
if diffusion_model is not None and isinstance(diffusion_model['model'], SwiftModel):
for adapter_name in diffusion_model['model'].adapters:
diffusion_model['model'].deactivate_adapter(adapter_name,
offload='cpu')
if isinstance(cond_stage_model['model'], SwiftModel):
if cond_stage_model is not None and isinstance(cond_stage_model['model'], SwiftModel):
for adapter_name in cond_stage_model['model'].adapters:
cond_stage_model['model'].deactivate_adapter(adapter_name,
offload='cpu')
+1 -1
View File
@@ -1,4 +1,4 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.backbone import (ace, autoencoder, flux, image,
from scepter.modules.model.backbone import (ace, autoencoder, flux, image, cogvideox,
mmdit, pixart, unet, utils, video)
@@ -0,0 +1,3 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.backbone.cogvideox.cogvideox import CogVideoXTransformer3DModel
@@ -0,0 +1,319 @@
# -*- coding: utf-8 -*-
# Copyright 2024 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team.
# All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from collections import OrderedDict
from typing import Any, Dict, Optional, Tuple, Union
import torch
from torch import nn
from scepter.modules.model.base_model import BaseModel
from scepter.modules.model.registry import BACKBONES
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from .layers import CogVideoXBlock, CogVideoXPatchEmbed, TimestepEmbedding, Timesteps, AdaLayerNorm
@BACKBONES.register_class()
class CogVideoXTransformer3DModel(BaseModel):
"""
A Transformer model for video-like data in [CogVideoX](https://github.com/THUDM/CogVideo).
Parameters:
num_attention_heads (`int`, defaults to `30`):
The number of heads to use for multi-head attention.
attention_head_dim (`int`, defaults to `64`):
The number of channels in each head.
in_channels (`int`, defaults to `16`):
The number of channels in the input.
out_channels (`int`, *optional*, defaults to `16`):
The number of channels in the output.
flip_sin_to_cos (`bool`, defaults to `True`):
Whether to flip the sin to cos in the time embedding.
time_embed_dim (`int`, defaults to `512`):
Output dimension of timestep embeddings.
text_embed_dim (`int`, defaults to `4096`):
Input dimension of text embeddings from the text encoder.
num_layers (`int`, defaults to `30`):
The number of layers of Transformer blocks to use.
dropout (`float`, defaults to `0.0`):
The dropout probability to use.
attention_bias (`bool`, defaults to `True`):
Whether or not to use bias in the attention projection layers.
sample_width (`int`, defaults to `90`):
The width of the input latents.
sample_height (`int`, defaults to `60`):
The height of the input latents.
sample_frames (`int`, defaults to `49`):
The number of frames in the input latents. Note that this parameter was incorrectly initialized to 49
instead of 13 because CogVideoX processed 13 latent frames at once in its default and recommended settings,
but cannot be changed to the correct value to ensure backwards compatibility. To create a transformer with
K latent frames, the correct value to pass here would be: ((K - 1) * temporal_compression_ratio + 1).
patch_size (`int`, defaults to `2`):
The size of the patches to use in the patch embedding layer.
temporal_compression_ratio (`int`, defaults to `4`):
The compression ratio across the temporal dimension. See documentation for `sample_frames`.
max_text_seq_length (`int`, defaults to `226`):
The maximum sequence length of the input text embeddings.
activation_fn (`str`, defaults to `"gelu-approximate"`):
Activation function to use in feed-forward.
timestep_activation_fn (`str`, defaults to `"silu"`):
Activation function to use when generating the timestep embeddings.
norm_elementwise_affine (`bool`, defaults to `True`):
Whether or not to use elementwise affine in normalization layers.
norm_eps (`float`, defaults to `1e-5`):
The epsilon value to use in normalization layers.
spatial_interpolation_scale (`float`, defaults to `1.875`):
Scaling factor to apply in 3D positional embeddings across spatial dimensions.
temporal_interpolation_scale (`float`, defaults to `1.0`):
Scaling factor to apply in 3D positional embeddings across temporal dimensions.
"""
def __init__(
self,
cfg,
logger=None
):
super().__init__(cfg, logger=logger)
num_attention_heads = cfg.get("NUM_ATTENTION_HEADS", 30)
attention_head_dim = cfg.get("ATTENTION_HEAD_DIM", 64)
in_channels = cfg.get("IN_CHANNELS", 16)
out_channels = cfg.get("OUT_CHANNELS", 16)
flip_sin_to_cos = cfg.get("FLIP_SIN_TO_COS", True)
freq_shift = cfg.get("FREQ_SHIFT", 0)
time_embed_dim = cfg.get("TIME_EMBED_DIM", 512)
text_embed_dim = cfg.get("TEXT_EMBED_DIM", 4096)
num_layers = cfg.get("NUM_LAYERS", 30)
dropout = cfg.get("DROPOUT", 0.0)
attention_bias = cfg.get("ATTENTION_BIAS", True)
sample_width = cfg.get("SAMPLE_WIDTH", 90)
sample_height = cfg.get("SAMPLE_HEIGHT", 60)
sample_frames = cfg.get("SAMPLE_FRAMES", 49)
patch_size = cfg.get("PATCH_SIZE", 2)
temporal_compression_ratio = cfg.get("TEMPORAL_COMPRESSION_RATIO", 4)
max_text_seq_length = cfg.get("MAX_TEXT_SEQ_LENGTH", 226)
activation_fn = cfg.get("ACTIVATION_FN", "gelu-approximate")
timestep_activation_fn = cfg.get("TIMESTEP_ACTIVATION_FN", "silu")
norm_elementwise_affine = cfg.get("NORM_ELEMENTWISE_AFFINE", True)
norm_eps = cfg.get("NORM_EPS", 1e-5)
spatial_interpolation_scale = cfg.get("SPATIAL_INTERPOLATION_SCALE", 1.875)
temporal_interpolation_scale = cfg.get("TEMPORAL_INTERPOLATION_SCALE", 1.0)
use_rotary_positional_embeddings = cfg.get("USE_ROTARY_POSITIONAL_EMBEDDINGS", False)
use_learned_positional_embeddings = cfg.get("USE_LEARNED_POSITIONAL_EMBEDDINGS", False)
self.gradient_checkpointing = cfg.get("GRADIENT_CHECKPOINTING", False)
inner_dim = num_attention_heads * attention_head_dim
self.patch_size = patch_size
self.use_rotary_positional_embeddings = use_rotary_positional_embeddings
if not use_rotary_positional_embeddings and use_learned_positional_embeddings:
raise ValueError(
"There are no CogVideoX checkpoints available with disable rotary embeddings and learned positional "
"embeddings. If you're using a custom model and/or believe this should be supported, please open an "
"issue at https://github.com/huggingface/diffusers/issues."
)
# 1. Patch embedding
self.patch_embed = CogVideoXPatchEmbed(
patch_size=patch_size,
in_channels=in_channels,
embed_dim=inner_dim,
text_embed_dim=text_embed_dim,
bias=True,
sample_width=sample_width,
sample_height=sample_height,
sample_frames=sample_frames,
temporal_compression_ratio=temporal_compression_ratio,
max_text_seq_length=max_text_seq_length,
spatial_interpolation_scale=spatial_interpolation_scale,
temporal_interpolation_scale=temporal_interpolation_scale,
use_positional_embeddings=not use_rotary_positional_embeddings,
use_learned_positional_embeddings=use_learned_positional_embeddings,
)
self.embedding_dropout = nn.Dropout(dropout)
# 2. Time embeddings
self.time_proj = Timesteps(inner_dim, flip_sin_to_cos, freq_shift)
self.time_embedding = TimestepEmbedding(inner_dim, time_embed_dim, timestep_activation_fn)
# 3. Define spatio-temporal transformers blocks
self.transformer_blocks = nn.ModuleList(
[
CogVideoXBlock(
dim=inner_dim,
num_attention_heads=num_attention_heads,
attention_head_dim=attention_head_dim,
time_embed_dim=time_embed_dim,
dropout=dropout,
activation_fn=activation_fn,
attention_bias=attention_bias,
norm_elementwise_affine=norm_elementwise_affine,
norm_eps=norm_eps,
)
for _ in range(num_layers)
]
)
self.norm_final = nn.LayerNorm(inner_dim, norm_eps, norm_elementwise_affine)
# 4. Output blocks
self.norm_out = AdaLayerNorm(
embedding_dim=time_embed_dim,
output_dim=2 * inner_dim,
norm_elementwise_affine=norm_elementwise_affine,
norm_eps=norm_eps,
chunk_dim=1,
)
self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * out_channels)
def forward(
self,
x: torch.Tensor = None,
t: Union[int, float, torch.LongTensor] = None,
cond: torch.Tensor = None,
timestep_cond: Optional[torch.Tensor] = None,
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
**kwargs
):
if 'image_latent' in kwargs and kwargs['image_latent'] is not None:
hidden_states = torch.cat([x, kwargs['image_latent']], dim=2)
else:
hidden_states = x
timestep = t
encoder_hidden_states = cond
batch_size, num_frames, channels, height, width = hidden_states.shape
# 1. Time embedding
timesteps = timestep
t_emb = self.time_proj(timesteps)
# timesteps does not contain any weights and will always return f32 tensors
# but time_embedding might actually be running in fp16. so we need to cast here.
# there might be better ways to encapsulate this.
t_emb = t_emb.to(dtype=encoder_hidden_states.dtype)
emb = self.time_embedding(t_emb, timestep_cond)
# 2. Patch embedding
hidden_states = self.patch_embed(encoder_hidden_states, hidden_states)
hidden_states = self.embedding_dropout(hidden_states)
text_seq_length = encoder_hidden_states.shape[1]
encoder_hidden_states = hidden_states[:, :text_seq_length]
hidden_states = hidden_states[:, text_seq_length:]
# 3. Transformer blocks
for i, block in enumerate(self.transformer_blocks):
if self.training and self.gradient_checkpointing:
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False}
hidden_states, encoder_hidden_states = torch.utils.checkpoint.checkpoint(
create_custom_forward(block),
hidden_states,
encoder_hidden_states,
emb,
image_rotary_emb,
**ckpt_kwargs,
)
else:
hidden_states, encoder_hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
temb=emb,
image_rotary_emb=image_rotary_emb,
)
if not self.use_rotary_positional_embeddings:
# CogVideoX-2B
hidden_states = self.norm_final(hidden_states)
else:
# CogVideoX-5B
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
hidden_states = self.norm_final(hidden_states)
hidden_states = hidden_states[:, text_seq_length:]
# 4. Final block
hidden_states = self.norm_out(hidden_states, temb=emb)
hidden_states = self.proj_out(hidden_states)
# 5. Unpatchify
# Note: we use `-1` instead of `channels`:
# - It is okay to `channels` use for CogVideoX-2b and CogVideoX-5b (number of input channels is equal to output channels)
# - However, for CogVideoX-5b-I2V also takes concatenated input image latents (number of input channels is twice the output channels)
p = self.patch_size
output = hidden_states.reshape(batch_size, num_frames, height // p, width // p, -1, p, p)
output = output.permute(0, 1, 4, 2, 5, 3, 6).flatten(5, 6).flatten(3, 4)
return output
def load_pretrained_model(self, pretrained_model):
if pretrained_model is not None:
pretrained_model_list = [pretrained_model] if isinstance(pretrained_model, str) else pretrained_model
ckpt_all = OrderedDict()
for pretrained_model in pretrained_model_list:
with FS.get_from(pretrained_model,
wait_finish=True) as local_model:
if local_model.endswith('safetensors'):
from safetensors.torch import load_file as load_safetensors
ckpt = load_safetensors(local_model)
else:
ckpt = torch.load(local_model, map_location='cpu')
ckpt_all.update(ckpt)
missing, unexpected = self.load_state_dict(ckpt_all, strict=False)
if we.rank == 0:
self.logger.info(
f'Restored from {pretrained_model_list} with {len(missing)} missing and {len(unexpected)} unexpected keys'
)
if len(missing) > 0:
self.logger.info(f'Missing Keys:\n {missing}')
if len(unexpected) > 0:
self.logger.info(f'\nUnexpected Keys:\n {unexpected}')
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
CogVideoXTransformer3DModel.para_dict,
set_name=True)
if __name__ == "__main__":
import argparse
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.config import Config
from scepter.modules.utils.logger import get_logger
parser = argparse.ArgumentParser()
cfg = Config(parser_ins=parser)
for file_sys in cfg.FILE_SYSTEM:
FS.init_fs_client(file_sys)
model = BACKBONES.build(cfg.DIFFUSION_MODEL, logger=get_logger()).eval().requires_grad_(False).to('cuda').to(torch.bfloat16)
hidden_states = torch.load(FS.get_from(cfg.HIDDEN_STATES))
encoder_hidden_states = torch.load(FS.get_from(cfg.ENCODER_HIDDEN_STATES))
timestep = torch.load(FS.get_from(cfg.TIMESTEP))
timestep_cond = None
image_rotary_emb = None
attention_kwargs = None
output = model(hidden_states, encoder_hidden_states, timestep, timestep_cond, image_rotary_emb, attention_kwargs)
print(output, torch.sum(output))
@@ -0,0 +1,554 @@
# -*- coding: utf-8 -*-
# Copyright 2024 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team.
# All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Optional, Tuple
import torch
from torch import nn
import torch.nn.functional as F
from .utils import get_activation, get_timestep_embedding, get_3d_sincos_pos_embed, apply_rotary_emb
from .utils import GELU, GEGLU, ApproximateGELU, SwiGLU
class TimestepEmbedding(nn.Module):
def __init__(
self,
in_channels: int,
time_embed_dim: int,
act_fn: str = "silu",
out_dim: int = None,
post_act_fn: Optional[str] = None,
cond_proj_dim=None,
sample_proj_bias=True,
):
super().__init__()
self.linear_1 = nn.Linear(in_channels, time_embed_dim, sample_proj_bias)
if cond_proj_dim is not None:
self.cond_proj = nn.Linear(cond_proj_dim, in_channels, bias=False)
else:
self.cond_proj = None
self.act = get_activation(act_fn)
if out_dim is not None:
time_embed_dim_out = out_dim
else:
time_embed_dim_out = time_embed_dim
self.linear_2 = nn.Linear(time_embed_dim, time_embed_dim_out, sample_proj_bias)
if post_act_fn is None:
self.post_act = None
else:
self.post_act = get_activation(post_act_fn)
def forward(self, sample, condition=None):
if condition is not None:
sample = sample + self.cond_proj(condition)
sample = self.linear_1(sample)
if self.act is not None:
sample = self.act(sample)
sample = self.linear_2(sample)
if self.post_act is not None:
sample = self.post_act(sample)
return sample
class Timesteps(nn.Module):
def __init__(self, num_channels: int, flip_sin_to_cos: bool, downscale_freq_shift: float, scale: int = 1):
super().__init__()
self.num_channels = num_channels
self.flip_sin_to_cos = flip_sin_to_cos
self.downscale_freq_shift = downscale_freq_shift
self.scale = scale
def forward(self, timesteps):
t_emb = get_timestep_embedding(
timesteps,
self.num_channels,
flip_sin_to_cos=self.flip_sin_to_cos,
downscale_freq_shift=self.downscale_freq_shift,
scale=self.scale,
)
return t_emb
class CogVideoXLayerNormZero(nn.Module):
def __init__(
self,
conditioning_dim: int,
embedding_dim: int,
elementwise_affine: bool = True,
eps: float = 1e-5,
bias: bool = True,
) -> None:
super().__init__()
self.silu = nn.SiLU()
self.linear = nn.Linear(conditioning_dim, 6 * embedding_dim, bias=bias)
self.norm = nn.LayerNorm(embedding_dim, eps=eps, elementwise_affine=elementwise_affine)
def forward(
self, hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor, temb: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
shift, scale, gate, enc_shift, enc_scale, enc_gate = self.linear(self.silu(temb)).chunk(6, dim=1)
hidden_states = self.norm(hidden_states) * (1 + scale)[:, None, :] + shift[:, None, :]
encoder_hidden_states = self.norm(encoder_hidden_states) * (1 + enc_scale)[:, None, :] + enc_shift[:, None, :]
return hidden_states, encoder_hidden_states, gate[:, None, :], enc_gate[:, None, :]
class AdaLayerNorm(nn.Module):
r"""
Norm layer modified to incorporate timestep embeddings.
Parameters:
embedding_dim (`int`): The size of each embedding vector.
num_embeddings (`int`, *optional*): The size of the embeddings dictionary.
output_dim (`int`, *optional*):
norm_elementwise_affine (`bool`, defaults to `False):
norm_eps (`bool`, defaults to `False`):
chunk_dim (`int`, defaults to `0`):
"""
def __init__(
self,
embedding_dim: int,
num_embeddings: Optional[int] = None,
output_dim: Optional[int] = None,
norm_elementwise_affine: bool = False,
norm_eps: float = 1e-5,
chunk_dim: int = 0,
):
super().__init__()
self.chunk_dim = chunk_dim
output_dim = output_dim or embedding_dim * 2
if num_embeddings is not None:
self.emb = nn.Embedding(num_embeddings, embedding_dim)
else:
self.emb = None
self.silu = nn.SiLU()
self.linear = nn.Linear(embedding_dim, output_dim)
self.norm = nn.LayerNorm(output_dim // 2, norm_eps, norm_elementwise_affine)
def forward(
self, x: torch.Tensor, timestep: Optional[torch.Tensor] = None, temb: Optional[torch.Tensor] = None
) -> torch.Tensor:
if self.emb is not None:
temb = self.emb(timestep)
temb = self.linear(self.silu(temb))
if self.chunk_dim == 1:
# This is a bit weird why we have the order of "shift, scale" here and "scale, shift" in the
# other if-branch. This branch is specific to CogVideoX for now.
shift, scale = temb.chunk(2, dim=1)
shift = shift[:, None, :]
scale = scale[:, None, :]
else:
scale, shift = temb.chunk(2, dim=0)
x = self.norm(x) * (1 + scale) + shift
return x
class CogVideoXPatchEmbed(nn.Module):
def __init__(
self,
patch_size: int = 2,
in_channels: int = 16,
embed_dim: int = 1920,
text_embed_dim: int = 4096,
bias: bool = True,
sample_width: int = 90,
sample_height: int = 60,
sample_frames: int = 49,
temporal_compression_ratio: int = 4,
max_text_seq_length: int = 226,
spatial_interpolation_scale: float = 1.875,
temporal_interpolation_scale: float = 1.0,
use_positional_embeddings: bool = True,
use_learned_positional_embeddings: bool = True,
) -> None:
super().__init__()
self.patch_size = patch_size
self.embed_dim = embed_dim
self.sample_height = sample_height
self.sample_width = sample_width
self.sample_frames = sample_frames
self.temporal_compression_ratio = temporal_compression_ratio
self.max_text_seq_length = max_text_seq_length
self.spatial_interpolation_scale = spatial_interpolation_scale
self.temporal_interpolation_scale = temporal_interpolation_scale
self.use_positional_embeddings = use_positional_embeddings
self.use_learned_positional_embeddings = use_learned_positional_embeddings
self.proj = nn.Conv2d(
in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size, bias=bias
)
self.text_proj = nn.Linear(text_embed_dim, embed_dim)
if use_positional_embeddings or use_learned_positional_embeddings:
persistent = use_learned_positional_embeddings
pos_embedding = self._get_positional_embeddings(sample_height, sample_width, sample_frames)
self.register_buffer("pos_embedding", pos_embedding, persistent=persistent)
def _get_positional_embeddings(self, sample_height: int, sample_width: int, sample_frames: int) -> torch.Tensor:
post_patch_height = sample_height // self.patch_size
post_patch_width = sample_width // self.patch_size
post_time_compression_frames = (sample_frames - 1) // self.temporal_compression_ratio + 1
num_patches = post_patch_height * post_patch_width * post_time_compression_frames
pos_embedding = get_3d_sincos_pos_embed(
self.embed_dim,
(post_patch_width, post_patch_height),
post_time_compression_frames,
self.spatial_interpolation_scale,
self.temporal_interpolation_scale,
)
pos_embedding = torch.from_numpy(pos_embedding).flatten(0, 1)
joint_pos_embedding = torch.zeros(
1, self.max_text_seq_length + num_patches, self.embed_dim, requires_grad=False
)
joint_pos_embedding.data[:, self.max_text_seq_length :].copy_(pos_embedding)
return joint_pos_embedding
def forward(self, text_embeds: torch.Tensor, image_embeds: torch.Tensor):
r"""
Args:
text_embeds (`torch.Tensor`):
Input text embeddings. Expected shape: (batch_size, seq_length, embedding_dim).
image_embeds (`torch.Tensor`):
Input image embeddings. Expected shape: (batch_size, num_frames, channels, height, width).
"""
text_embeds = self.text_proj(text_embeds)
batch, num_frames, channels, height, width = image_embeds.shape
image_embeds = image_embeds.reshape(-1, channels, height, width)
image_embeds = self.proj(image_embeds)
image_embeds = image_embeds.view(batch, num_frames, *image_embeds.shape[1:])
image_embeds = image_embeds.flatten(3).transpose(2, 3) # [batch, num_frames, height x width, channels]
image_embeds = image_embeds.flatten(1, 2) # [batch, num_frames x height x width, channels]
embeds = torch.cat(
[text_embeds, image_embeds], dim=1
).contiguous() # [batch, seq_length + num_frames x height x width, channels]
if self.use_positional_embeddings or self.use_learned_positional_embeddings:
if self.use_learned_positional_embeddings and (self.sample_width != width or self.sample_height != height):
raise ValueError(
"It is currently not possible to generate videos at a different resolution that the defaults. This should only be the case with 'THUDM/CogVideoX-5b-I2V'."
"If you think this is incorrect, please open an issue at https://github.com/huggingface/diffusers/issues."
)
pre_time_compression_frames = (num_frames - 1) * self.temporal_compression_ratio + 1
if (
self.sample_height != height
or self.sample_width != width
or self.sample_frames != pre_time_compression_frames
):
pos_embedding = self._get_positional_embeddings(height, width, pre_time_compression_frames)
pos_embedding = pos_embedding.to(embeds.device, dtype=embeds.dtype)
else:
pos_embedding = self.pos_embedding
embeds = embeds + pos_embedding
return embeds
class FeedForward(nn.Module):
r"""
A feed-forward layer.
Parameters:
dim (`int`): The number of channels in the input.
dim_out (`int`, *optional*): The number of channels in the output. If not given, defaults to `dim`.
mult (`int`, *optional*, defaults to 4): The multiplier to use for the hidden dimension.
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
final_dropout (`bool` *optional*, defaults to False): Apply a final dropout.
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
"""
def __init__(
self,
dim: int,
dim_out: Optional[int] = None,
mult: int = 4,
dropout: float = 0.0,
activation_fn: str = "geglu",
final_dropout: bool = False,
inner_dim=None,
bias: bool = True,
):
super().__init__()
if inner_dim is None:
inner_dim = int(dim * mult)
dim_out = dim_out if dim_out is not None else dim
if activation_fn == "gelu":
act_fn = GELU(dim, inner_dim, bias=bias)
if activation_fn == "gelu-approximate":
act_fn = GELU(dim, inner_dim, approximate="tanh", bias=bias)
elif activation_fn == "geglu":
act_fn = GEGLU(dim, inner_dim, bias=bias)
elif activation_fn == "geglu-approximate":
act_fn = ApproximateGELU(dim, inner_dim, bias=bias)
elif activation_fn == "swiglu":
act_fn = SwiGLU(dim, inner_dim, bias=bias)
self.net = nn.ModuleList([])
# project in
self.net.append(act_fn)
# project dropout
self.net.append(nn.Dropout(dropout))
# project out
self.net.append(nn.Linear(inner_dim, dim_out, bias=bias))
# FF as used in Vision Transformer, MLP-Mixer, etc. have a final dropout
if final_dropout:
self.net.append(nn.Dropout(dropout))
def forward(self, hidden_states: torch.Tensor, *args, **kwargs) -> torch.Tensor:
if len(args) > 0 or kwargs.get("scale", None) is not None:
deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`."
print(deprecation_message)
for module in self.net:
hidden_states = module(hidden_states)
return hidden_states
class Attention(nn.Module):
def __init__(
self,
query_dim: int,
dim_head: int = 64,
heads: int = 8,
kv_heads: Optional[int] = None,
qk_norm: Optional[str] = None,
eps: float = 1e-5,
bias: bool = False,
out_bias: bool = True,
dropout: float = 0.0,
out_dim: int = None,
cross_attention_dim: Optional[int] = None,
):
super().__init__()
self.inner_dim = out_dim if out_dim is not None else dim_head * heads
self.inner_kv_dim = self.inner_dim if kv_heads is None else dim_head * kv_heads
self.query_dim = query_dim
self.cross_attention_dim = cross_attention_dim if cross_attention_dim is not None else query_dim
self.is_cross_attention = cross_attention_dim is not None
self.out_dim = out_dim if out_dim is not None else query_dim
self.heads = out_dim // dim_head if out_dim is not None else heads
if qk_norm is None:
self.norm_q = None
self.norm_k = None
elif qk_norm == "layer_norm":
self.norm_q = nn.LayerNorm(dim_head, eps=eps)
self.norm_k = nn.LayerNorm(dim_head, eps=eps)
self.to_q = nn.Linear(query_dim, self.inner_dim, bias=bias)
self.to_k = nn.Linear(self.cross_attention_dim, self.inner_kv_dim, bias=bias)
self.to_v = nn.Linear(self.cross_attention_dim, self.inner_kv_dim, bias=bias)
self.to_out = nn.ModuleList([])
self.to_out.append(nn.Linear(self.inner_dim, self.out_dim, bias=out_bias))
self.to_out.append(nn.Dropout(dropout))
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
image_rotary_emb: Optional[torch.Tensor] = None,
) -> torch.Tensor:
text_seq_length = encoder_hidden_states.size(1)
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
batch_size, sequence_length, _ = (
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
)
query = self.to_q(hidden_states)
key = self.to_k(hidden_states)
value = self.to_v(hidden_states)
inner_dim = key.shape[-1]
head_dim = inner_dim // self.heads
query = query.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
key = key.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
value = value.view(batch_size, -1, self.heads, head_dim).transpose(1, 2)
if self.norm_q is not None:
query = self.norm_q(query)
if self.norm_k is not None:
key = self.norm_k(key)
# Apply RoPE if needed
if image_rotary_emb is not None:
query[:, :, text_seq_length:] = apply_rotary_emb(query[:, :, text_seq_length:], image_rotary_emb)
if not self.is_cross_attention:
key[:, :, text_seq_length:] = apply_rotary_emb(key[:, :, text_seq_length:], image_rotary_emb)
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.heads * head_dim)
# linear proj
hidden_states = self.to_out[0](hidden_states)
# dropout
hidden_states = self.to_out[1](hidden_states)
encoder_hidden_states, hidden_states = hidden_states.split(
[text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
)
return hidden_states, encoder_hidden_states
class CogVideoXBlock(nn.Module):
r"""
Transformer block used in [CogVideoX](https://github.com/THUDM/CogVideo) model.
Parameters:
dim (`int`):
The number of channels in the input and output.
num_attention_heads (`int`):
The number of heads to use for multi-head attention.
attention_head_dim (`int`):
The number of channels in each head.
time_embed_dim (`int`):
The number of channels in timestep embedding.
dropout (`float`, defaults to `0.0`):
The dropout probability to use.
activation_fn (`str`, defaults to `"gelu-approximate"`):
Activation function to be used in feed-forward.
attention_bias (`bool`, defaults to `False`):
Whether or not to use bias in attention projection layers.
qk_norm (`bool`, defaults to `True`):
Whether or not to use normalization after query and key projections in Attention.
norm_elementwise_affine (`bool`, defaults to `True`):
Whether to use learnable elementwise affine parameters for normalization.
norm_eps (`float`, defaults to `1e-5`):
Epsilon value for normalization layers.
final_dropout (`bool` defaults to `False`):
Whether to apply a final dropout after the last feed-forward layer.
ff_inner_dim (`int`, *optional*, defaults to `None`):
Custom hidden dimension of Feed-forward layer. If not provided, `4 * dim` is used.
ff_bias (`bool`, defaults to `True`):
Whether or not to use bias in Feed-forward layer.
attention_out_bias (`bool`, defaults to `True`):
Whether or not to use bias in Attention output projection layer.
"""
def __init__(
self,
dim: int,
num_attention_heads: int,
attention_head_dim: int,
time_embed_dim: int,
dropout: float = 0.0,
activation_fn: str = "gelu-approximate",
attention_bias: bool = False,
qk_norm: bool = True,
norm_elementwise_affine: bool = True,
norm_eps: float = 1e-5,
final_dropout: bool = True,
ff_inner_dim: Optional[int] = None,
ff_bias: bool = True,
attention_out_bias: bool = True,
):
super().__init__()
# 1. Self Attention
self.norm1 = CogVideoXLayerNormZero(time_embed_dim, dim, norm_elementwise_affine, norm_eps, bias=True)
self.attn1 = Attention(
query_dim=dim,
dim_head=attention_head_dim,
heads=num_attention_heads,
qk_norm="layer_norm" if qk_norm else None,
eps=1e-6,
bias=attention_bias,
out_bias=attention_out_bias
)
# 2. Feed Forward
self.norm2 = CogVideoXLayerNormZero(time_embed_dim, dim, norm_elementwise_affine, norm_eps, bias=True)
self.ff = FeedForward(
dim,
dropout=dropout,
activation_fn=activation_fn,
final_dropout=final_dropout,
inner_dim=ff_inner_dim,
bias=ff_bias,
)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
) -> torch.Tensor:
text_seq_length = encoder_hidden_states.size(1)
# norm & modulate
norm_hidden_states, norm_encoder_hidden_states, gate_msa, enc_gate_msa = self.norm1(
hidden_states, encoder_hidden_states, temb
)
# attention
attn_hidden_states, attn_encoder_hidden_states = self.attn1(
hidden_states=norm_hidden_states,
encoder_hidden_states=norm_encoder_hidden_states,
image_rotary_emb=image_rotary_emb,
)
hidden_states = hidden_states + gate_msa * attn_hidden_states
encoder_hidden_states = encoder_hidden_states + enc_gate_msa * attn_encoder_hidden_states
# norm & modulate
norm_hidden_states, norm_encoder_hidden_states, gate_ff, enc_gate_ff = self.norm2(
hidden_states, encoder_hidden_states, temb
)
# feed-forward
norm_hidden_states = torch.cat([norm_encoder_hidden_states, norm_hidden_states], dim=1)
ff_output = self.ff(norm_hidden_states)
hidden_states = hidden_states + gate_ff * ff_output[:, text_seq_length:]
encoder_hidden_states = encoder_hidden_states + enc_gate_ff * ff_output[:, :text_seq_length]
return hidden_states, encoder_hidden_states
@@ -0,0 +1,544 @@
# -*- coding: utf-8 -*-
# Copyright 2024 The CogVideoX team, Tsinghua University & ZhipuAI and The HuggingFace Team.
# All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import math
from typing import Optional, Tuple, Union, List
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
ACTIVATION_FUNCTIONS = {
"swish": nn.SiLU(),
"silu": nn.SiLU(),
"mish": nn.Mish(),
"gelu": nn.GELU(),
"relu": nn.ReLU(),
}
def get_activation(act_fn: str) -> nn.Module:
"""Helper function to get activation function from string.
Args:
act_fn (str): Name of activation function.
Returns:
nn.Module: Activation function.
"""
act_fn = act_fn.lower()
if act_fn in ACTIVATION_FUNCTIONS:
return ACTIVATION_FUNCTIONS[act_fn]
else:
raise ValueError(f"Unsupported activation function: {act_fn}")
class FP32SiLU(nn.Module):
r"""
SiLU activation function with input upcasted to torch.float32.
"""
def __init__(self):
super().__init__()
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
return F.silu(inputs.float(), inplace=False).to(inputs.dtype)
class GELU(nn.Module):
r"""
GELU activation function with tanh approximation support with `approximate="tanh"`.
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
approximate (`str`, *optional*, defaults to `"none"`): If `"tanh"`, use tanh approximation.
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
"""
def __init__(self, dim_in: int, dim_out: int, approximate: str = "none", bias: bool = True):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out, bias=bias)
self.approximate = approximate
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
if gate.device.type != "mps":
return F.gelu(gate, approximate=self.approximate)
# mps: gelu is not implemented for float16
return F.gelu(gate.to(dtype=torch.float32), approximate=self.approximate).to(dtype=gate.dtype)
def forward(self, hidden_states):
hidden_states = self.proj(hidden_states)
hidden_states = self.gelu(hidden_states)
return hidden_states
class GEGLU(nn.Module):
r"""
A [variant](https://arxiv.org/abs/2002.05202) of the gated linear unit activation function.
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
"""
def __init__(self, dim_in: int, dim_out: int, bias: bool = True):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out * 2, bias=bias)
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
if gate.device.type != "mps":
return F.gelu(gate)
# mps: gelu is not implemented for float16
return F.gelu(gate.to(dtype=torch.float32)).to(dtype=gate.dtype)
def forward(self, hidden_states, *args, **kwargs):
if len(args) > 0 or kwargs.get("scale", None) is not None:
deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`."
print("scale", "1.0.0", deprecation_message)
hidden_states = self.proj(hidden_states)
hidden_states, gate = hidden_states.chunk(2, dim=-1)
return hidden_states * self.gelu(gate)
class SwiGLU(nn.Module):
r"""
A [variant](https://arxiv.org/abs/2002.05202) of the gated linear unit activation function. It's similar to `GEGLU`
but uses SiLU / Swish instead of GeLU.
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
"""
def __init__(self, dim_in: int, dim_out: int, bias: bool = True):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out * 2, bias=bias)
self.activation = nn.SiLU()
def forward(self, hidden_states):
hidden_states = self.proj(hidden_states)
hidden_states, gate = hidden_states.chunk(2, dim=-1)
return hidden_states * self.activation(gate)
class ApproximateGELU(nn.Module):
r"""
The approximate form of the Gaussian Error Linear Unit (GELU). For more details, see section 2 of this
[paper](https://arxiv.org/abs/1606.08415).
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
"""
def __init__(self, dim_in: int, dim_out: int, bias: bool = True):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out, bias=bias)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.proj(x)
return x * torch.sigmoid(1.702 * x)
def randn_tensor(
shape: Union[Tuple, List],
generator: Optional[Union[List["torch.Generator"], "torch.Generator"]] = None,
device: Optional["torch.device"] = None,
dtype: Optional["torch.dtype"] = None,
layout: Optional["torch.layout"] = None,
):
"""A helper function to create random tensors on the desired `device` with the desired `dtype`. When
passing a list of generators, you can seed each batch size individually. If CPU generators are passed, the tensor
is always created on the CPU.
"""
# device on which tensor is created defaults to device
rand_device = device
batch_size = shape[0]
layout = layout or torch.strided
device = device or torch.device("cpu")
if generator is not None:
gen_device_type = generator.device.type if not isinstance(generator, list) else generator[0].device.type
if gen_device_type != device.type and gen_device_type == "cpu":
rand_device = "cpu"
if device != "mps":
print(
f"The passed generator was created on 'cpu' even though a tensor on {device} was expected."
f" Tensors will be created on 'cpu' and then moved to {device}. Note that one can probably"
f" slighly speed up this function by passing a generator that was created on the {device} device."
)
elif gen_device_type != device.type and gen_device_type == "cuda":
raise ValueError(f"Cannot generate a {device} tensor from a generator of type {gen_device_type}.")
# make sure generator list of length 1 is treated like a non-list
if isinstance(generator, list) and len(generator) == 1:
generator = generator[0]
if isinstance(generator, list):
shape = (1,) + shape[1:]
latents = [
torch.randn(shape, generator=generator[i], device=rand_device, dtype=dtype, layout=layout)
for i in range(batch_size)
]
latents = torch.cat(latents, dim=0).to(device)
else:
latents = torch.randn(shape, generator=generator, device=rand_device, dtype=dtype, layout=layout).to(device)
return latents
def get_timestep_embedding(
timesteps: torch.Tensor,
embedding_dim: int,
flip_sin_to_cos: bool = False,
downscale_freq_shift: float = 1,
scale: float = 1,
max_period: int = 10000,
):
"""
This matches the implementation in Denoising Diffusion Probabilistic Models: Create sinusoidal timestep embeddings.
Args
timesteps (torch.Tensor):
a 1-D Tensor of N indices, one per batch element. These may be fractional.
embedding_dim (int):
the dimension of the output.
flip_sin_to_cos (bool):
Whether the embedding order should be `cos, sin` (if True) or `sin, cos` (if False)
downscale_freq_shift (float):
Controls the delta between frequencies between dimensions
scale (float):
Scaling factor applied to the embeddings.
max_period (int):
Controls the maximum frequency of the embeddings
Returns
torch.Tensor: an [N x dim] Tensor of positional embeddings.
"""
assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array"
half_dim = embedding_dim // 2
exponent = -math.log(max_period) * torch.arange(
start=0, end=half_dim, dtype=torch.float32, device=timesteps.device
)
exponent = exponent / (half_dim - downscale_freq_shift)
emb = torch.exp(exponent)
emb = timesteps[:, None].float() * emb[None, :]
# scale embeddings
emb = scale * emb
# concat sine and cosine embeddings
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)
# flip sine and cosine embeddings
if flip_sin_to_cos:
emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1)
# zero pad
if embedding_dim % 2 == 1:
emb = torch.nn.functional.pad(emb, (0, 1, 0, 0))
return emb
def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
"""
embed_dim: output dimension for each position pos: a list of positions to be encoded: size (M,) out: (M, D)
"""
if embed_dim % 2 != 0:
raise ValueError("embed_dim must be divisible by 2")
omega = np.arange(embed_dim // 2, dtype=np.float64)
omega /= embed_dim / 2.0
omega = 1.0 / 10000**omega # (D/2,)
pos = pos.reshape(-1) # (M,)
out = np.einsum("m,d->md", pos, omega) # (M, D/2), outer product
emb_sin = np.sin(out) # (M, D/2)
emb_cos = np.cos(out) # (M, D/2)
emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
return emb
def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):
if embed_dim % 2 != 0:
raise ValueError("embed_dim must be divisible by 2")
# use half of dimensions to encode grid_h
emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2)
emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2)
emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D)
return emb
def get_3d_sincos_pos_embed(
embed_dim: int,
spatial_size: Union[int, Tuple[int, int]],
temporal_size: int,
spatial_interpolation_scale: float = 1.0,
temporal_interpolation_scale: float = 1.0,
) -> np.ndarray:
r"""
Args:
embed_dim (`int`):
spatial_size (`int` or `Tuple[int, int]`):
temporal_size (`int`):
spatial_interpolation_scale (`float`, defaults to 1.0):
temporal_interpolation_scale (`float`, defaults to 1.0):
"""
if embed_dim % 4 != 0:
raise ValueError("`embed_dim` must be divisible by 4")
if isinstance(spatial_size, int):
spatial_size = (spatial_size, spatial_size)
embed_dim_spatial = 3 * embed_dim // 4
embed_dim_temporal = embed_dim // 4
# 1. Spatial
grid_h = np.arange(spatial_size[1], dtype=np.float32) / spatial_interpolation_scale
grid_w = np.arange(spatial_size[0], dtype=np.float32) / spatial_interpolation_scale
grid = np.meshgrid(grid_w, grid_h) # here w goes first
grid = np.stack(grid, axis=0)
grid = grid.reshape([2, 1, spatial_size[1], spatial_size[0]])
pos_embed_spatial = get_2d_sincos_pos_embed_from_grid(embed_dim_spatial, grid)
# 2. Temporal
grid_t = np.arange(temporal_size, dtype=np.float32) / temporal_interpolation_scale
pos_embed_temporal = get_1d_sincos_pos_embed_from_grid(embed_dim_temporal, grid_t)
# 3. Concat
pos_embed_spatial = pos_embed_spatial[np.newaxis, :, :]
pos_embed_spatial = np.repeat(pos_embed_spatial, temporal_size, axis=0) # [T, H*W, D // 4 * 3]
pos_embed_temporal = pos_embed_temporal[:, np.newaxis, :]
pos_embed_temporal = np.repeat(pos_embed_temporal, spatial_size[0] * spatial_size[1], axis=1) # [T, H*W, D // 4]
pos_embed = np.concatenate([pos_embed_temporal, pos_embed_spatial], axis=-1) # [T, H*W, D]
return pos_embed
def apply_rotary_emb(
x: torch.Tensor,
freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]],
use_real: bool = True,
use_real_unbind_dim: int = -1,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Apply rotary embeddings to input tensors using the given frequency tensor. This function applies rotary embeddings
to the given query or key 'x' tensors using the provided frequency tensor 'freqs_cis'. The input tensors are
reshaped as complex numbers, and the frequency tensor is reshaped for broadcasting compatibility. The resulting
tensors contain rotary embeddings and are returned as real tensors.
Args:
x (`torch.Tensor`):
Query or key tensor to apply rotary embeddings. [B, H, S, D] xk (torch.Tensor): Key tensor to apply
freqs_cis (`Tuple[torch.Tensor]`): Precomputed frequency tensor for complex exponentials. ([S, D], [S, D],)
Returns:
Tuple[torch.Tensor, torch.Tensor]: Tuple of modified query tensor and key tensor with rotary embeddings.
"""
if use_real:
cos, sin = freqs_cis # [S, D]
cos = cos[None, None]
sin = sin[None, None]
cos, sin = cos.to(x.device), sin.to(x.device)
if use_real_unbind_dim == -1:
# Used for flux, cogvideox, hunyuan-dit
x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, S, H, D//2]
x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3)
elif use_real_unbind_dim == -2:
# Used for Stable Audio
x_real, x_imag = x.reshape(*x.shape[:-1], 2, -1).unbind(-2) # [B, S, H, D//2]
x_rotated = torch.cat([-x_imag, x_real], dim=-1)
else:
raise ValueError(f"`use_real_unbind_dim={use_real_unbind_dim}` but should be -1 or -2.")
out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype)
return out
else:
# used for lumina
x_rotated = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
freqs_cis = freqs_cis.unsqueeze(2)
x_out = torch.view_as_real(x_rotated * freqs_cis).flatten(3)
return x_out.type_as(x)
def get_1d_rotary_pos_embed(
dim: int,
pos: Union[np.ndarray, int],
theta: float = 10000.0,
use_real=False,
linear_factor=1.0,
ntk_factor=1.0,
repeat_interleave_real=True,
freqs_dtype=torch.float32, # torch.float32, torch.float64 (flux)
):
"""
Precompute the frequency tensor for complex exponentials (cis) with given dimensions.
This function calculates a frequency tensor with complex exponentials using the given dimension 'dim' and the end
index 'end'. The 'theta' parameter scales the frequencies. The returned tensor contains complex values in complex64
data type.
Args:
dim (`int`): Dimension of the frequency tensor.
pos (`np.ndarray` or `int`): Position indices for the frequency tensor. [S] or scalar
theta (`float`, *optional*, defaults to 10000.0):
Scaling factor for frequency computation. Defaults to 10000.0.
use_real (`bool`, *optional*):
If True, return real part and imaginary part separately. Otherwise, return complex numbers.
linear_factor (`float`, *optional*, defaults to 1.0):
Scaling factor for the context extrapolation. Defaults to 1.0.
ntk_factor (`float`, *optional*, defaults to 1.0):
Scaling factor for the NTK-Aware RoPE. Defaults to 1.0.
repeat_interleave_real (`bool`, *optional*, defaults to `True`):
If `True` and `use_real`, real part and imaginary part are each interleaved with themselves to reach `dim`.
Otherwise, they are concateanted with themselves.
freqs_dtype (`torch.float32` or `torch.float64`, *optional*, defaults to `torch.float32`):
the dtype of the frequency tensor.
Returns:
`torch.Tensor`: Precomputed frequency tensor with complex exponentials. [S, D/2]
"""
assert dim % 2 == 0
if isinstance(pos, int):
pos = torch.arange(pos)
if isinstance(pos, np.ndarray):
pos = torch.from_numpy(pos) # type: ignore # [S]
theta = theta * ntk_factor
freqs = (
1.0
/ (theta ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=pos.device)[: (dim // 2)] / dim))
/ linear_factor
) # [D/2]
freqs = torch.outer(pos, freqs) # type: ignore # [S, D/2]
if use_real and repeat_interleave_real:
# flux, hunyuan-dit, cogvideox
freqs_cos = freqs.cos().repeat_interleave(2, dim=1).float() # [S, D]
freqs_sin = freqs.sin().repeat_interleave(2, dim=1).float() # [S, D]
return freqs_cos, freqs_sin
elif use_real:
# stable audio
freqs_cos = torch.cat([freqs.cos(), freqs.cos()], dim=-1).float() # [S, D]
freqs_sin = torch.cat([freqs.sin(), freqs.sin()], dim=-1).float() # [S, D]
return freqs_cos, freqs_sin
else:
# lumina
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 # [S, D/2]
return freqs_cis
def get_3d_rotary_pos_embed(
embed_dim, crops_coords, grid_size, temporal_size, theta: int = 10000, use_real: bool = True
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
"""
RoPE for video tokens with 3D structure.
Args:
embed_dim: (`int`):
The embedding dimension size, corresponding to hidden_size_head.
crops_coords (`Tuple[int]`):
The top-left and bottom-right coordinates of the crop.
grid_size (`Tuple[int]`):
The grid size of the spatial positional embedding (height, width).
temporal_size (`int`):
The size of the temporal dimension.
theta (`float`):
Scaling factor for frequency computation.
Returns:
`torch.Tensor`: positional embedding with shape `(temporal_size * grid_size[0] * grid_size[1], embed_dim/2)`.
"""
if use_real is not True:
raise ValueError(" `use_real = False` is not currently supported for get_3d_rotary_pos_embed")
start, stop = crops_coords
grid_size_h, grid_size_w = grid_size
grid_h = np.linspace(start[0], stop[0], grid_size_h, endpoint=False, dtype=np.float32)
grid_w = np.linspace(start[1], stop[1], grid_size_w, endpoint=False, dtype=np.float32)
grid_t = np.linspace(0, temporal_size, temporal_size, endpoint=False, dtype=np.float32)
# Compute dimensions for each axis
dim_t = embed_dim // 4
dim_h = embed_dim // 8 * 3
dim_w = embed_dim // 8 * 3
# Temporal frequencies
freqs_t = get_1d_rotary_pos_embed(dim_t, grid_t, use_real=True)
# Spatial frequencies for height and width
freqs_h = get_1d_rotary_pos_embed(dim_h, grid_h, use_real=True)
freqs_w = get_1d_rotary_pos_embed(dim_w, grid_w, use_real=True)
# BroadCast and concatenate temporal and spaial frequencie (height and width) into a 3d tensor
def combine_time_height_width(freqs_t, freqs_h, freqs_w):
freqs_t = freqs_t[:, None, None, :].expand(
-1, grid_size_h, grid_size_w, -1
) # temporal_size, grid_size_h, grid_size_w, dim_t
freqs_h = freqs_h[None, :, None, :].expand(
temporal_size, -1, grid_size_w, -1
) # temporal_size, grid_size_h, grid_size_2, dim_h
freqs_w = freqs_w[None, None, :, :].expand(
temporal_size, grid_size_h, -1, -1
) # temporal_size, grid_size_h, grid_size_2, dim_w
freqs = torch.cat(
[freqs_t, freqs_h, freqs_w], dim=-1
) # temporal_size, grid_size_h, grid_size_w, (dim_t + dim_h + dim_w)
freqs = freqs.view(
temporal_size * grid_size_h * grid_size_w, -1
) # (temporal_size * grid_size_h * grid_size_w), (dim_t + dim_h + dim_w)
return freqs
t_cos, t_sin = freqs_t # both t_cos and t_sin has shape: temporal_size, dim_t
h_cos, h_sin = freqs_h # both h_cos and h_sin has shape: grid_size_h, dim_h
w_cos, w_sin = freqs_w # both w_cos and w_sin has shape: grid_size_w, dim_w
cos = combine_time_height_width(t_cos, h_cos, w_cos)
sin = combine_time_height_width(t_sin, h_sin, w_sin)
return cos, sin
def get_resize_crop_region_for_grid(src, tgt_width, tgt_height):
tw = tgt_width
th = tgt_height
h, w = src
r = h / w
if r > (th / tw):
resize_height = th
resize_width = int(round(th / h * w))
else:
resize_width = tw
resize_height = int(round(tw / w * h))
crop_top = int(round((th - resize_height) / 2.0))
crop_left = int(round((tw - resize_width) / 2.0))
return (crop_top, crop_left), (crop_top + resize_height, crop_left + resize_width)
+128 -2
View File
@@ -12,7 +12,7 @@ from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from torch import Tensor, nn
from torch.utils.checkpoint import checkpoint_sequential
from torch.nn.utils.rnn import pad_sequence
from .layers import (DoubleStreamBlock, EmbedND, LastLayer, MLPEmbedder,
SingleStreamBlock, timestep_embedding)
@@ -245,7 +245,133 @@ class Flux(BaseModel):
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
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):
context, y = cond["context"].to(x), cond["y"].to(x)
batch_frames, batch_frames_ids = [], []
for ix, shape in zip(x, cond["x_shapes"]):
# unpack image from sequence
ix = ix[:, :shape[0] * shape[1]].view(-1, shape[0], shape[1])
c, h, w = ix.shape
ix = rearrange(ix, "c (h ph) (w pw) -> (h w) (c ph pw)", ph=2, pw=2)
ix_id = torch.zeros(h // 2, w // 2, 3)
ix_id[..., 1] = ix_id[..., 1] + torch.arange(h // 2)[:, None]
ix_id[..., 2] = ix_id[..., 2] + torch.arange(w // 2)[None, :]
ix_id = rearrange(ix_id, "h w c -> (h w) c")
batch_frames.append([ix])
batch_frames_ids.append([ix_id])
x_list, x_id_list, mask_x_list, x_seq_length = [], [], [], []
for frames, frame_ids in zip(batch_frames, batch_frames_ids):
proj_frames = []
for idx, one_frame in enumerate(frames):
one_frame = self.img_in(one_frame)
proj_frames.append(one_frame)
ix = torch.cat(proj_frames, dim=0)
if_id = torch.cat(frame_ids, dim=0)
x_list.append(ix)
x_id_list.append(if_id)
mask_x_list.append(torch.ones(ix.shape[0]).to(ix.device, non_blocking=True).bool())
x_seq_length.append(ix.shape[0])
x = pad_sequence(tuple(x_list), batch_first=True)
x_ids = pad_sequence(tuple(x_id_list), batch_first=True).to(x) # [b,pad_seq,2] pad (0.,0.) at dim2
mask_x = pad_sequence(tuple(mask_x_list), batch_first=True)
txt = self.txt_in(context)
txt_ids = torch.zeros(context.shape[0], context.shape[1], 3).to(x)
mask_txt = torch.ones(context.shape[0], context.shape[1]).to(x.device, non_blocking=True).bool()
return x, x_ids, txt, txt_ids, y, mask_x, mask_txt, x_seq_length
def unpack(self, x: Tensor, cond: dict = None, x_seq_length: list = None) -> Tensor:
x_list = []
image_shapes = cond["x_shapes"]
for u, shape, seq_length in zip(x, image_shapes, x_seq_length):
height, width = shape
h, w = math.ceil(height / 2), math.ceil(width / 2)
u = rearrange(
u[seq_length-h*w:seq_length, ...],
"(h w) (c ph pw) -> (h ph w pw) c",
h=h,
w=w,
ph=2,
pw=2,
)
x_list.append(u)
x = pad_sequence(tuple(x_list), batch_first=True).permute(0, 2, 1)
return x
def forward(
self,
x: Tensor,
t: Tensor,
cond: dict = {},
guidance: Tensor | None = None,
gc_seg: int = 0,
**kwargs
) -> Tensor:
x, x_ids, txt, txt_ids, y, mask_x, mask_txt, seq_length_list = self.prepare_input(x, cond)
# running on sequences img
vec = self.time_in(timestep_embedding(t, 256))
if self.guidance_embed:
if guidance is None:
raise ValueError("Didn't get guidance strength for guidance distilled model.")
vec = vec + self.guidance_in(timestep_embedding(guidance, 256))
vec = vec + self.vector_in(y)
ids = torch.cat((txt_ids, x_ids), dim=1)
pe = self.pe_embedder(ids)
mask_aside = torch.cat((mask_txt, mask_x), dim=1)
mask = mask_aside[:, None, :] * mask_aside[:, :, None]
kwargs = dict(
vec=vec,
pe=pe,
mask=mask,
txt_length = txt.shape[1],
)
x = torch.cat((txt, x), 1)
if self.use_grad_checkpoint and gc_seg >= 0:
x = checkpoint_sequential(
functions=[partial(block, **kwargs) for block in self.double_blocks],
segments=gc_seg if gc_seg > 0 else len(self.double_blocks),
input=x,
use_reentrant=False
)
else:
for block in self.double_blocks:
x = block(x, **kwargs)
kwargs = dict(
vec=vec,
pe=pe,
mask=mask,
)
if self.use_grad_checkpoint and gc_seg >= 0:
x = checkpoint_sequential(
functions=[partial(block, **kwargs) for block in self.single_blocks],
segments=gc_seg if gc_seg > 0 else len(self.single_blocks),
input=x,
use_reentrant=False
)
else:
for block in self.single_blocks:
x = block(x, **kwargs)
x = x[:, txt.shape[1]:, ...]
x = self.final_layer(x, vec) # (N, T, patch_size ** 2 * out_channels) 6 64 64
x = self.unpack(x, cond, seq_length_list)
return x
@staticmethod
def get_config_template():
return dict_to_yaml('BACKBONE',
__class__.__name__,
FluxMR.para_dict,
set_name=True)
+77 -63
View File
@@ -4,24 +4,66 @@ from __future__ import annotations
import math
from dataclasses import dataclass
from torch import Tensor, nn
import torch
from einops import rearrange, repeat
from torch import Tensor, nn
from torch import Tensor
from torch.nn.utils.rnn import pad_sequence
try:
from flash_attn import (
flash_attn_varlen_func
)
FLASHATTN_IS_AVAILABLE = True
except ImportError:
FLASHATTN_IS_AVAILABLE = False
flash_attn_varlen_func = None
def attention(q: Tensor,
k: Tensor,
v: Tensor,
pe: Tensor,
mask: Tensor | None = None) -> Tensor:
def attention(q: Tensor, k: Tensor, v: Tensor, pe: Tensor, mask: Tensor | None = None, backend = 'pytorch') -> Tensor:
q, k = apply_rope(q, k, pe)
x = torch.nn.functional.scaled_dot_product_attention(q,
k,
v,
attn_mask=mask)
x = torch.nan_to_num(x, nan=0.0, posinf=1e10, neginf=-1e10)
x = rearrange(x, 'B H L D -> B L (H D)')
if backend == 'pytorch':
if mask is not None and mask.dtype == torch.bool:
mask = torch.zeros_like(mask).to(q).masked_fill_(mask.logical_not(), -1e20)
x = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
# x = torch.nan_to_num(x, nan=0.0, posinf=1e10, neginf=-1e10)
x = rearrange(x, "B H L D -> B L (H D)")
elif backend == 'flash_attn':
# q: (B, H, L, D)
# k: (B, H, S, D) now L = S
# v: (B, H, S, D)
b, h, lq, d = q.shape
_, _, lk, _ = k.shape
q = rearrange(q, "B H L D -> B L H D")
k = rearrange(k, "B H S D -> B S H D")
v = rearrange(v, "B H S D -> B S H D")
if mask is None:
q_lens = torch.tensor([lq] * b, dtype=torch.int32).to(q.device, non_blocking=True)
k_lens = torch.tensor([lk] * b, dtype=torch.int32).to(k.device, non_blocking=True)
else:
q_lens = torch.sum(mask[:, 0, :, 0], dim=1).int()
k_lens = torch.sum(mask[:, 0, 0, :], dim=1).int()
q = torch.cat([q_v[:q_l] for q_v, q_l in zip(q, q_lens)])
k = torch.cat([k_v[:k_l] for k_v, k_l in zip(k, k_lens)])
v = torch.cat([v_v[:v_l] for v_v, v_l in zip(v, k_lens)])
cu_seqlens_q = torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(0, dtype=torch.int32)
cu_seqlens_k = torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(0, dtype=torch.int32)
max_seqlen_q = q_lens.max()
max_seqlen_k = k_lens.max()
x = flash_attn_varlen_func(
q,
k,
v,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
max_seqlen_q=max_seqlen_q,
max_seqlen_k=max_seqlen_k
)
x_list = [x[cu_seqlens_q[i]:cu_seqlens_q[i+1]] for i in range(b)]
x = pad_sequence(tuple(x_list), batch_first=True)
x = rearrange(x, "B L H D -> B L (H D)")
else:
raise NotImplementedError
return x
@@ -173,11 +215,8 @@ class Modulation(nn.Module):
self.multiplier = 6 if double else 3
self.lin = nn.Linear(dim, self.multiplier * dim, bias=True)
def forward(self,
vec: Tensor) -> tuple[ModulationOut, ModulationOut | None]:
out = self.lin(nn.functional.silu(vec))[:,
None, :].chunk(self.multiplier,
dim=-1)
def forward(self, vec: Tensor) -> tuple[ModulationOut, ModulationOut | None]:
out = self.lin(nn.functional.silu(vec))[:, None, :].chunk(self.multiplier, dim=-1)
return (
ModulationOut(*out[:3]),
@@ -186,56 +225,37 @@ class Modulation(nn.Module):
class DoubleStreamBlock(nn.Module):
def __init__(self,
hidden_size: int,
num_heads: int,
mlp_ratio: float,
qkv_bias: bool = False):
def __init__(self, hidden_size: int, num_heads: int, mlp_ratio: float, qkv_bias: bool = False, backend = 'pytorch'):
super().__init__()
mlp_hidden_dim = int(hidden_size * mlp_ratio)
self.num_heads = num_heads
self.hidden_size = hidden_size
self.img_mod = Modulation(hidden_size, double=True)
self.img_norm1 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.img_attn = SelfAttention(dim=hidden_size,
num_heads=num_heads,
qkv_bias=qkv_bias)
self.img_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.img_attn = SelfAttention(dim=hidden_size, num_heads=num_heads, qkv_bias=qkv_bias)
self.img_norm2 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.img_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.img_mlp = nn.Sequential(
nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
nn.GELU(approximate='tanh'),
nn.GELU(approximate="tanh"),
nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
)
self.backend = backend
self.txt_mod = Modulation(hidden_size, double=True)
self.txt_norm1 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.txt_attn = SelfAttention(dim=hidden_size,
num_heads=num_heads,
qkv_bias=qkv_bias)
self.txt_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.txt_attn = SelfAttention(dim=hidden_size, num_heads=num_heads, qkv_bias=qkv_bias)
self.txt_norm2 = nn.LayerNorm(hidden_size,
elementwise_affine=False,
eps=1e-6)
self.txt_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.txt_mlp = nn.Sequential(
nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
nn.GELU(approximate='tanh'),
nn.GELU(approximate="tanh"),
nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
)
def forward(self,
x: Tensor,
vec: Tensor,
pe: Tensor,
mask: Tensor = None,
txt_length=None):
def forward(self, x: Tensor, vec: Tensor, pe: Tensor, mask: Tensor = None, txt_length = None):
img_mod1, img_mod2 = self.img_mod(vec)
txt_mod1, txt_mod2 = self.txt_mod(vec)
@@ -245,19 +265,13 @@ class DoubleStreamBlock(nn.Module):
img_modulated = self.img_norm1(img)
img_modulated = (1 + img_mod1.scale) * img_modulated + img_mod1.shift
img_qkv = self.img_attn.qkv(img_modulated)
img_q, img_k, img_v = rearrange(img_qkv,
'B L (K H D) -> K B H L D',
K=3,
H=self.num_heads)
img_q, img_k, img_v = rearrange(img_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
img_q, img_k = self.img_attn.norm(img_q, img_k, img_v)
# prepare txt for attention
txt_modulated = self.txt_norm1(txt)
txt_modulated = (1 + txt_mod1.scale) * txt_modulated + txt_mod1.shift
txt_qkv = self.txt_attn.qkv(txt_modulated)
txt_q, txt_k, txt_v = rearrange(txt_qkv,
'B L (K H D) -> K B H L D',
K=3,
H=self.num_heads)
txt_q, txt_k, txt_v = rearrange(txt_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
txt_q, txt_k = self.txt_attn.norm(txt_q, txt_k, txt_v)
# run actual attention
@@ -266,18 +280,16 @@ class DoubleStreamBlock(nn.Module):
v = torch.cat((txt_v, img_v), dim=2)
if mask is not None:
mask = repeat(mask, 'B L S-> B H L S', H=self.num_heads)
attn = attention(q, k, v, pe=pe, mask=mask)
txt_attn, img_attn = attn[:, :txt.shape[1]], attn[:, txt.shape[1]:]
attn = attention(q, k, v, pe=pe, mask = mask, backend = self.backend)
txt_attn, img_attn = attn[:, : txt.shape[1]], attn[:, txt.shape[1] :]
# calculate the img bloks
img = img + img_mod1.gate * self.img_attn.proj(img_attn)
img = img + img_mod2.gate * self.img_mlp(
(1 + img_mod2.scale) * self.img_norm2(img) + img_mod2.shift)
img = img + img_mod2.gate * self.img_mlp((1 + img_mod2.scale) * self.img_norm2(img) + img_mod2.shift)
# calculate the txt bloks
txt = txt + txt_mod1.gate * self.txt_attn.proj(txt_attn)
txt = txt + txt_mod2.gate * self.txt_mlp(
(1 + txt_mod2.scale) * self.txt_norm2(txt) + txt_mod2.shift)
txt = txt + txt_mod2.gate * self.txt_mlp((1 + txt_mod2.scale) * self.txt_norm2(txt) + txt_mod2.shift)
x = torch.cat((txt, img), 1)
return x
@@ -293,6 +305,7 @@ class SingleStreamBlock(nn.Module):
num_heads: int,
mlp_ratio: float = 4.0,
qk_scale: float | None = None,
backend='pytorch'
):
super().__init__()
self.hidden_dim = hidden_size
@@ -317,6 +330,7 @@ class SingleStreamBlock(nn.Module):
self.mlp_act = nn.GELU(approximate='tanh')
self.modulation = Modulation(hidden_size, double=False)
self.backend = backend
def forward(self,
x: Tensor,
+17 -39
View File
@@ -19,10 +19,6 @@ class BaseDiffusion(object):
para_dict = {
'NOISE_SCHEDULER': {},
'SAMPLER_SCHEDULER': {},
'MIN_SNR_GAMMA': {
'value': None,
'description': 'The minimum SNR gamma value for the loss function.'
},
'PREDICTION_TYPE': {
'value': 'eps',
'description':
@@ -37,7 +33,6 @@ class BaseDiffusion(object):
self.init_params()
def init_params(self):
self.min_snr_gamma = self.cfg.get('MIN_SNR_GAMMA', None)
self.prediction_type = self.cfg.get('PREDICTION_TYPE', 'eps')
self.noise_scheduler = NOISE_SCHEDULERS.build(self.cfg.NOISE_SCHEDULER,
logger=self.logger)
@@ -67,17 +62,19 @@ class BaseDiffusion(object):
show_progress=False,
return_intermediate=None,
intermediate_callback=None,
reverse_scale = -1.,
x = None,
**kwargs):
assert isinstance(steps, (int, torch.LongTensor))
assert return_intermediate in (None, 'x0', 'xt')
assert isinstance(sampler, (str, dict, Config))
intermediates = []
def callback_fn(x_t, t, sigma=None, alpha=None):
def callback_fn(x_t, t, sigma=None, alpha_bar=None):
timestamp = t
t = t.repeat(len(x_t)).round().long().to(x_t.device)
sigma = sigma.repeat(len(x_t), *([1] * (len(sigma.shape) - 1)))
alpha = alpha.repeat(len(x_t), *([1] * (len(alpha.shape) - 1)))
alpha_bar = alpha_bar.repeat(len(x_t), *([1] * (len(alpha_bar.shape) - 1)))
if guide_scale is None or guide_scale == 1.0:
out = model(x=x_t, t=t, **model_kwargs)
@@ -101,15 +98,12 @@ class BaseDiffusion(object):
if self.prediction_type == 'x0':
x0 = out
elif self.prediction_type == 'eps':
x0 = (x_t - sigma * out) / alpha
x0 = (x_t - sigma * out) / alpha_bar
elif self.prediction_type == 'v':
x0 = alpha * x_t - sigma * out
x0 = alpha_bar * x_t - sigma * out
else:
raise NotImplementedError(
f'prediction_type {self.prediction_type} not implemented')
# print("torch.sum(y_out):", torch.sum(y_out), "torch.sum(u_out):", torch.sum(u_out), "torch.sum(out):",
# torch.sum(out), "torch.sum(x0):", torch.sum(x0), "sigmas", sigma, "alphas", alpha)
return x0
sampler_ins = self.get_sampler(sampler)
@@ -117,12 +111,14 @@ class BaseDiffusion(object):
# this is ignored for schnell
sampler_output = sampler_ins.preprare_sampler(
noise,
x = x,
steps=steps,
reverse_scale= reverse_scale,
prediction_type=self.prediction_type,
scheduler_ins=self.sampler_scheduler,
callback_fn=callback_fn)
for _ in trange(steps, disable=not show_progress):
for _ in trange(sampler_output.steps, disable=not show_progress):
trange.desc = sampler_output.msg
sampler_output = sampler_ins.step(sampler_output)
if return_intermediate == 'x_0':
@@ -145,30 +141,19 @@ class BaseDiffusion(object):
if noise is None:
noise = torch.randn_like(x_0)
schedule_output = self.noise_scheduler.add_noise(x_0, noise, **kwargs)
x_t, t, sigma, alpha = schedule_output.x_t, schedule_output.t, schedule_output.sigma, schedule_output.alpha
x_t, t, sigma, alpha_bar = schedule_output.x_t, schedule_output.t, schedule_output.sigma, schedule_output.alpha_bar
out = model(x=x_t, t=t, **model_kwargs)
# mse loss
target = {
'eps': noise,
'x0': x_0,
'v': alpha * noise - sigma * x_0
'v': alpha_bar * noise - sigma * x_0
}[self.prediction_type]
loss = (out - target).pow(2)
if reduction == 'mean':
loss = loss.flatten(1).mean(dim=1)
if self.min_snr_gamma is not None:
alphas = self.noise_scheduler.alphas.to(x_0.device)[t]
sigmas = self.noise_scheduler.sigmas.pow(2).to(x_0.device)[t]
snrs = (alphas / sigmas).clamp(min=1e-20)
min_snrs = snrs.clamp(max=self.min_snr_gamma)
weights = min_snrs / snrs
else:
weights = 1
loss = loss * weights
return loss
def get_sampler(self, sampler):
@@ -248,17 +233,6 @@ class DiffusionFluxRF(BaseDiffusion):
loss = (target - out)**2
if reduction == 'mean':
loss = loss.flatten(1).mean(dim=1)
if self.min_snr_gamma is not None:
alphas = self.noise_scheduler.alphas.to(x_0.device)[t]
sigmas = self.noise_scheduler.sigmas.pow(2).to(x_0.device)[t]
snrs = (alphas / sigmas).clamp(min=1e-20)
min_snrs = snrs.clamp(max=self.min_snr_gamma)
weights = min_snrs / snrs
else:
weights = 1
loss = loss * weights
return loss
@torch.no_grad()
@@ -271,6 +245,8 @@ class DiffusionFluxRF(BaseDiffusion):
show_progress=False,
return_intermediate=None,
intermediate_callback=None,
reverse_scale=-1.,
x=None,
**kwargs):
# sanity check
assert isinstance(steps, (int, torch.LongTensor))
@@ -278,7 +254,7 @@ class DiffusionFluxRF(BaseDiffusion):
assert isinstance(sampler, (str, dict, Config))
intermediates = []
def callback_fn(x_t, t, sigma=None, alpha=None):
def callback_fn(x_t, t, sigma=None, alpha_bar=None):
sigma = torch.full((x_t.shape[0], ),
sigma,
dtype=x_t.dtype,
@@ -291,12 +267,14 @@ class DiffusionFluxRF(BaseDiffusion):
# this is ignored for schnell
sampler_output = sampler_ins.preprare_sampler(
noise,
x=x,
steps=steps,
reverse_scale=reverse_scale,
prediction_type=self.prediction_type,
scheduler_ins=self.sampler_scheduler,
callback_fn=callback_fn)
for _ in trange(steps, disable=not show_progress):
for _ in trange(sampler_output.steps, disable=not show_progress):
trange.desc = sampler_output.msg
sampler_output = sampler_ins.step(sampler_output)
if return_intermediate == 'x_0':
+71 -17
View File
@@ -15,15 +15,18 @@ class SamplerOutput(object):
callback_fn: callable
prediction_type: str
alphas: torch.Tensor
alphas_bar: torch.Tensor
betas: torch.Tensor
sigmas: torch.Tensor
alphas_init: torch.Tensor
alphas_bar_init: torch.Tensor
betas_init: torch.Tensor
sigmas_init: torch.Tensor
ts: torch.Tensor
x_t: torch.Tensor
x_0: torch.Tensor
step: int
steps: int
msg: str
def add_custom_field(self, key: str, value) -> None:
@@ -49,7 +52,7 @@ class BaseDiffusionSampler(object):
self.t_max = self.cfg.get('T_MAX', None)
self.t_min = self.cfg.get('T_MIN', None)
def discretization(self, steps=20, num_timesteps=1000, **kwargs):
def discretization(self, steps=20, num_timesteps=1000, reverse_scale = -1., **kwargs):
# get timesteps
if isinstance(steps, int):
steps += 1 if self.discard_penultimate_step else 0
@@ -74,17 +77,23 @@ class BaseDiffusionSampler(object):
steps = steps.clamp_(t_min, t_max)
elif isinstance(steps, list):
steps = torch.tensor(steps)
timesteps = torch.as_tensor(steps, dtype=torch.float32)
return timesteps
if reverse_scale >=0:
img2img_step = int((1 - reverse_scale) * len(steps))
timesteps = torch.as_tensor(steps[img2img_step:], dtype=torch.float32)
return timesteps
return torch.as_tensor(steps, dtype=torch.float32)
def preprare_sampler(self,
noise,
x=None,
steps=20,
reverse_scale=-1.,
scheduler_ins=None,
prediction_type='',
sigmas=None,
betas=None,
alphas=None,
alphas_bar=None,
callback_fn=None,
**kwargs):
'''
@@ -96,36 +105,52 @@ class BaseDiffusionSampler(object):
4. To ensure the safety of threading, use the instance of SamplerOutput as the manager,
which manage all necessary information.
'''
if reverse_scale >= 0:
assert x is not None
num_timesteps = scheduler_ins.num_timesteps if scheduler_ins is not None else 1000
timestamps = self.discretization(steps,
num_timesteps=num_timesteps,
reverse_scale=reverse_scale,
**kwargs)
alphas = scheduler_ins.t_to_alpha(
timestamps, **kwargs) if scheduler_ins is not None else alphas
alphas_bar = scheduler_ins.t_to_alpha_bar(
timestamps, **kwargs) if scheduler_ins is not None else alphas_bar
betas = scheduler_ins.t_to_beta(
timestamps, **kwargs) if scheduler_ins is not None else betas
sigmas = scheduler_ins.t_to_sigma(
timestamps, **kwargs) if scheduler_ins is not None else sigmas
alphas_init = scheduler_ins.t_to_alpha_init(
timestamps, **kwargs) if scheduler_ins is not None else alphas
alphas_bar_init = scheduler_ins.t_to_alpha_bar_init(
timestamps, **kwargs) if scheduler_ins is not None else alphas_bar
betas_init = scheduler_ins.t_to_beta_init(
timestamps, **kwargs) if scheduler_ins is not None else betas
sigmas_init = scheduler_ins.t_to_sigma_init(
timestamps, **kwargs) if scheduler_ins is not None else sigmas
if reverse_scale >= 0:
x_t = x_0 = scheduler_ins.add_noise(x, noise=noise, t=timestamps[0].repeat(x.size(0)).to(x.device)).x_t if len(timestamps) > 0 else x
else:
x_t = x_0 = noise
# Consider the sigma's list is from sigma_ to zero. the steps equal to len(timestamps)
output = SamplerOutput(callback_fn=callback_fn,
prediction_type=prediction_type,
alphas=alphas,
alphas_bar=alphas_bar,
betas=betas,
sigmas=sigmas,
alphas_init=alphas_init,
alphas_bar_init=alphas_bar_init,
betas_init=betas_init,
sigmas_init=sigmas_init,
ts=timestamps,
x_t=noise,
x_0=noise,
x_t=x_t,
x_0=x_0,
step=0,
msg='step 0')
msg='step 0',
steps=len(timestamps) - 1)
return output
def step(self, sampler_ouput):
@@ -159,22 +184,35 @@ class DDIMSampler(BaseDiffusionSampler):
def preprare_sampler(self,
noise,
x=None,
steps=20,
reverse_scale = -1.,
scheduler_ins=None,
prediction_type='',
sigmas=None,
betas=None,
alphas=None,
alphas_bar=None,
callback_fn=None,
**kwargs):
output = super().preprare_sampler(noise, steps, scheduler_ins,
prediction_type, sigmas, betas,
alphas, callback_fn, **kwargs)
output = super().preprare_sampler(noise,
x = x,
steps = steps,
reverse_scale = reverse_scale,
scheduler_ins = scheduler_ins,
prediction_type = prediction_type,
sigmas = sigmas,
betas = betas,
alphas = alphas,
alphas_bar = alphas_bar,
callback_fn = callback_fn,
**kwargs)
sigmas = output.sigmas
sigmas = torch.cat([sigmas, sigmas.new_zeros([1])])
sigmas_vp = (sigmas**2 / (1 + sigmas**2))**0.5
sigmas_vp[sigmas == float('inf')] = 1.
output.add_custom_field('sigmas_vp', sigmas_vp)
output.steps += 1
return output
def step(self, sampler_output):
@@ -182,10 +220,10 @@ class DDIMSampler(BaseDiffusionSampler):
step = sampler_output.step
t = sampler_output.ts[step]
sigmas_vp = sampler_output.sigmas_vp.to(x_t.device)
alpha_init = _i(sampler_output.alphas_init, step, x_t[:1])
alpha_bar_init = _i(sampler_output.alphas_bar_init, step, x_t[:1])
sigma_init = _i(sampler_output.sigmas_init, step, x_t[:1])
x = sampler_output.callback_fn(x_t, t, sigma_init, alpha_init)
x = sampler_output.callback_fn(x_t, t, sigma_init, alpha_bar_init)
noise_factor = self.eta * (sigmas_vp[step + 1]**2 /
sigmas_vp[step]**2 *
(1 - (1 - sigmas_vp[step]**2) /
@@ -202,16 +240,19 @@ class DDIMSampler(BaseDiffusionSampler):
return sampler_output
@DIFFUSION_SAMPLERS.register_class('flow_eluer')
@DIFFUSION_SAMPLERS.register_class('flow_euler')
class FlowEluerSampler(BaseDiffusionSampler):
def preprare_sampler(self,
noise,
x=None,
steps=20,
reverse_scale = -1.,
scheduler_ins=None,
prediction_type='',
sigmas=None,
betas=None,
alphas=None,
alphas_bar=None,
callback_fn=None,
**kwargs):
if noise.ndim == 3:
@@ -220,9 +261,18 @@ class FlowEluerSampler(BaseDiffusionSampler):
n, _, h, w = noise.shape
seq_len = (h // 2 * w // 2)
kwargs['seq_len'] = seq_len
output = super().preprare_sampler(noise, steps, scheduler_ins,
prediction_type, sigmas, betas,
alphas, callback_fn, **kwargs)
output = super().preprare_sampler(noise,
x = x,
steps = steps,
reverse_scale = reverse_scale,
scheduler_ins = scheduler_ins,
prediction_type = prediction_type,
sigmas = sigmas,
betas = betas,
alphas = alphas,
alphas_bar = alphas_bar,
callback_fn = callback_fn,
**kwargs)
return output
def step(self, sampler_output):
@@ -241,9 +291,13 @@ class FlowEluerSampler(BaseDiffusionSampler):
sampler_output.msg = f'step {step}, sigma_curr: {sigma_curr}, sigma_prev: {sigma_prev}'
return sampler_output
def discretization(self, steps=20, num_timesteps=1000, **kwargs):
def discretization(self, steps=20, num_timesteps=1000, reverse_scale=-1., **kwargs):
# extra step for zero
timesteps = torch.linspace(num_timesteps, 0, steps + 1)
if reverse_scale >= 0:
img2img_step = int((1 - reverse_scale) * len(timesteps))
timesteps = timesteps[img2img_step:]
return timesteps
return timesteps
@staticmethod
+52 -11
View File
@@ -21,7 +21,7 @@ class ScheduleOutput(object):
x_0: torch.Tensor
t: torch.Tensor
sigma: torch.Tensor
alpha: torch.Tensor
alpha_bar: torch.Tensor
custom_fields: dict = field(default_factory=dict)
def add_custom_field(self, key: str, value) -> None:
@@ -30,6 +30,21 @@ class ScheduleOutput(object):
@NOISE_SCHEDULERS.register_class()
class BaseNoiseScheduler(object):
'''
In the diffusion model, the parameters related to the noise schedule are alpha, beta,
and sigma. The following are the definitions of the above three parameters, which should
be the basic property for the instance of noise scheduler.
\alpha_{t} = \sqrt{1 - \beta_{t}^2} \alpha is the strength of signal and \beta is the strength of noise
\sigma_{t} = \sqrt{1 - \overline\alpha} = \sqrt{1 - \prod_{i=1}^{t}\alpha^2_{i}} (P(x_{t}|x_{0}) ~ N(\overline\alpha x_{0}, \sigma^2))
\alpha_bar_{t} = \sqrt{\overline\alpha} = \sqrt{\prod_{i=1}^{t}\alpha^2_{i}} (P(x_{t}|x_{0}) ~ N(\overline\alpha x_{0}, \sigma^2))
where sigma_{t} is the var of p(x_{t-1}|x_{t}, x_{0}).
(reference to https://arxiv.org/abs/2010.02502)
let sigma transfer to beta:
square_\beta = 1 - \frac{1 - square_\sigma_{t}}{1 - square_\sigma_{t - 1 }}
'''
para_dict = {
'NUM_TIMESTEPS': {
'value': 1000,
@@ -48,7 +63,7 @@ class BaseNoiseScheduler(object):
self.num_timesteps = self.cfg.get('NUM_TIMESTEPS', 1000)
self._sample_steps = torch.arange(self.num_timesteps,
dtype=torch.float32)
self._sigmas, self._betas, self._alphas, self._timesteps = None, None, None, None
self._sigmas, self._betas, self._alphas, self._alphas_bar, self._timesteps = None, None, None, None, None
def check_function(self):
try:
@@ -128,6 +143,10 @@ class BaseNoiseScheduler(object):
square_beta = self.sigmas_to_square_betas(sigma)
return torch.sqrt(1 - square_beta)
def t_to_alpha_bar(self, t, **kwargs):
sigma = self.t_to_sigma(t)
return torch.sqrt(1 - sigma**2)
def t_to_beta(self, t, **kwargs):
sigma = self.t_to_sigma(t)
square_beta = self.sigmas_to_square_betas(sigma)
@@ -138,11 +157,11 @@ class BaseNoiseScheduler(object):
t = torch.randint(0,
self.num_timesteps, (x_0.shape[0], ),
device=x_0.device).long()
alpha = _i(self.alphas, t, x_0)
alpha = _i(self.alphas_bar, t, x_0)
sigma = _i(self.sigmas, t, x_0)
x_t = alpha * x_0 + sigma * noise
return ScheduleOutput(x_0=x_0, x_t=x_t, t=t, alpha=alpha, sigma=sigma)
return ScheduleOutput(x_0=x_0, x_t=x_t, t=t, alpha_bar=alpha, sigma=sigma)
def t_to_alpha_init(self, t, **kwargs):
indices = t.long()
@@ -153,6 +172,16 @@ class BaseNoiseScheduler(object):
alpha = self.alphas[step_indices].flatten().to(t)
return alpha
def t_to_alpha_bar_init(self, t, **kwargs):
indices = t.long()
indices[indices >= self.num_timesteps] = self.num_timesteps - 1
timesteps = self.timesteps.to(t)[indices]
step_indices = [(self.timesteps.to(t) == t).nonzero().item()
for t in timesteps]
alpha_bar = self.alphas_bar[step_indices].flatten().to(t)
return alpha_bar
def t_to_beta_init(self, t, **kwargs):
indices = t.long()
indices[indices >= self.num_timesteps] = self.num_timesteps - 1
@@ -205,6 +234,10 @@ class BaseNoiseScheduler(object):
def alphas(self):
return self._alphas
@property
def alphas_bar(self):
return self._alphas_bar
@property
def timesteps(self):
return self._timesteps
@@ -221,6 +254,10 @@ class BaseNoiseScheduler(object):
'data': self._alphas.cpu().numpy(),
'label': 'alphas'
}, {
'data': self._alphas_bar.cpu().numpy(),
'label': 'alphas_bar'
},
{
'data': self._timesteps.cpu().numpy() / self.num_timesteps,
'label': 'timesteps'
}]
@@ -280,7 +317,8 @@ class ScaledLinearScheduler(BaseNoiseScheduler):
self.snr_shift_scale,
self.rescale_betas_zero_snr)
self._betas = torch.sqrt(square_betas)
self._alphas = torch.sqrt(1 - self._sigmas**2)
self._alphas = torch.sqrt(1 - square_betas)
self._alphas_bar = torch.sqrt(1 - self._sigmas**2)
self._timesteps = torch.arange(len(self._sigmas), dtype=torch.float32)
@@ -304,7 +342,8 @@ class LinearScheduler(BaseNoiseScheduler):
sigmas = self.betas_to_sigmas(betas)
self._sigmas = sigmas
self._betas = betas
self._alphas = torch.sqrt(1 - sigmas**2)
self._alphas = torch.sqrt(1 - betas**2)
self._alphas_bar = torch.sqrt(1 - sigmas**2)
self._timesteps = torch.arange(len(sigmas), dtype=torch.float32)
@@ -319,7 +358,8 @@ class FlowMatchUniformScheduler(BaseNoiseScheduler):
self._timesteps = timesteps
self._sigmas = self.t_to_sigma(timesteps)
self._betas = torch.sqrt(self.sigmas_to_square_betas(self._sigmas))
self._alphas = torch.sqrt(1 - self.betas**2)
self._alphas = torch.sqrt(1 - self._betas**2)
self._alphas_bar = torch.sqrt(1 - self._sigmas ** 2)
def add_noise(self, x_0, noise=None, t=None, **kwargs):
if t is None:
@@ -332,7 +372,7 @@ class FlowMatchUniformScheduler(BaseNoiseScheduler):
x_t=x_t,
t=t,
sigma=sigma,
alpha=self.t_to_alpha(t))
alpha_bar=self.t_to_alpha_bar(t))
def sigma_to_t(self, sigma, **kwargs):
return sigma * self.num_timesteps
@@ -406,7 +446,7 @@ class FlowMatchShiftScheduler(FlowMatchUniformScheduler):
x_t=x_t,
t=t,
sigma=sigma,
alpha=self.t_to_alpha(t))
alpha_bar=self.t_to_alpha_bar(t))
def sigma_to_t(self, sigma, **kwargs):
t = sigma / (sigma - self.shift * sigma + self.shift)
@@ -486,7 +526,7 @@ class FlowMatchFluxShiftScheduler(FlowMatchUniformScheduler):
x_t=x_t,
t=t,
sigma=sigma,
alpha=self.t_to_alpha(t))
alpha_bar=self.t_to_alpha_bar(t))
def sigma_to_t(self, sigma, **kwargs):
seq_len = kwargs.get('seq_len', 256)
@@ -570,6 +610,7 @@ class FlowMatchSigmaScheduler(FlowMatchUniformScheduler):
(self.shift - 1) * timesteps)
self._betas = torch.sqrt(self.sigmas_to_square_betas(self._sigmas))
self._alphas = torch.sqrt(1 - self.betas**2)
self._alphas_bar = torch.sqrt(1 - self._sigmas ** 2)
def add_noise(self, x_0, noise=None, t=None, **kwargs):
if t is None:
@@ -589,7 +630,7 @@ class FlowMatchSigmaScheduler(FlowMatchUniformScheduler):
x_t=x_t,
t=t,
sigma=sigma,
alpha=self.t_to_alpha(t))
alpha_bar=self.t_to_alpha_bar(t))
def compute_density_for_timestep_sampling(self, t):
"""Compute the density for sampling the timesteps when doing SD3 training.
+30 -73
View File
@@ -832,22 +832,28 @@ class T5EmbedderHF(BaseEmbedder):
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.t5_dtype = cfg.get('T5_DTYPE', 'bfloat16')
self.use_grad = cfg.get('USE_GRAD', False)
self.clean = cfg.get('CLEAN', 'whitespace')
self.added_identifier = cfg.get('ADDED_IDENTIFIER', None)
tokenizer_path = cfg.get('TOKENIZER_PATH', None)
pretrained_path = cfg.get('PRETRAINED_MODEL', None)
if pretrained_path:
with FS.get_dir_to_local_dir(pretrained_path,
wait_finish=True) as local_path:
if self.t5_dtype is not None:
self.model = T5EncoderModel.from_pretrained(
local_path,
torch_dtype=getattr(
torch,
'float' if self.t5_dtype == 'float32' else self.t5_dtype))
else:
self.model = T5EncoderModel.from_pretrained(local_path)
else:
self.model = None
if tokenizer_path:
self.tokenize_kargs = {'return_tensors': 'pt'}
with FS.get_dir_to_local_dir(tokenizer_path,
@@ -869,9 +875,6 @@ class T5EmbedderHF(BaseEmbedder):
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():
@@ -888,14 +891,10 @@ class T5EmbedderHF(BaseEmbedder):
else:
x = self.model(tokens.input_ids.to(we.device_id))
x = x.last_hidden_state
# if not self.return_pooled:
# return x.detach()
# else:
# return x.detach(), self.pool(x, tokens.input_ids)
if return_mask:
return x.detach() + 0.0, tokens.attention_mask.to(we.device_id)
else:
return x.detach() + 0.0, None
return x.detach() + 0.0
def pool(self, x, tokens):
# take features from the eot embedding (eot_token is the highest number in each sequence)
@@ -921,6 +920,15 @@ class T5EmbedderHF(BaseEmbedder):
return self(tokens, return_mask=return_mask)
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
tokens = self.tokenizer(text, **self.tokenize_kargs)
return self(tokens, return_mask=return_mask, use_mask=use_mask)
def encode_list(self, text, return_mask=False, use_mask=True):
if isinstance(text, str):
text = [text]
if self.clean:
@@ -942,62 +950,11 @@ class T5EmbedderHF(BaseEmbedder):
else:
return torch.cat(cont, dim=0)
def encode_longlist(self, text_list, return_mask=True):
text_max_len = max([len(p) for p in text_list]) * self.length
cont_list, cont_mask_list = [], []
for pp in text_list:
cont, cont_mask = self.encode(pp, return_mask=return_mask)
cont_channel, cont_dim = cont.shape[0] * cont.shape[1], cont.shape[
2]
cont = cont.view(cont_channel, cont_dim)
cont_mask_channel = cont_mask.shape[0] * cont_mask.shape[1]
cont_mask = cont_mask.view(cont_mask_channel)
select_cont = cont[cont_mask == 1]
select_cont_mask, _ = torch.sort(cont_mask, dim=0, descending=True)
if select_cont.shape[0] != text_max_len:
select_cont = F.pad(
select_cont,
(0, 0, 0, text_max_len - select_cont.shape[0]))
if select_cont_mask.shape[0] != text_max_len:
select_cont_mask = F.pad(
select_cont_mask,
(0, text_max_len - select_cont_mask.shape[0]))
cont_list.append(select_cont)
cont_mask_list.append(select_cont_mask)
return torch.stack(cont_list), torch.stack(cont_mask_list)
def encode_longlist_v1(self, text_list, return_mask=True):
cont_list = []
max_len = 0
for pp in text_list:
cont, cont_mask = self.encode(pp, return_mask=True)
txt_lens = cont_mask.flatten(start_dim=1).sum(dim=-1)
pp_cont = torch.cat(
[c[:txt_len] for c, txt_len in zip(cont, txt_lens)], dim=0)
max_len = pp_cont.size(0) if pp_cont.size(0) > max_len else max_len
cont_list.append(pp_cont)
cont = torch.cat([
torch.cat([c, c.new_zeros(max_len - c.size(0), c.size(1))],
dim=0).unsqueeze(0) for c in cont_list
],
dim=0)
if return_mask:
cont_mask = torch.cat([
torch.cat(
[c.new_ones(c.size(0)),
c.new_zeros(max_len - c.size(0))],
dim=-1).unsqueeze(0) for c in cont_list
],
dim=0).type(torch.long, non_blocking=True)
return cont, cont_mask
else:
return cont
def encode_list(self, text_list, return_mask=True):
def encode_list_of_list(self, text_list, return_mask=True, use_mask=True):
cont_list = []
mask_list = []
for pp in text_list:
cont, cont_mask = self.encode(pp, return_mask=return_mask)
cont, cont_mask = self.encode_list(pp, return_mask=return_mask, use_mask=use_mask)
cont_list.append(cont)
mask_list.append(cont_mask)
if return_mask:
@@ -1,3 +1,4 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.network.autoencoder.ae_kl import AutoencoderKL
from scepter.modules.model.network.autoencoder.ae_kl_cogvideox import AutoencoderKLCogVideoX
File diff suppressed because it is too large Load Diff
@@ -1,7 +1,8 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.model.network.ldm.ldm import LatentDiffusion
from scepter.modules.model.network.ldm.ldm_ace import LatentDiffusionACE
from scepter.modules.model.network.ldm.ldm_ace import (LatentDiffusionACE,
LatentDiffusionACERefiner)
from scepter.modules.model.network.ldm.ldm_edit import LatentDiffusionEdit
from scepter.modules.model.network.ldm.ldm_pixart import LatentDiffusionPixart
from scepter.modules.model.network.ldm.ldm_sce import (
@@ -9,3 +10,6 @@ from scepter.modules.model.network.ldm.ldm_sce import (
LatentDiffusionXLSCEControl, LatentDiffusionXLSCETuning)
from scepter.modules.model.network.ldm.ldm_sd3 import LatentDiffusionSD3
from scepter.modules.model.network.ldm.ldm_xl import LatentDiffusionXL
from scepter.modules.model.network.ldm.ldm_cogvideox import LatentDiffusionCogVideoX
from scepter.modules.model.network.ldm.ldm_flux import (LatentDiffusionFlux,
LatentDiffusionFluxMR)
+258 -5
View File
@@ -1,6 +1,7 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
import math
import random
from contextlib import nullcontext
@@ -10,6 +11,7 @@ from torch import nn
from scepter.modules.model.network.ldm import LatentDiffusion
from scepter.modules.model.registry import MODELS
import torchvision.transforms as T
from scepter.modules.model.utils.basic_utils import check_list_of_list
from scepter.modules.model.utils.basic_utils import \
pack_imagelist_into_tensor_v2 as pack_imagelist_into_tensor
@@ -67,10 +69,10 @@ class LatentDiffusionACE(LatentDiffusion):
if self.use_text_pos_embeddings and not torch.sum(
self.text_position_embeddings.pos) > 0:
identifier_cont, identifier_cont_mask = getattr(
self.cond_stage_model, 'encode')(self.text_indentifers,
self.cond_stage_model, 'encode_list_of_list')(self.text_indentifers,
return_mask=True)
self.text_position_embeddings.load_state_dict(
{'pos': identifier_cont[:, 0, :]})
{'pos': torch.cat( [one_id[0][0, :].unsqueeze(0) for one_id in identifier_cont], dim=0)})
cont_, cont_mask_ = [], []
for pp, edit, c, cm in zip(prompt, edit_image, cont, cont_mask):
if isinstance(pp, list):
@@ -138,7 +140,7 @@ class LatentDiffusionACE(LatentDiffusion):
prompt_ = [[pp] if isinstance(pp, str) else pp for pp in prompt]
try:
cont, cont_mask = getattr(self.cond_stage_model,
'encode_list')(prompt_, return_mask=True)
'encode_list_of_list')(prompt_, return_mask=True)
except Exception as e:
print(e, prompt_)
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont,
@@ -240,11 +242,11 @@ class LatentDiffusionACE(LatentDiffusion):
# with torch.autocast(device_type="cuda", enabled=True, dtype=torch.bfloat16):
prompt_ = [[pp] if isinstance(pp, str) else pp for pp in prompt]
cont, cont_mask = getattr(self.cond_stage_model,
'encode_list')(prompt_, return_mask=True)
'encode_list_of_list')(prompt_, return_mask=True)
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont,
cont_mask)
null_cont, null_cont_mask = getattr(self.cond_stage_model,
'encode_list')(n_prompt,
'encode_list_of_list')(n_prompt,
return_mask=True)
null_cont, null_cont_mask = self.cond_stage_embeddings(
prompt, edit_image, null_cont, null_cont_mask)
@@ -349,3 +351,254 @@ class LatentDiffusionACE(LatentDiffusion):
__class__.__name__,
LatentDiffusionACE.para_dict,
set_name=True)
@MODELS.register_class()
class LatentDiffusionACERefiner(LatentDiffusionACE):
def init_params(self):
super().init_params()
self.enhence_model_cfg = self.cfg.get("ENHENCE_MODEL", None)
self.enhence_sampler_cfg = self.cfg.get("ENHENCE_SAMPLER_CFG", {})
def construct_network(self):
super().construct_network()
if self.enhence_model_cfg:
self.enhence_model = MODELS.build(self.enhence_model_cfg, logger=self.logger).eval().requires_grad_(False)
self.enhence_sampler_cfg = {key.lower(): value for key, value in self.enhence_sampler_cfg.items()}
else:
self.enhence_model = None
self.enhence_sampler_cfg = None
def forward_sample(self,
edit_image=[],
edit_mask=[],
noise=None,
cond_mask=[],
x_shapes=[],
prompt=[],
n_prompt=[],
sampler='ddim',
sample_steps=20,
seed=2023,
guide_scale=4.5,
guide_rescale=0.5,
discretization='trailing',
**kwargs
):
'''
Args:
edit_image: list of list of edit_image
edit_image_mask: list of list of edit_image_mask
image: target image
image_mask: target image mask
prompt: list of list of text
n_prompt: list of list of text
sampler:
sample_steps:
seed:
guide_scale:
guide_rescale:
discretization:
log_num:
**kwargs:
Returns:
'''
# prepare data
context, null_context = {}, {}
context['x_shapes'] = null_context['x_shapes'] = x_shapes
# process image mask
context['x_mask'] = null_context['x_mask'] = cond_mask
# process text
# with torch.autocast(device_type="cuda", enabled=True, dtype=torch.bfloat16):
cont, cont_mask = getattr(self.cond_stage_model, 'encode_list')(prompt, return_mask=True)
cont, cont_mask = self.cond_stage_embeddings(prompt, edit_image, cont, cont_mask)
null_cont, null_cont_mask = getattr(self.cond_stage_model, 'encode_list')(n_prompt, return_mask=True)
null_cont, null_cont_mask = self.cond_stage_embeddings(prompt, edit_image, null_cont, null_cont_mask)
context['crossattn'] = cont
null_context['crossattn'] = null_cont
null_context['edit'] = context['edit'] = edit_image
null_context['edit_mask'] = context['edit_mask'] = edit_mask
# process sample
model = self.model_ema if self.use_ema and self.eval_ema else self.model
embedding_context = model.no_sync if isinstance(model, torch.distributed.fsdp.FullyShardedDataParallel) \
else nullcontext
with embedding_context():
samples = self.diffusion.sample(solver=sampler,
noise=noise,
model=model,
model_kwargs=[{
'cond': context,
'mask': cont_mask,
'text_position_embeddings': self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
}, {
'cond': null_context,
'mask': null_cont_mask,
'text_position_embeddings': self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
}] if guide_scale is not None and guide_scale > 1 else {
'cond': context,
'mask': cont_mask,
'text_position_embeddings': self.text_position_embeddings.pos if hasattr(
self.text_position_embeddings, 'pos') else None
},
cat_uc=False,
steps=sample_steps,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
discretization=discretization,
show_progress=True,
seed=seed,
condition_fn=None,
clamp=None,
percentile=None,
t_max=None,
t_min=None,
discard_penultimate_step=None,
return_intermediate=None,
**kwargs)
samples = unpack_tensor_into_imagelist(samples, x_shapes)
x_samples = self.decode_first_stage(samples)
return x_samples
def upscale_resize(self, image, interpolation=T.InterpolationMode.BILINEAR):
_, c, H, W = image.shape
scale = max(1.0, math.sqrt(4096 / ((H / 16) * (W / 16))))
rH = int(H * scale) // 16 * 16 # ensure divisible by self.d
rW = int(W * scale) // 16 * 16
image = T.Resize((rH, rW), interpolation=interpolation, antialias=True)(image)
return image
@torch.no_grad()
def forward_test(self,
edit_image=[],
edit_image_mask=[],
image=None,
image_mask=None,
prompt=[],
n_prompt=[],
sampler='ddim',
sample_steps=20,
seed=2023,
guide_scale=4.5,
guide_rescale=0.5,
discretization='trailing',
enhance_scale=0.99,
log_num=-1,
**kwargs):
assert check_list_of_list(prompt) and check_list_of_list(edit_image) and check_list_of_list(edit_image_mask)
assert len(edit_image) == len(edit_image_mask) == len(prompt)
assert self.cond_stage_model is not None
# gc_seg is unused
kwargs.pop("gc_seg", -1)
prompt, n_prompt, image, image_mask, edit_image, edit_image_mask = self.limit_batch_data(
[prompt, n_prompt, image, image_mask, edit_image, edit_image_mask], log_num)
prompt = [[pp] if isinstance(pp, str) else pp for pp in prompt]
g = torch.Generator(device=we.device_id)
seed = seed if seed >= 0 else random.randint(0, 2 ** 32 - 1)
g.manual_seed(seed)
n_prompt = copy.deepcopy(prompt)
# only modify the last prompt to be zero
for nn_p_id, nn_p in enumerate(n_prompt):
if isinstance(nn_p, str):
n_prompt[nn_p_id] = [""]
elif isinstance(nn_p, list):
n_prompt[nn_p_id][-1] = ""
else:
raise NotImplementedError
# process image
image = to_device(image)
x = self.encode_first_stage(image, **kwargs)
noise = [torch.empty(*i.shape, device=we.device_id).normal_(generator=g) for i in x]
noise, x_shapes = pack_imagelist_into_tensor(noise)
image_mask = to_device(image_mask, strict=False)
cond_mask = [self.interpolate_func(i) for i in image_mask] if image_mask is not None else [None] * len(image)
# processe edit image & edit image mask
edit_image = [to_device(i, strict=False) for i in edit_image]
edit_image_mask = [to_device(i, strict=False) for i in edit_image_mask]
e_img, e_mask = [], []
for u, m in zip(edit_image, edit_image_mask):
if u is None:
continue
if m is None:
m = [None] * len(u)
e_img.append(self.encode_first_stage(u, **kwargs))
e_mask.append([self.interpolate_func(i) for i in m])
x_samples = self.forward_sample(
edit_image=e_img,
edit_mask=e_mask,
noise=noise,
cond_mask=cond_mask,
x_shapes=x_shapes,
prompt=prompt,
n_prompt=n_prompt,
sampler=sampler,
sample_steps=sample_steps,
seed=seed,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
discretization='trailing',
**kwargs)
if self.enhence_model and enhance_scale > 0:
x_samples = [self.upscale_resize(x) for x in x_samples]
x_start = self.enhence_model.encode_first_stage(x_samples, **kwargs)
noise = []
for i, x in enumerate(x_start):
noise_ = self.enhence_model.noise_sample(1, x_samples[i].shape[2], x_samples[i].shape[3], seed)
noise.append(noise_)
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
x_samples = self.enhence_model.forward_sample(noise = noise,
x = x_start,
reverse_scale = enhance_scale,
prompt =[kwargs.pop("enhance_prompt", "") for _ in noise],
**self.enhence_sampler_cfg)
outputs = list()
for i in range(len(prompt)):
rec_img = torch.clamp((x_samples[i].float() + 1.0) / 2.0 + self.decoder_bias / 255, min=0.0, max=1.0)
rec_img = rec_img.squeeze(0)
edit_imgs, edit_img_masks = [], []
if edit_image is not None and edit_image[i] is not None:
if edit_image_mask[i] is None:
edit_image_mask[i] = [None] * len(edit_image[i])
for edit_img, edit_mask in zip(edit_image[i], edit_image_mask[i]):
edit_img = torch.clamp((edit_img + 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__,
LatentDiffusionACERefiner.para_dict,
set_name=True)
@@ -0,0 +1,225 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import random
import torch
from typing import Tuple
from scepter.modules.model.network.ldm import LatentDiffusion
from scepter.modules.model.registry import MODELS
from scepter.modules.model.utils.basic_utils import default
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.model.backbone.cogvideox.utils import get_3d_rotary_pos_embed, get_resize_crop_region_for_grid
@MODELS.register_class()
class LatentDiffusionCogVideoX(LatentDiffusion):
para_dict = LatentDiffusion.para_dict
def init_params(self):
super().init_params()
self.latent_channels = self.model_config.get('LATENT_CHANNELS', self.model_config.IN_CHANNELS)
self.scale_factor_spatial = self.cfg.get('SCALE_FACTOR_SPATIAL', 8)
self.scale_factor_temporal = self.cfg.get('SCALE_FACTOR_TEMPORAL', 4)
self.scaling_factor_image = self.cfg.get('SCALING_FACTOR_IMAGE', 0.7)
self.use_rotary_positional_embeddings = self.model_config.get('USE_ROTARY_POSITIONAL_EMBEDDINGS', False)
self.attention_head_dim = self.model_config.get('ATTENTION_HEAD_DIM', 64)
self.patch_size = self.model_config.get('PATCH_SIZE', 2)
self.sample_height = self.first_stage_config.get('SAMPLE_HEIGHT', 480)
self.sample_width = self.first_stage_config.get('SAMPLE_WIDTH', 720)
self.noised_image_dropout = self.cfg.get('NOISED_IMAGE_DROPOUT', 0.05)
def construct_network(self):
super().construct_network()
self.model = self.model.to(getattr(torch, self.model_config.DTYPE))
@torch.no_grad()
def encode_first_stage(self, x, **kwargs):
if isinstance(x, list):
x = torch.stack(x, dim=0) # [B, C, F, H, W]
latents = self.scaling_factor_image * self.first_stage_model.encode(x).sample()
return latents
@torch.no_grad()
def decode_first_stage(self, latents):
latents = latents.permute(0, 2, 1, 3, 4) # [batch_size, num_channels, num_frames, height, width]
latents = 1 / self.scaling_factor_image * latents
frames = self.first_stage_model.decode(latents)
return frames
def get_image_latent(self, image, video, noise):
latent = torch.zeros_like(noise)
if isinstance(image, list):
image = torch.stack(image, dim=0) # [B, C, F, H, W]
if len(image.shape) == 4: # [B, C, H, W]
image = image.unsqueeze(2) # [B, C, F, H, W]
image_latent = self.encode_first_stage(image) # [B, C, F, H, W]
image_latent = image_latent.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
latent[:, :1, :, :, :] = image_latent
return latent, image
def noise_sample(self, batch_size, num_frames, height, width, generator, dtype=torch.bfloat16):
shape = (batch_size,
(num_frames - 1) // self.scale_factor_temporal + 1,
self.latent_channels,
height // self.scale_factor_spatial,
width // self.scale_factor_spatial
)
noise = torch.randn(shape, generator=generator, dtype=dtype, device='cpu').to(we.device_id)
return noise
def _prepare_rotary_positional_embeddings(
self,
height: int,
width: int,
num_frames: int,
device: torch.device,
) -> Tuple[torch.Tensor, torch.Tensor]:
grid_height = height // (self.scale_factor_spatial * self.patch_size)
grid_width = width // (self.scale_factor_spatial * self.patch_size)
base_size_width = self.sample_width // (self.scale_factor_spatial * self.patch_size)
base_size_height = self.sample_height // (self.scale_factor_spatial * self.patch_size)
grid_crops_coords = get_resize_crop_region_for_grid(
(grid_height, grid_width), base_size_width, base_size_height
)
freqs_cos, freqs_sin = get_3d_rotary_pos_embed(
embed_dim=self.attention_head_dim,
crops_coords=grid_crops_coords,
grid_size=(grid_height, grid_width),
temporal_size=num_frames,
)
freqs_cos = freqs_cos.to(device=device)
freqs_sin = freqs_sin.to(device=device)
return freqs_cos, freqs_sin
def forward_train(self, video=None, video_latent=None, image=None, noise=None, prompt=None, image_size=None, **kwargs):
# video: [B, C, F, H, W]
if image_size is None: image_size = [480, 720]
if video_latent is not None:
x_start = torch.stack(video_latent)
else:
x_start = self.encode_first_stage(video, **kwargs)
x_start = x_start.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
t = torch.randint(low=0, high=self.num_timesteps, size=(len(video),), device=we.device_id)
if prompt and self.cond_stage_model:
with torch.autocast(device_type='cuda', enabled=True, dtype=torch.bfloat16):
cont = getattr(self.cond_stage_model, 'encode')(prompt, return_mask=False, use_mask=False)
if noise is None:
noise = torch.randn_like(x_start)
if image is not None:
if random.random() < self.noised_image_dropout:
image_latent = torch.zeros_like(noise)
else:
image_latent, _ = self.get_image_latent(image, video, noise)
else:
image_latent = None
height, width = image_size
image_rotary_emb = (
self._prepare_rotary_positional_embeddings(height, width, noise.size(1), we.device_id)
if self.use_rotary_positional_embeddings
else None
)
loss = self.diffusion.loss(x_0=x_start,
t=t,
model=self.model,
model_kwargs={"cond": cont,
'image_latent': image_latent,
'image_rotary_emb': image_rotary_emb},
noise=noise,
**kwargs)
loss = loss.mean()
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
return ret
@torch.no_grad()
@torch.autocast('cuda', dtype=torch.bfloat16)
def forward_test(self,
video=None,
image=None,
prompt=None,
n_prompt=None,
sampler='ddim',
sample_steps=50,
seed=42,
guide_scale=6.0,
guide_rescale=0.0,
num_frames=49,
image_size=None,
show_process=False,
**kwargs):
if image_size is None:
image_size = [480, 720]
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
generator = torch.Generator().manual_seed(seed)
prompt = [prompt] if isinstance(prompt, str) else prompt
num_samples = len(prompt)
n_prompt = default(n_prompt, [self.default_n_prompt] * len(prompt))
if prompt and self.cond_stage_model:
with torch.autocast(device_type='cuda', enabled=True, dtype=torch.bfloat16):
cont = getattr(self.cond_stage_model, 'encode')(prompt, return_mask=False, use_mask=False)
null_cont = getattr(self.cond_stage_model, 'encode')(n_prompt, return_mask=False, use_mask=False)
height, width = image_size
noise = self.noise_sample(num_samples, num_frames, height, width, generator)
image_rotary_emb = (
self._prepare_rotary_positional_embeddings(height, width, noise.size(1), we.device_id)
if self.use_rotary_positional_embeddings
else None
)
image_latent, image = self.get_image_latent(image, video, noise) if image is not None else (None, None)
samples = self.diffusion.sample(noise=noise,
sampler=sampler,
model=self.model,
model_kwargs=[{
'cond': cont,
'image_latent': image_latent,
'image_rotary_emb': image_rotary_emb,
}, {
'cond': null_cont,
'image_latent': image_latent,
'image_rotary_emb': image_rotary_emb,
}],
steps=sample_steps,
show_progress=True,
use_dynamic_cfg=True,
guide_scale=guide_scale,
guide_rescale=guide_rescale,
return_intermediate=None,
**kwargs).float()
x_frames = self.decode_first_stage(samples).float()
outputs = []
for batch_idx in range(num_samples):
rec_video = torch.clamp(x_frames[batch_idx] / 2 + 0.5, min=0.0, max=1.0)
one_tup = {
'reconstruct_video': rec_video.squeeze(0).float(),
'instruction': prompt[batch_idx]
}
if image is not None:
ori_image = torch.clamp(image[batch_idx] / 2 + 0.5, min=0.0, max=1.0)
one_tup['edit_image'] = ori_image
if video is not None:
ori_video = torch.clamp(video[batch_idx] / 2 + 0.5, min=0.0, max=1.0)
one_tup['target_video'] = ori_video.squeeze(0)
outputs.append(one_tup)
return outputs
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
LatentDiffusionCogVideoX.para_dict,
set_name=True)
+167 -2
View File
@@ -4,15 +4,19 @@ import copy
import math
import numbers
import random
from contextlib import nullcontext
import torch
from scepter.modules.model.network.ldm import LatentDiffusion
from scepter.modules.model.registry import MODELS, BACKBONES, LOSSES, TOKENIZERS, EMBEDDERS, DIFFUSIONS
from scepter.modules.model.utils.basic_utils import disabled_train
from scepter.modules.model.utils.basic_utils import disabled_train, check_list_of_list, to_device, \
pack_imagelist_into_tensor, unpack_tensor_into_imagelist, limit_batch_data
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.model.utils.basic_utils import count_params
@MODELS.register_class()
class LatentDiffusionFlux(LatentDiffusion):
para_dict = LatentDiffusion.para_dict
@@ -137,7 +141,7 @@ class LatentDiffusionFlux(LatentDiffusion):
def forward_test(self,
image=None,
prompt=None,
sampler='flow_eluer',
sampler='flow_euler',
sample_steps=20,
seed=2023,
guide_scale=4.5,
@@ -218,3 +222,164 @@ class LatentDiffusionFlux(LatentDiffusion):
@torch.no_grad()
def decode_first_stage(self, z):
return self.first_stage_model.decode(z)
@MODELS.register_class()
class LatentDiffusionFluxMR(LatentDiffusionFlux):
para_dict = {
}
para_dict.update(LatentDiffusion.para_dict)
def forward_train(self,
image=None,
noise=None,
prompt=[],
**kwargs):
if check_list_of_list(prompt):
prompt = [pp[0] for pp in prompt]
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
context = getattr(self.cond_stage_model, 'encode')(prompt)
image = to_device(image)
x_start = self.encode_first_stage(image, **kwargs)
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
guide_scale = self.guide_scale
if guide_scale is not None:
guide_scale = torch.full((x_start.shape[0],), guide_scale, device=x_start.device, dtype=x_start.dtype)
else:
guide_scale = None
loss = self.diffusion.loss(x_0=x_start,
model=self.model,
model_kwargs={"cond": context,
"gc_seg": gc_seg,
"guidance": guide_scale},
noise=None,
reduction='none',
**kwargs)
loss = loss[loss_mask].mean()
ret = {'loss': loss, 'probe_data': {'prompt': prompt}}
return ret
@torch.no_grad()
def forward_sample(self,
noise = None,
prompt=None,
sampler='flow_euler',
sample_steps=20,
guide_scale=3.5,
show_process=True,
x = None,
reverse_scale = 0.,
**kwargs
):
noise, x_shapes = pack_imagelist_into_tensor(noise)
if x is not None:
x, _ = pack_imagelist_into_tensor(x)
context = getattr(self.cond_stage_model, 'encode')(prompt)
context["x_shapes"] = x_shapes
guide_scale = guide_scale or self.guide_scale
if guide_scale is not None:
guide_scale = torch.full((noise.shape[0],), guide_scale, device=noise.device, dtype=noise.dtype)
else:
guide_scale = None
# UNet use input n_prompt
model = self.model_ema if self.use_ema and self.eval_ema else self.model
embedding_context = model.no_sync if isinstance(model, torch.distributed.fsdp.FullyShardedDataParallel) \
else nullcontext
with embedding_context():
x_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,
reverse_scale = reverse_scale,
x = x,
**kwargs).float()
x_samples = unpack_tensor_into_imagelist(x_samples, x_shapes)
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
x_samples = self.decode_first_stage(x_samples)
return x_samples
@torch.no_grad()
def forward_test(self,
image=None,
prompt=[],
sampler='flow_euler',
sample_steps=20,
seed=2023,
guide_scale=3.5,
guide_rescale=0.0,
show_process=True,
log_num = -1,
**kwargs):
if check_list_of_list(prompt):
prompt = [pp[0] for pp in prompt]
assert self.cond_stage_model is not None
# gc_seg is unused
prompt, image = limit_batch_data([prompt, image], log_num)
seed = seed if seed >= 0 else random.randint(0, 2**32 - 1)
if 'index' in kwargs:
kwargs.pop('index')
if image is not None:
noise = [self.noise_sample(1, ix.shape[1], ix.shape[2], seed) for ix in image]
else:
image_size = None
if 'meta' in kwargs:
meta = kwargs.pop('meta')
if 'image_size' in meta:
h = int(meta['image_size'][0][0])
w = int(meta['image_size'][1][0])
image_size = [h, w]
if 'image_size' in kwargs:
image_size = kwargs.pop('image_size')
if isinstance(image_size, numbers.Number):
image_size = [image_size, image_size]
if image_size is None:
image_size = [1024, 1024]
height, width = image_size
noise = [self.noise_sample(1, height, width, seed) for _ in prompt]
x_samples = self.forward_sample(
prompt=prompt,
sampler=sampler,
sample_steps=sample_steps,
guide_scale=guide_scale,
show_process=show_process,
noise=noise,
)
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)
one_tup = {'prompt': prompt[i], 'n_prompt': '', 'image': rec_img}
outputs.append(one_tup)
return outputs
@staticmethod
def get_config_template():
return dict_to_yaml('MODEL',
__class__.__name__,
LatentDiffusionFlux.para_dict,
set_name=True)
@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]
@@ -102,3 +102,25 @@ def to_device(inputs, strict=True):
def check_list_of_list(ll):
return isinstance(ll, list) and all(isinstance(i, list) for i in ll)
def pack_imagelist_into_tensor(image_list):
image_tensor, shapes = [], []
for img in image_list:
_, c, h, w = img.size()
image_tensor.append(img.view(c, h * w).transpose(1, 0)) # h*w, c
shapes.append((h, w))
image_tensor = pad_sequence(image_tensor, batch_first=True).permute(0, 2, 1) # b, c, l
return image_tensor, shapes
def limit_batch_data(batch_data_list, log_num):
if log_num and log_num > 0:
batch_data_list_limited = []
for sub_data in batch_data_list:
if sub_data is not None:
sub_data = sub_data[:log_num]
batch_data_list_limited.append(sub_data)
return batch_data_list_limited
else:
return batch_data_list
+2 -1
View File
@@ -1,7 +1,8 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
from scepter.modules.solver import hooks
from scepter.modules.solver.ace_solver import ACESolver
from scepter.modules.solver.base_solver import BaseSolver
from scepter.modules.solver.diffusion_solver import LatentDiffusionSolver
from scepter.modules.solver.train_val_solver import TrainValSolver
from scepter.modules.solver.ace_solver import ACESolver
from scepter.modules.solver.diffusion_video_solver import LatentDiffusionVideoSolver
+108 -88
View File
@@ -10,24 +10,23 @@ from functools import partial
import numpy as np
import torch
import torch.nn as nn
from torch.distributed.fsdp import FullStateDictConfig
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import (MixedPrecision, ShardingStrategy,
StateDictType)
from torch.distributed.fsdp.wrap import lambda_auto_wrap_policy
from torch.nn.parallel import DistributedDataParallel
from tqdm import tqdm
from scepter.modules.data.dataset import DATASETS
from scepter.modules.opt.lr_schedulers import LR_SCHEDULERS
from scepter.modules.opt.optimizers import OPTIMIZERS
from scepter.modules.solver import BaseSolver
from scepter.modules.solver.registry import SOLVERS
from scepter.modules.utils.config import Config, dict_to_yaml
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 .base_solver import BaseSolver
from torch.distributed.fsdp import FullStateDictConfig
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import (MixedPrecision, ShardingStrategy,
StateDictType)
from torch.distributed.fsdp.wrap import (lambda_auto_wrap_policy,
size_based_auto_wrap_policy)
from torch.nn.parallel import DistributedDataParallel
from tqdm import tqdm
sharding_strategy_map = {
'full_shard': ShardingStrategy.FULL_SHARD,
@@ -38,17 +37,18 @@ sharding_strategy_map = {
def shard_model(model,
device_id,
process_group=None,
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
buffer_dtype=torch.float32,
fsdp_group=['blocks'],
sharding_strategy=ShardingStrategy.FULL_SHARD,
sync_module_states=False):
sync_module_states=False,
use_orig_params=False):
wrap_modules = []
for module_name in fsdp_group:
if hasattr(model, module_name):
if isinstance(getattr(model, module_name),
(list, tuple, nn.ModuleList)):
if isinstance(getattr(model, module_name), (list, tuple, nn.ModuleList)):
wrap_modules.extend([m for m in getattr(model, module_name)])
else:
wrap_modules.extend([getattr(model, module_name)])
@@ -56,7 +56,7 @@ def shard_model(model,
warnings.warn("Can't find module {} in model".format(module_name))
return FSDP(
module=model,
process_group=None,
process_group=process_group,
sharding_strategy=sharding_strategy,
auto_wrap_policy=partial(
# size_based_auto_wrap_policy, min_num_params=int(1e6),
@@ -66,7 +66,8 @@ def shard_model(model,
reduce_dtype=reduce_dtype,
buffer_dtype=buffer_dtype),
device_id=device_id,
sync_module_states=sync_module_states)
sync_module_states=sync_module_states,
use_orig_params=use_orig_params)
def get_module(instance, sub_module):
@@ -195,7 +196,11 @@ class LatentDiffusionSolver(BaseSolver):
self.logger.info('Use fsdp as the backend of ddp.')
else:
self.logger.info('Use default backend.')
self.use_scaler = cfg.get('USE_SCALER', True)
self.enable_gradscaler = cfg.get('ENABLE_GRADSCALER', False)
self.use_orig_params = cfg.get('USE_ORIG_PARAMS', False)
self.model_shard = cfg.get('SHARDING_STRATEGY', 'full_shard')
self.sharding_size = cfg.get('SHARDING_SIZE', None)
self.reduce_dtype = getattr(torch,
cfg.get('FSDP_REDUCE_DTYPE', 'float32'))
self.buffer_dtype = getattr(torch,
@@ -211,7 +216,7 @@ class LatentDiffusionSolver(BaseSolver):
self.sample_args = cfg.get('SAMPLE_ARGS', None)
self.tuner_cfg = cfg.get('TUNER', None)
self.freeze_cfg = cfg.get('FREEZE', None)
self.log_train_num = cfg.get('LOG_TRAIN_NUM', -1)
self.log_train_num = cfg.get("LOG_TRAIN_NUM", -1)
def set_up(self):
self.construct_data()
@@ -220,6 +225,7 @@ class LatentDiffusionSolver(BaseSolver):
self.model_to_device()
self.init_lr()
self.init_opti()
self.logger.info(self.model)
def construct_hook(self):
# initialize data
@@ -281,54 +287,79 @@ class LatentDiffusionSolver(BaseSolver):
def init_opti(self):
import torch.cuda.amp as amp
import torch.distributed as dist
if we.is_distributed:
if self.use_fairscale:
from fairscale.nn.data_parallel import ShardedDataParallel
from fairscale.optim.oss import OSS
if hasattr(self.model, 'ignored_parameters'):
train_params, _ = self.model.parameters(
train_params, ignored_params = self.model.parameters(
), self.model.ignored_parameters()
else:
train_params, _ = self.model.parameters(), None
train_params, ignored_params = self.model.parameters(
), None
self.optimizer = OSS(params=train_params,
optim=torch.optim.AdamW,
lr=self.cfg.OPTIMIZER.LEARNING_RATE)
self.model = ShardedDataParallel(self.model, self.optimizer)
elif self.use_fsdp:
shard_fn = partial
if self.model_shard == 'hybrid_shard' and self.sharding_size is not None and self.sharding_size > 1:
if self.sharding_size > we.world_size:
self.logger.info(f'Reset sharding_size ({self.sharding_size}) to world_size ({we.world_size})')
sharding_size = min(self.sharding_size, we.world_size)
assert we.world_size % sharding_size == 0
# mesh to facilitate rank indexing
mesh = torch.arange(we.world_size).view(-1, sharding_size)
# sharding groups
for ranks in mesh.tolist():
group = dist.new_group(ranks=ranks)
if we.rank in ranks:
sharding_group = group
# replication groups
for ranks in mesh.t().tolist():
group = dist.new_group(ranks=ranks)
if we.rank in ranks:
replication_group = group
# fsdp group tuple
fsdp_group = (sharding_group, replication_group)
fsdp_rank0 = we.rank // sharding_size * sharding_size
else:
fsdp_group = None
fsdp_rank0 = 0
if self.shard_modules is not None:
for module in self.shard_modules:
if isinstance(module, str):
sub_module = get_module(self.model, module)
if sub_module is not None:
sub_module = shard_model(
sub_module,
device_id=we.device_id,
param_dtype=self.dtype,
reduce_dtype=self.reduce_dtype,
buffer_dtype=self.buffer_dtype,
sharding_strategy=sharding_strategy_map[
self.model_shard],
sync_module_states=True)
sub_module,
process_group=fsdp_group,
device_id=we.device_id,
param_dtype=self.dtype,
reduce_dtype=self.reduce_dtype,
buffer_dtype=self.buffer_dtype,
sharding_strategy=sharding_strategy_map[self.model_shard],
sync_module_states=True,
use_orig_params=self.use_orig_params)
set_module(self.model, module, sub_module)
elif isinstance(module, (dict, Config)):
sub_module = get_module(self.model,
module['MODULE'])
sub_module = get_module(self.model, module["MODULE"])
if sub_module is not None:
sub_module = shard_model(
sub_module,
process_group=fsdp_group,
device_id=we.device_id,
param_dtype=self.dtype,
reduce_dtype=self.reduce_dtype,
buffer_dtype=self.buffer_dtype,
fsdp_group=module.get(
'FSDP_GROUP', ['blocks']),
sharding_strategy=sharding_strategy_map[
self.model_shard],
sync_module_states=True)
set_module(self.model, module['MODULE'],
sub_module)
fsdp_group=module.get("FSDP_GROUP", ["blocks"]),
sharding_strategy=sharding_strategy_map[self.model_shard],
sync_module_states=module.get("SYNC_MODULE_STATES", True),
use_orig_params=self.use_orig_params)
set_module(self.model, module["MODULE"], sub_module)
else:
self.logger.warning(
'FSDP_SHARD_MODULES is None, which means wraping the whold model as the '
@@ -380,22 +411,23 @@ class LatentDiffusionSolver(BaseSolver):
logger=self.logger,
optimizer=self.optimizer)
if self.cfg.DTYPE in ['float16']:
if self.use_scaler and self.cfg.DTYPE in ['float16', 'bfloat16']:
if we.is_distributed:
if self.use_fairscale:
from fairscale.optim.grad_scaler import ShardedGradScaler
self.scaler = ShardedGradScaler(enabled=True)
self.scaler = ShardedGradScaler(enabled=self.enable_gradscaler)
elif self.use_fsdp:
from torch.distributed.fsdp.sharded_grad_scaler import ShardedGradScaler
self.scaler = ShardedGradScaler(enabled=True,
self.scaler = ShardedGradScaler(enabled=self.enable_gradscaler,
process_group=None)
else:
self.scaler = amp.GradScaler()
else:
self.scaler = amp.GradScaler(enabled=self.enable_gradscaler)
elif self.cfg.DTYPE in ['float16']:
self.scaler = amp.GradScaler()
else:
self.scaler = None
else:
self.scaler = None
self.logger.info(self.model)
def load_checkpoint(self, checkpoint: dict):
"""
@@ -450,10 +482,10 @@ class LatentDiffusionSolver(BaseSolver):
f'Load checkpoint for optimizer {module}.')
else:
self.optimizer.load_state_dict(checkpoint['optimizer'])
self.logger.info('Load checkpoint for optimizer.')
self.logger.info(f'Load checkpoint for optimizer.')
if 'scaler' in checkpoint and self.scaler:
self.scaler.load_state_dict(checkpoint['scaler'])
self.logger.info('Load checkpoint for scaler.')
self.logger.info(f'Load checkpoint for scaler.')
self.logger.info('Load checkpoint finished.')
def save_checkpoint(self) -> dict:
@@ -522,7 +554,8 @@ class LatentDiffusionSolver(BaseSolver):
ckpt['model'][module] = current_module.state_dict()
else:
ckpt['model'] = model.state_dict()
if self.optimizer and not self.use_fairscale:
if (self.optimizer and not self.use_fairscale
and self.save_modules and "optimizer" in self.save_modules):
if self.use_fsdp and we.is_distributed:
ckpt['optimizer'] = OrderedDict()
for module in self.train_modules:
@@ -586,9 +619,6 @@ class LatentDiffusionSolver(BaseSolver):
'batch_size': len(batch_data['prompt'])
})
self.current_batch_data[self.mode] = batch_data
if self.sample_args:
self.current_batch_data[self.mode].update(
self.sample_args.get_lowercase_dict())
with torch.autocast(device_type='cuda',
enabled=self.use_amp,
dtype=self.dtype):
@@ -631,12 +661,12 @@ class LatentDiffusionSolver(BaseSolver):
# the inference image use
ret_images, ret_labels = [], []
if 'hint' in result:
ret_images.append(
(result['hint'][:result['image'].shape[0]].permute(
1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
ret_labels.append('Control Image')
ret_images.append((result['image'].permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_images.append((result['hint'][:result['image'].shape[0]].permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_labels.append(f"Control Image")
ret_images.append(
(result['image'].permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_labels.append(result['prompt'] +
" <font color='red'> |NegPrompt| </font> " +
result['n_prompt'])
@@ -678,15 +708,15 @@ class LatentDiffusionSolver(BaseSolver):
# the inference image use
ret_images, ret_labels = [], []
if 'hint' in result:
ret_images.append(
(result['hint'][:result['image'].shape[0]].permute(
1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
ret_labels.append('Control Image')
ret_images.append((result['image'].permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_images.append((result['hint'][:result['image'].shape[0]].permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_labels.append(f"Control Image")
ret_images.append(
(result['image'].permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_labels.append(result['prompt'] +
" <font color='red'> |NegPrompt| </font> " +
result['n_prompt'])
" <font color='red'> |NegPrompt| </font> " +
result['n_prompt'])
log_data.append(ret_images)
log_label.append(ret_labels)
ori_label.append(result['prompt'])
@@ -715,11 +745,8 @@ class LatentDiffusionSolver(BaseSolver):
swift_cfg_dict[f'{t_id}_{cfg_name}'] = init_config
if len(swift_cfg_dict) > 0:
from swift import Swift
model = Swift.prepare_model(self.model, config=swift_cfg_dict)
self.logger.info([(key, param.shape)
for key, param in model.named_parameters()
if param.requires_grad])
model = Swift.prepare_model(self.model, config=swift_cfg_dict, autocast_adapter_dtype=False)
self.logger.info([(key, param.shape) for key, param in model.named_parameters() if param.requires_grad])
return model
def freeze(self, freeze_cfg, model=None):
@@ -781,9 +808,7 @@ class LatentDiffusionSolver(BaseSolver):
for name, param in freeze_model.named_parameters():
if re.match(train_part, name):
param.requires_grad = True
self.logger.info([(key, param.shape)
for key, param in freeze_model.named_parameters()
if param.requires_grad])
self.logger.info([(key, param.shape) for key, param in freeze_model.named_parameters() if param.requires_grad])
return model
@torch.no_grad()
@@ -826,37 +851,32 @@ class LatentDiffusionSolver(BaseSolver):
@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])
batch_data = self.current_batch_data[self.mode]
if self.sample_args:
batch_data.update(self.sample_args.get_lowercase_dict())
self.eval_mode()
with torch.autocast(device_type='cuda',
enabled=self.use_amp,
dtype=self.dtype):
batch_data['log_num'] = self.log_train_num
results = self.run_step_eval(batch_data)
images = batch_data['image'] if 'image' in batch_data else [
None
] * len(results)
results = self.run_step_eval(transfer_data_to_cuda(batch_data))
images = batch_data['image'] if 'image' in batch_data else [None] * len(results)
self.train_mode()
log_data, log_label = [], []
for result, image in zip(results, images):
ret_images, ret_labels = [], []
if 'hint' in result:
ret_images.append(
(result['hint'][:result['image'].shape[0]].permute(
1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
ret_images.append((result['hint'][:result['image'].shape[0]].permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
if image is not None:
image = torch.clamp((image + 1.0) / 2.0, min=0.0, max=1.0)
ret_images.append((image.permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_labels.append('target image')
ret_images.append((image.permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
ret_labels.append(f'target image')
ret_images.append(
(result['image'].permute(1, 2, 0).cpu().numpy() *
255).astype(np.uint8))
ret_labels.append(result['prompt'] +
" <font color='red'> |NegPrompt| </font> " +
result['n_prompt'])
ret_images.append((result['image'].permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8))
ret_labels.append(result['prompt']
+ " <font color='red'> |NegPrompt| </font> "
+ result['n_prompt'])
log_data.append(ret_images)
log_label.append(ret_labels)
self.register_probe({
@@ -947,4 +967,4 @@ class LatentDiffusionSolver(BaseSolver):
logger.info(
f'Load ema frozen params {ema_param_numel} / {all_param_numel} = '
f'{ema_param_numel / all_param_numel:.2%}, '
f'frozen part: {ema_param_dict}.')
f'frozen part: {ema_param_dict}.')
@@ -0,0 +1,190 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import torch
import numpy as np
from tqdm import tqdm
from scepter.modules.utils.config import dict_to_yaml
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
@SOLVERS.register_class()
class LatentDiffusionVideoSolver(LatentDiffusionSolver):
para_dict = LatentDiffusionSolver.para_dict
def __init__(self, cfg, logger=None):
super().__init__(cfg, logger=logger)
self.fps = cfg.get("FPS", 8)
def save_results(self, results):
log_data, log_label = [], []
for result in results:
ret_videos, ret_labels = [], []
if 'edit_video' in result:
ret_videos.append((result['edit_video'].permute(1, 2, 3, 0).cpu().numpy() *
255).astype(np.uint8))
ret_labels.append("left: edit video")
if 'edit_image' in result:
ret_videos.append((result['edit_image'].permute(1, 2, 3, 0).cpu().numpy() *
255).astype(np.uint8))
ret_labels.append("left: edit image")
if 'target_video' in result:
if len(ret_videos) > 0:
ret_labels.append("middle: target video")
else:
ret_labels.append("left: target video")
ret_videos.append((result['target_video'].permute(1, 2, 3, 0).cpu().numpy() *
255).astype(np.uint8))
ret_videos.append((result['reconstruct_video'].permute(1, 2, 3, 0).cpu().numpy() *
255).astype(np.uint8))
ret_labels.append("right: generation video" + " Prompt: " + result['instruction'])
log_data.append(ret_videos)
log_label.append(ret_labels)
return log_data, log_label
def run_train(self):
self.train_mode()
self.before_all_iter(self.hooks_dict[self._mode])
data_iter = iter(self.datas[self._mode].dataloader)
self.print_memory_status()
for step in range(self.max_steps):
if 'eval' in self._mode_set and (self.eval_interval > 0 and
step % self.eval_interval == 0):
self.run_eval()
self.train_mode()
batch_data = next(data_iter)
self.before_iter(self.hooks_dict[self._mode])
if 'meta' in batch_data and isinstance(batch_data['meta'], dict):
self.register_probe({
'data_key':
ProbeData(batch_data['meta'].get('data_key', []),
view_distribute=True)
})
self.register_probe({
'prompt': batch_data['prompt'],
'batch_size': len(batch_data['prompt'])
})
self.current_batch_data[self.mode] = batch_data
if self.sample_args:
self.current_batch_data[self.mode].update(
self.sample_args.get_lowercase_dict())
batch_data = transfer_data_to_cuda(batch_data)
with torch.autocast(device_type='cuda',
enabled=self.use_amp,
dtype=self.dtype):
results = self.run_step_train(
batch_data,
step,
step=self.total_iter,
rank=we.rank)
self._iter_outputs[self._mode] = self._reduce_scalar(results)
self.after_iter(self.hooks_dict[self._mode])
if we.debug:
self.print_trainable_params_status(prefix='model.')
if 'eval' in self._mode_set and (self.eval_interval > 0
and step == self.max_steps - 1):
self.run_eval()
self.train_mode()
self.after_all_iter(self.hooks_dict[self._mode])
@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_video':
ProbeData(log_data,
is_image=False,
is_video=True,
fps=self.fps,
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_video':
ProbeData(log_data,
is_image=False,
is_video=True,
fps=self.fps,
build_html=True,
build_label=log_label)
})
self.after_all_iter(self.hooks_dict[self._mode])
@property
def probe_data(self):
if not we.debug and self.mode == 'train':
batch_data = self.current_batch_data[self.mode]
if self.sample_args is not None:
batch_data.update(self.sample_args.get_lowercase_dict())
self.eval_mode()
with torch.autocast(device_type='cuda',
enabled=self.use_amp,
dtype=self.dtype):
batch_data['log_train_num'] = self.log_train_num
all_results = self.run_step_eval(transfer_data_to_cuda(batch_data))
self.train_mode()
log_data, log_label = self.save_results(all_results)
self.register_probe({
'train_video':
ProbeData(log_data,
is_image=False,
is_video=True,
fps=self.fps,
build_html=True,
build_label=log_label)
})
self.register_probe({'train_label': log_label})
return super(LatentDiffusionSolver, self).probe_data
@staticmethod
def get_config_template():
return dict_to_yaml('SOLVER',
__class__.__name__,
LatentDiffusionVideoSolver.para_dict,
set_name=True)
+5 -5
View File
@@ -127,25 +127,25 @@ class BackwardHook(Hook):
if solver.scaler is not None:
solver.scaler.scale(solver.loss /
self.accumulate_step).backward()
if self.gradient_clip > 0:
solver.scaler.unscale_(solver.optimizer)
self.grad_clip(solver.train_parameters())
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:
solver.scaler.unscale_(solver.optimizer)
self.grad_clip(solver.train_parameters())
self.profile(solver)
solver.scaler.step(solver.optimizer)
solver.scaler.update()
solver.optimizer.zero_grad()
else:
(solver.loss / self.accumulate_step).backward()
if self.gradient_clip > 0:
self.grad_clip(solver.train_parameters())
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.train_parameters())
self.profile(solver)
solver.optimizer.step()
solver.optimizer.zero_grad()
+18 -7
View File
@@ -128,13 +128,24 @@ class CheckpointHook(Hook):
solver.work_dir,
'checkpoints/{}-{}'.format(self.save_name_prefix,
solver.total_iter + 1))
if we.rank == 0:
local_folder, _ = FS.map_to_local(save_path)
if hasattr(solver.model, 'module'):
solver.model.module.save_pretrained(local_folder)
else:
solver.model.save_pretrained(local_folder)
FS.put_dir_from_local_dir(local_folder, save_path)
solver_model = solver.model.module if hasattr(solver.model, 'module') else solver.model
if isinstance(solver_model.base_model.model, torch.distributed.fsdp.FullyShardedDataParallel):
full_state_dict_config = torch.distributed.fsdp.FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
with torch.distributed.fsdp.FullyShardedDataParallel.state_dict_type(solver_model.base_model, torch.distributed.fsdp.StateDictType.FULL_STATE_DICT, full_state_dict_config):
state_dict = solver_model.base_model.state_dict()
if we.rank == 0:
state_dict_new = {}
local_folder, _ = FS.map_to_local(save_path)
for adapter_name in solver_model.adapters.keys():
state_dict_adapter = solver_model.adapters[adapter_name].state_dict_callback(state_dict, adapter_name, replace_key=False)
state_dict_new.update(state_dict_adapter)
solver_model.save_pretrained(local_folder, state_dict=state_dict_new)
FS.put_dir_from_local_dir(local_folder, save_path)
else:
if we.rank == 0:
local_folder, _ = FS.map_to_local(save_path)
solver_model.save_pretrained(local_folder)
FS.put_dir_from_local_dir(local_folder, save_path)
else:
if hasattr(solver, 'save_pretrained'):
save_path = osp.join(
+8 -6
View File
@@ -117,6 +117,7 @@ class LogHook(Hook):
super(LogHook, self).__init__(cfg, logger=logger)
self.priority = cfg.get('PRIORITY', _DEFAULT_LOG_PRIORITY)
self.log_interval = cfg.get('LOG_INTERVAL', 10)
self.interval = cfg.get('INTERVAL', self.log_interval)
self.show_gpu_mem = cfg.get('SHOW_GPU_MEM', False)
self.log_agg_dict = defaultdict(LogAgg)
@@ -147,18 +148,18 @@ class LogHook(Hook):
outputs['time'] = iter_time
outputs['data_time'] = self.data_time
if solver.mode in self.batch_size:
outputs['throughput'] = int(self.batch_size[solver.mode] * we.world_size / iter_time * 86400)
outputs['throughput'] = int(self.batch_size[solver.mode] * we.data_group_world_size / iter_time * 86400)
log_agg.update(outputs, 1)
log_agg = log_agg.aggregate(self.log_interval)
log_agg = log_agg.aggregate(self.interval)
if 'throughput' in log_agg:
log_agg['throughput'] = f"{int(log_agg['throughput'][-1])}/day"
if solver.mode in self.batch_size:
log_agg['all_throughput'] = (solver.iter + 1) * we.world_size * self.batch_size[solver.mode]
log_agg['all_throughput'] = (solver.iter + 1) * we.data_group_world_size * self.batch_size[solver.mode]
if self.show_gpu_mem:
log_agg['nvidia-smi'] = str(print_memory_status()) +"MiB"
if (solver.iter + 1) % self.log_interval == 0:
if (solver.iter + 1) % self.interval == 0:
_print_iter_log(solver,
log_agg,
start_time=self.start_time,
@@ -206,7 +207,7 @@ class LogHook(Hook):
solver.logger.info(f'Current Epoch {mode} Summary:')
log_agg = self.log_agg_dict[mode]
_print_iter_log(solver,
log_agg.aggregate(self.log_interval),
log_agg.aggregate(self.interval),
start_time=self.start_time,
mode=mode)
if not mode == 'train':
@@ -242,6 +243,7 @@ class TensorboardLogHook(Hook):
self.priority = cfg.get('PRIORITY', _DEFAULT_LOG_PRIORITY)
self.log_dir = cfg.get('LOG_DIR', None)
self.log_interval = cfg.get('LOG_INTERVAL', 1000)
self.interval = cfg.get('INTERVAL', self.log_interval)
self._local_log_dir = None
self.writer: Optional[SummaryWriter] = None
@@ -286,7 +288,7 @@ class TensorboardLogHook(Hook):
self.writer.add_scalar(f'{mode}/iter/{key}',
value,
global_step=solver.total_iter)
if solver.total_iter % self.log_interval:
if solver.total_iter % self.interval:
self.writer.flush()
# Put to remote file systems every epoch
FS.put_dir_from_local_dir(self._local_log_dir, self.log_dir)
+8 -4
View File
@@ -721,6 +721,7 @@ class Workenv(object):
torch.backends.cudnn.benchmark = config.ENV.get(
'CUDNN_BENCHMARK', False)
fn(config)
return
else:
import torch.multiprocessing as mp
if 'MASTER_ADDR' not in os.environ:
@@ -741,10 +742,13 @@ class Workenv(object):
if self.is_distributed:
self.backend = config.ENV.get('BACKEND', 'nccl')
self.sync_bn = config.ENV.get('SYNC_BN', False)
mp.spawn(mp_worker,
nprocs=ngpus_per_node,
args=(ngpus_per_node, config, fn, pmi_rank, world_size,
self))
spawn_join = config.ENV.get('SPAWN_JOIN', True)
context = mp.spawn(mp_worker,
nprocs=ngpus_per_node,
join=spawn_join,
args=(ngpus_per_node, config, fn, pmi_rank, world_size,
self))
return context
def get_env(self):
ret_dict = {}
+139 -117
View File
@@ -1,6 +1,10 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import copy
from enum import Enum
import os
from scepter.modules.utils.file_system import FS
class Media(Enum):
@@ -9,18 +13,21 @@ class Media(Enum):
VIDEO = 3
AUDIO = 4
IMAGE_PAIR = 5
VIDEO_PAIR = 6
class HtmlVisualization(object):
def __init__(self,
allow_annotation=False,
slice_size=1000,
align='center',
width_scale='60%',
title='Visualization',
height=600,
width=None,
text_cols=40):
def __init__(
self,
allow_annotation=False,
slice_size=1000,
align='center',
width_scale='60%',
title='Visualization',
height=600,
width=None,
text_cols=40
):
self.content_list = []
self.rows_meta = []
self.allow_annotation = allow_annotation
@@ -30,9 +37,9 @@ class HtmlVisualization(object):
self.title = title
self.html_start = '<html>'
self.html_head = f'<head><meta charset="utf-8"><title>{title}</title></head>'
self.height = height if height is not None else '600'
self.width = width if width is not None else 'auto'
self.text_cols = text_cols if text_cols is not None else 'auto'
self.height = height if height is not None else "600"
self.width = width if width is not None else "auto"
self.text_cols = text_cols if text_cols is not None else "auto"
self.html_style = ('''
<style> \n
.container {
@@ -52,11 +59,21 @@ class HtmlVisualization(object):
height: 100%; \n
transition: 0.4s ease; \n
}\n
.image img {
.image img { \n
width: 100%; \n
height: 100%; \n
object-fit: contain; \n
} \n
.video { \n
display:flex; \n
position:absolute; \n
width:100%; \n
height:100%; \n
transition:0.4s ease; \n
object-fit:contain; \n
} \n
.slider {
position: absolute; \n
cursor: ew-resize; \n
@@ -64,13 +81,6 @@ class HtmlVisualization(object):
background-color: rgba(255, 255, 255, 0.5); \n
z-index: 10; \n
} \n
video { \n
width: auto;
height: 100%;
margin: 0px; \n
border: 0px solid #ccc; \n
padding: 0px; \n
} \n
textarea { \n
margin: 0px; \n
border: 0px; \n
@@ -78,10 +88,12 @@ class HtmlVisualization(object):
resize: none; \n
border: 1px solid #ccc; \n
} \n
.large-checkbox {transform: scale(2.5); margin-left: 20px; margin-bottom: 20px; vertical-align: middle;} \n
</style> \n
\n
'''.replace('{width_scale}', self.width_scale).replace(
'{align}', self.align).replace('{pair_height}', f'{self.height}'))
'''.replace('{width_scale}',
self.width_scale).replace('{align}', self.align)
.replace('{pair_height}', f'{self.height}'))
self.html_body_script = '''
<script>\n
@@ -89,7 +101,7 @@ class HtmlVisualization(object):
containers.forEach(container => {\n
let isDragging = true;\n
const slider = container.querySelector('.slider')\n
const image2 = container.querySelector('#image2')\n
const media2 = container.querySelector('#media2')\n
container.addEventListener('mousedown', () => {\n
isDragging = true;\n
@@ -113,7 +125,7 @@ class HtmlVisualization(object):
percentage = Math.max(0, Math.min(100, percentage));\n
image2.style.clipPath = `inset(0 ${100 - percentage}% 0 0)`;\n
media2.style.clipPath = `inset(0 ${100 - percentage}% 0 0)`;\n
slider.style.left = `${percentage}%`;\n
@@ -124,36 +136,6 @@ class HtmlVisualization(object):
slider.style.left = '50%';\n
});\n
</script>\n
<script> \n
const videos = document.querySelectorAll('video'); \n
\n
const observer = new IntersectionObserver((entries) => { \n
entries.forEach(entry => { \n
if (entry.isIntersecting) { \n
const video = entry.target; \n
video.src = video.dataset.src; \n
video.load(); \n
observer.unobserve(video); \n
} \n
}); \n
}); \n
\n
videos.forEach(video => { \n
observer.observe(video); \n
}); \n
\n
function adjustHeight() { \n
const textareas = document.querySelectorAll('textarea'); \n
textareas.forEach(textarea => { \n
const td = textarea.parentNode; \n
const tdHeight = td.clientHeight; \n
textarea.style.height = tdHeight + 'px'; \n
}); \n
} \n
window.onload = adjustHeight; \n
window.onresize = adjustHeight; \n
</script> \n
'''
self.html_body = '<body>{BODY}\n' + self.html_body_script + '</body>\n'
@@ -183,26 +165,27 @@ class HtmlVisualization(object):
'''
self.label_button = (
'<table><tr><td>' +
"<button style='height: 50px;' type=\"button\" onclick=\"saveSamples()\">Save Samples</button>"
+ '</td></tr></table>')
'<table><tr><td>' +
"<button style='height: 50px;' type=\"button\" onclick=\"saveSamples()\">Save Samples</button>"
+ '</td></tr></table>')
def format_col(self,
content='',
label='',
type=Media.TEXT,
show_label=True,
cols_span=1):
cols_span=1
):
if type == Media.TEXT:
ret_str = '<textarea' # noqa: E501
if self.height is not None:
rows = f"rows={self.height//30}"
rows = f"rows={self.height // 30}"
ret_str += f" {rows}"
if self.width is not None:
cols = f"cols={self.text_cols * cols_span}"
ret_str += f" {cols}"
ret_str += f'>"{content}"</textarea>'
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
elif type == Media.IMAGE:
ret_str = f'<img src="{content}"'
if self.height is not None:
@@ -212,7 +195,7 @@ class HtmlVisualization(object):
width = f'width="{self.width}"'
ret_str += f" {width}"
ret_str += ' >'
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
elif type == Media.VIDEO:
ret_str = '<video' # noqa
if self.height is not None:
@@ -221,61 +204,87 @@ class HtmlVisualization(object):
if self.width is not None:
width = f'width="{self.width}"'
ret_str += f" {width}"
ret_str += ' preload="none" controls>'
ret_str += ' preload="none" autoplay muted loop>'
ret_str += f'<source src="{content}" type="video/mp4"></video>'
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
elif type == Media.AUDIO:
ret_str = f'<audio src="{content}" controls>'
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
elif type == Media.IMAGE_PAIR:
assert isinstance(content, (list, tuple)) and len(content) == 2
ret_str = '\n'
ret_str += ' <div class="container"'
ret_str += (
f'> \n'
f' <div class="image" id="image1">' # noqa
f' <img src="{content[1]}" alt="before">\n' # noqa
f' </div>\n' # noqa
f' <div class="image" id="image2" style="clip-path: inset(0 50% 0 0);">\n' # noqa
f' <img src="{content[0]}" alt="after">\n' # noqa
f' </div>\n' # noqa
f' <div class="slider" id="slider"></div>\n' # noqa
f'')
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
ret_str = f'\n'
ret_str += f' <div class="container"'
ret_str += (f'> \n'
f' <div class="image" id="media1">'
f' <img src="{content[1]}" alt="before">\n'
f' </div>\n'
f' <div class="image" id="media2" style="clip-path: inset(0 50% 0 0);">\n'
f' <img src="{content[0]}" alt="after">\n'
f' </div>\n'
f' <div class="slider" id="slider"></div>\n'
f'')
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
elif type == Media.VIDEO_PAIR:
assert isinstance(content, (list, tuple)) and len(content) == 2
ret_str = f'\n'
ret_str += f' <div class="container"'
ret_str += (f'> \n'
f' <video autoplay muted loop class="video" id="media1"><source src="{content[1]}" type="video/mp4"></video>\n'
f' <video autoplay muted loop class="video" id="media2" style="clip-path: inset(0 50% 0 0);"><source src="{content[0]}" type="video/mp4"></video>\n'
f' <div class="slider" id="slider"></div>\n'
f'</div>')
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
else:
raise NotImplementedError
if self.allow_annotation:
ret_str = f'<label for="sample#sample_id#">{ret_str}</label>'
if cols_span > 1:
ret_str = f'<th colspan="{cols_span}">{ret_str}</th>\n'
sec_ret_str = f'<th colspan="{cols_span}">{sec_ret_str}</th>\n' if not sec_ret_str == '' else sec_ret_str
sec_ret_str = f'<th colspan="{cols_span}">{sec_ret_str}</th>\n' if not sec_ret_str == "" else sec_ret_str
else:
ret_str = f'<td>{ret_str}</td>\n'
sec_ret_str = f'<td align="center">{sec_ret_str}</td>\n' if not sec_ret_str == '' else sec_ret_str
sec_ret_str = f'<td align="center">{sec_ret_str}</td>\n' if not sec_ret_str == "" else sec_ret_str
return [ret_str, sec_ret_str]
def format_row(self):
sample_id = 0
all_sample_html = []
for one_content, one_row_meta in zip(self.content_list,
self.rows_meta):
one_row_str = '<tr>'
one_row_str += '\n'.join([v[0] for v in one_content])
if self.allow_annotation:
row_meta = '#;#'.join(one_row_meta)
one_row_str += (
f'<td><input type="checkbox" class="large-checkbox" '
f'id="sample{sample_id}" name="sample[]" value="{row_meta}"></td>\n'
)
one_row_str += '</tr><tr>'
one_row_str += '\n'.join([v[1] for v in one_content]) # noqa
if self.allow_annotation: # noqa
one_row_str += f'<td></td>\n' # noqa
one_row_str += '</tr>'
if self.allow_annotation:
one_row_str = f'<label for="sample{sample_id}">{one_row_str}</label>'
all_sample_html.append(one_row_str)
sample_id += 1
return '<table>' + '\n'.join(all_sample_html) + '</table>'
all_sample_html = []
current_content_list = copy.deepcopy(self.content_list)
current_rows_meta = copy.deepcopy(self.rows_meta)
while len(current_content_list) > 0:
sample_id = 0
batch_content_list = current_content_list[:self.slice_size]
current_content_list = current_content_list[self.slice_size:]
batch_rows_meta = current_rows_meta[:self.slice_size]
current_rows_meta = current_rows_meta[self.slice_size:]
current_sample_html = []
for one_content, one_row_meta in zip(batch_content_list,
batch_rows_meta):
one_row_str = '<tr>'
if not self.allow_annotation:
one_row_str += '\n'.join([v[0] for v in one_content])
else:
one_row_str += '\n'.join([v[0].replace('#sample_id#', f'{sample_id}') for v in one_content])
row_meta = '#;#'.join(one_row_meta)
one_row_str += (
f'<td><input type="checkbox" class="large-checkbox" '
f'id="sample{sample_id}" name="sample[]" value="{row_meta}"></td>\n'
)
one_row_str += '</tr><tr>'
one_row_str += '\n'.join([v[1] for v in one_content]) # noqa
if self.allow_annotation: # noqa
one_row_str += f'<td></td>\n' # noqa
one_row_str += '</tr>'
# if self.allow_annotation:
# one_row_str = f'<label for="sample{sample_id}">{one_row_str}</label>'
current_sample_html.append(one_row_str)
sample_id += 1
all_sample_html.append("<table>" + '\n'.join(current_sample_html) + "</table>")
return all_sample_html
def add_record(self,
content,
@@ -295,11 +304,8 @@ class HtmlVisualization(object):
if col_id > len(self.content_list[row_id]):
raise RuntimeError(
'col_id should be next number of the last col_id.')
format_col = self.format_col(content,
f"{row_id}-{col_id}: {label}",
type,
show_label=show_label,
cols_span=cols_span)
format_col = self.format_col(content, f"{row_id}-{col_id}: {label}",
type, show_label=show_label, cols_span=cols_span)
annotation_meta = annotation_meta if annotation_meta else ''
if col_id == len(self.content_list[row_id]):
@@ -311,14 +317,30 @@ class HtmlVisualization(object):
def save_html(self, path):
html_body = self.format_row()
ret_html_list = [
self.html_start, self.html_head, self.html_style,
self.html_body.replace('{BODY}', html_body)
]
if self.allow_annotation:
ret_html_list.append(self.label_button)
ret_html_list.append(self.html_script)
ret_html_list.append(self.html_end)
ret_html = '\n'.join(ret_html_list)
with open(path, 'w') as f:
f.write(ret_html)
if isinstance(html_body, list) and len(html_body) > 1:
try:
os.makedirs(path, exist_ok=True)
except:
print("Create folder path failed.")
for html_id, one_html in enumerate(html_body):
ret_html_list = [
self.html_start, self.html_head, self.html_style,
self.html_body.replace('{BODY}', one_html)
]
if self.allow_annotation:
ret_html_list.append(self.label_button)
ret_html_list.append(self.html_script)
ret_html_list.append(self.html_end)
ret_html = '\n'.join(ret_html_list)
FS.put_object(ret_html.encode(), os.path.join(path, f"{html_id}.html"))
else:
ret_html_list = [
self.html_start, self.html_head, self.html_style,
self.html_body.replace('{BODY}', html_body[0])
]
if self.allow_annotation:
ret_html_list.append(self.label_button)
ret_html_list.append(self.html_script)
ret_html_list.append(self.html_end)
ret_html = '\n'.join(ret_html_list)
FS.put_object(ret_html.encode(), path)
+263 -72
View File
@@ -12,15 +12,13 @@ import re
import string
import sys
import threading
import warnings
import cv2
import gradio as gr
import numpy as np
import torch
import transformers
from diffusers import CogVideoXImageToVideoPipeline
from diffusers.utils import export_to_video
from gradio_imageslider import ImageSlider
from PIL import Image
from transformers import AutoModel, AutoTokenizer
@@ -29,6 +27,7 @@ from scepter.modules.utils.config import Config
from scepter.modules.utils.directory import get_md5
from scepter.modules.utils.file_system import FS
from scepter.studio.utils.env import init_env
from importlib.metadata import version
from .example import get_examples
from .utils import load_image
@@ -51,32 +50,41 @@ class ChatBotUI(object):
is_debug=False,
language='en',
root_work_dir='./'):
cfg = Config(cfg_file=cfg_general_file)
try:
from diffusers import CogVideoXImageToVideoPipeline
from diffusers.utils import export_to_video
except Exception as e:
print(f"Import diffusers failed, please install or upgrade diffusers. Error information: {e}")
if isinstance(cfg_general_file, str):
cfg = Config(cfg_file=cfg_general_file)
else:
cfg = cfg_general_file
cfg.WORK_DIR = os.path.join(root_work_dir, cfg.WORK_DIR)
if not FS.exists(cfg.WORK_DIR):
FS.make_dir(cfg.WORK_DIR)
cfg = init_env(cfg)
self.cache_dir = cfg.WORK_DIR
self.chatbot_examples = get_examples(self.cache_dir)
self.chatbot_examples = get_examples(self.cache_dir) if not cfg.get('SKIP_EXAMPLES', False) else []
self.model_cfg_dir = cfg.MODEL.EDIT_MODEL.MODEL_CFG_DIR
self.model_yamls = glob.glob(os.path.join(self.model_cfg_dir,
'*.yaml'))
self.model_choices = dict()
self.default_model_name = ''
for i in self.model_yamls:
model_name = '.'.join(i.split('/')[-1].split('.')[:-1])
self.model_choices[model_name] = i
print('Models: ', self.model_choices)
self.model_name = cfg.MODEL.EDIT_MODEL.DEFAULT
assert self.model_name in self.model_choices
model_cfg = Config(load=True,
cfg_file=self.model_choices[self.model_name])
model_cfg = Config(load=True, cfg_file=i)
model_name = model_cfg.NAME
if model_cfg.IS_DEFAULT: self.default_model_name = model_name
self.model_choices[model_name] = model_cfg
print('Models: ', self.model_choices.keys())
assert len(self.model_choices) > 0
if self.default_model_name == "": self.default_model_name = list(self.model_choices.keys())[0]
self.model_name = self.default_model_name
self.pipe = ACEInference()
self.pipe.init_from_cfg(model_cfg)
self.pipe.init_from_cfg(self.model_choices[self.default_model_name])
self.max_msgs = 20
self.enable_i2v = cfg.get('ENABLE_I2V', False)
self.gradio_version = version('gradio')
if self.enable_i2v:
self.i2v_model_dir = cfg.MODEL.I2V.MODEL_DIR
self.i2v_model_name = cfg.MODEL.I2V.MODEL_NAME
@@ -167,6 +175,7 @@ class ChatBotUI(object):
]
def create_ui(self):
css = '.chatbot.prose.md {opacity: 1.0 !important} #chatbot {opacity: 1.0 !important}'
with gr.Blocks(css=css,
title='Chatbot',
@@ -177,7 +186,8 @@ class ChatBotUI(object):
self.history_result = gr.State(value={})
self.retry_msg = gr.State(value='')
with gr.Group():
with gr.Row(equal_height=True):
self.ui_mode = gr.State(value='legacy')
with gr.Row(equal_height=True, visible=False) as self.chat_group:
with gr.Column(visible=True) as self.chat_page:
self.chatbot = gr.Chatbot(
height=600,
@@ -192,7 +202,7 @@ class ChatBotUI(object):
size='sm')
with gr.Column(visible=False) as self.editor_page:
with gr.Tabs():
with gr.Tabs(visible=False) as self.upload_tabs:
with gr.Tab(id='ImageUploader',
label='Image Uploader',
visible=True) as self.upload_tab:
@@ -201,7 +211,7 @@ class ChatBotUI(object):
interactive=True,
type='pil',
image_mode='RGB',
sources='upload',
sources=['upload'],
elem_id='image_uploader',
format='png')
with gr.Row():
@@ -209,10 +219,9 @@ class ChatBotUI(object):
value='Submit',
elem_id='upload_submit')
self.ext_btn_1 = gr.Button(value='Exit')
with gr.Tabs(visible=False) as self.edit_tabs:
with gr.Tab(id='ImageEditor',
label='Image Editor',
visible=False) as self.edit_tab:
label='Image Editor') as self.edit_tab:
self.mask_type = gr.Dropdown(
label='Mask Type',
choices=[
@@ -275,13 +284,23 @@ class ChatBotUI(object):
self.ext_btn_2 = gr.Button(value='Exit')
with gr.Tab(id='ImageViewer',
label='Image Viewer',
visible=False) as self.image_view_tab:
self.image_viewer = ImageSlider(
label='Image',
type='pil',
show_download_button=True,
elem_id='image_viewer')
label='Image Viewer') as self.image_view_tab:
if self.gradio_version >= '5.0.0':
self.image_viewer = gr.Image(
label='Image',
type='pil',
show_download_button=True,
elem_id='image_viewer')
else:
try:
from gradio_imageslider import ImageSlider
except Exception as e:
print(f"Import gradio_imageslider failed, please install.")
self.image_viewer = ImageSlider(
label='Image',
type='pil',
show_download_button=True,
elem_id='image_viewer')
self.ext_btn_3 = gr.Button(value='Exit')
@@ -300,11 +319,30 @@ class ChatBotUI(object):
self.ext_btn_4 = gr.Button(value='Exit')
with gr.Row(equal_height=True, visible=True) as self.legacy_group:
with gr.Column():
self.legacy_image_uploader = gr.Image(
height=550,
interactive=True,
type='pil',
image_mode='RGB',
elem_id='legacy_image_uploader',
format='png')
with gr.Column():
self.legacy_image_viewer = gr.Image(
label='Image',
height=550,
type='pil',
interactive=False,
show_download_button=True,
elem_id='image_viewer')
with gr.Accordion(label='Setting', open=False):
with gr.Row():
self.model_name_dd = gr.Dropdown(
choices=self.model_choices,
value=self.model_name,
value=self.default_model_name,
label='Model Version')
with gr.Row():
@@ -315,39 +353,63 @@ class ChatBotUI(object):
label='Negative Prompt',
container=False)
with gr.Row():
# REFINER_PROMPT
self.refiner_prompt = gr.Textbox(
value=self.pipe.input.get("refiner_prompt", ""),
visible=self.pipe.input.get("refiner_prompt", None) is not None,
placeholder=
'Prompt used for refiner',
label='Refiner Prompt',
container=False)
with gr.Row():
with gr.Column(scale=8, min_width=500):
with gr.Row():
self.step = gr.Slider(minimum=1,
maximum=1000,
value=20,
value=self.pipe.input.get("sample_steps", 20),
visible=self.pipe.input.get("sample_steps", None) is not None,
label='Sample Step')
self.cfg_scale = gr.Slider(
minimum=1.0,
maximum=20.0,
value=4.5,
value=self.pipe.input.get("guide_scale", 4.5),
visible=self.pipe.input.get("guide_scale", None) is not None,
label='Guidance Scale')
self.rescale = gr.Slider(minimum=0.0,
maximum=1.0,
value=0.5,
value=self.pipe.input.get("guide_rescale", 0.5),
visible=self.pipe.input.get("guide_rescale", None) is not None,
label='Rescale')
self.refiner_scale = gr.Slider(minimum=-0.1,
maximum=1.0,
value=self.pipe.input.get("refiner_scale", -1),
visible=self.pipe.input.get("refiner_scale", None) is not None,
label='Refiner Scale')
self.seed = gr.Slider(minimum=-1,
maximum=10000000,
value=-1,
label='Seed')
self.output_height = gr.Slider(
minimum=256,
maximum=1024,
value=512,
maximum=1440,
value=self.pipe.input.get("output_height", 1024),
visible=self.pipe.input.get("output_height", None) is not None,
label='Output Height')
self.output_width = gr.Slider(
minimum=256,
maximum=1024,
value=512,
maximum=1440,
value=self.pipe.input.get("output_width", 1024),
visible=self.pipe.input.get("output_width", None) is not None,
label='Output Width')
with gr.Column(scale=1, min_width=50):
self.use_history = gr.Checkbox(value=False,
label='Use History')
self.use_ace = gr.Checkbox(value=self.pipe.input.get("use_ace", True),
visible=self.pipe.input.get("use_ace", None) is not None,
label='Use ACE')
self.video_auto = gr.Checkbox(
value=False,
label='Auto Gen Video',
@@ -384,7 +446,7 @@ class ChatBotUI(object):
visible=True)
with gr.Row():
inst = """
self.chatbot_inst = """
**Instruction**:
1. Click 'Upload' button to upload one or more images as input images.
@@ -398,12 +460,25 @@ class ChatBotUI(object):
8. If you find our work valuable, we invite you to refer to the [ACE Page](https://ali-vilab.github.io/ace-page/) for comprehensive information.
"""
gr.Markdown(value=inst)
self.legacy_inst = """
**Instruction**:
1. You can edit the image by uploading it; if no image is uploaded, an image will be generated from text..
2. Enter '@' in the text box will exhibit all images in the gallery.
3. Select the image you wish to edit from the gallery, and its Image ID will be displayed in the text box.
4. **Important** To render text on an image, please ensure to include a space between each letter. For instance, "add text 'g i r l' on the mask area of @xxxxx".
5. To perform multi-step editing, partial editing, inpainting, outpainting, and other operations, please click the Chatbot Checkbox to enable the conversational editing mode and follow the relevant instructions..
6. If you find our work valuable, we invite you to refer to the [ACE Page](https://ali-vilab.github.io/ace-page/) for comprehensive information.
"""
self.instruction = gr.Markdown(value=self.legacy_inst)
with gr.Row(variant='panel',
equal_height=True,
show_progress=False):
with gr.Column(scale=1, min_width=100):
with gr.Column(scale=1, min_width=100, visible=False) as self.upload_panel:
self.upload_btn = gr.Button(value=upload_sty +
' Upload',
variant='secondary')
@@ -413,12 +488,16 @@ class ChatBotUI(object):
label='Instruction',
container=False)
with gr.Column(scale=1, min_width=100):
self.chat_btn = gr.Button(value=chat_sty + ' Chat',
self.chat_btn = gr.Button(value='Generate',
variant='primary')
with gr.Column(scale=1, min_width=100):
self.retry_btn = gr.Button(value=refresh_sty +
' Retry',
variant='secondary')
with gr.Column(scale=1, min_width=100):
self.mode_checkbox = gr.Checkbox(
value=False,
label='ChatBot')
with gr.Column(scale=(1 if self.enable_i2v else 0),
min_width=0):
self.video_gen_btn = gr.Button(value=video_sty +
@@ -453,19 +532,78 @@ class ChatBotUI(object):
lock.acquire()
del self.pipe
torch.cuda.empty_cache()
model_cfg = Config(load=True,
cfg_file=self.model_choices[model_name])
torch.cuda.ipc_collect()
self.pipe = ACEInference()
self.pipe.init_from_cfg(model_cfg)
self.pipe.init_from_cfg(self.model_choices[model_name])
self.model_name = model_name
lock.release()
return model_name, gr.update(), gr.update()
return (model_name, gr.update(), gr.update(),
gr.Slider(
value=self.pipe.input.get("sample_steps", 20),
visible=self.pipe.input.get("sample_steps", None) is not None),
gr.Slider(
value=self.pipe.input.get("guide_scale", 4.5),
visible=self.pipe.input.get("guide_scale", None) is not None),
gr.Slider(
value=self.pipe.input.get("guide_rescale", 0.5),
visible=self.pipe.input.get("guide_rescale", None) is not None),
gr.Slider(
value=self.pipe.input.get("output_height", 1024),
visible=self.pipe.input.get("output_height", None) is not None),
gr.Slider(
value=self.pipe.input.get("output_width", 1024),
visible=self.pipe.input.get("output_width", None) is not None),
gr.Textbox(
value=self.pipe.input.get("refiner_prompt", ""),
visible=self.pipe.input.get("refiner_prompt", None) is not None),
gr.Slider(
value=self.pipe.input.get("refiner_scale", -1),
visible=self.pipe.input.get("refiner_scale", None) is not None
),
gr.Checkbox(
value=self.pipe.input.get("use_ace", True),
visible=self.pipe.input.get("use_ace", None) is not None
)
)
self.model_name_dd.change(
change_model,
inputs=[self.model_name_dd],
outputs=[self.model_name_dd, self.chatbot, self.text])
outputs=[
self.model_name_dd, self.chatbot, self.text,
self.step,
self.cfg_scale, self.rescale, self.output_height,
self.output_width, self.refiner_prompt, self.refiner_scale,
self.use_ace])
def mode_change(mode_check):
if mode_check:
# ChatBot
return (
gr.Row(visible=False),
gr.Row(visible=True),
gr.Button(value='Generate'),
gr.State(value='chatbot'),
gr.Column(visible=True),
gr.Markdown(value=self.chatbot_inst)
)
else:
# Legacy
return (
gr.Row(visible=True),
gr.Row(visible=False),
gr.Button(value=chat_sty + ' Chat'),
gr.State(value='legacy'),
gr.Column(visible=False),
gr.Markdown(value=self.legacy_inst)
)
self.mode_checkbox.change(mode_change, inputs=[self.mode_checkbox],
outputs=[self.legacy_group, self.chat_group,
self.chat_btn, self.ui_mode,
self.upload_panel, self.instruction])
########################################
def generate_gallery(text, images):
@@ -522,6 +660,9 @@ class ChatBotUI(object):
fps,
seed,
progress=gr.Progress(track_tqdm=True)):
from diffusers.utils import export_to_video
generator = torch.Generator(device='cuda').manual_seed(seed)
img_ids = re.findall('@(.*?)[ ,;.?$]', message)
if len(img_ids) == 0:
@@ -592,7 +733,11 @@ class ChatBotUI(object):
outputs=[self.history, self.chatbot, self.text, self.gallery])
########################################
def run_chat(message,
def run_chat(
message,
legacy_image,
ui_mode,
use_ace,
extend_prompt,
history,
images,
@@ -601,6 +746,8 @@ class ChatBotUI(object):
negative_prompt,
cfg_scale,
rescale,
refiner_prompt,
refiner_scale,
step,
seed,
output_h,
@@ -612,12 +759,25 @@ class ChatBotUI(object):
video_fps,
video_seed,
progress=gr.Progress(track_tqdm=True)):
legacy_img_ids = []
if ui_mode == 'legacy':
if legacy_image is not None:
history, images, img_id = self.add_uploaded_image_to_history(
legacy_image, history, images)
legacy_img_ids.append(img_id)
retry_msg = message
gen_id = get_md5(message)[:12]
save_path = os.path.join(self.cache_dir, f'{gen_id}.png')
img_ids = re.findall('@(.*?)[ ,;.?$]', message)
history_io = None
if len(img_ids) < 1:
img_ids = legacy_img_ids
for img_id in img_ids:
if f'@{img_id}' not in message:
message = f'@{img_id} ' + message
new_message = message
if len(img_ids) > 0:
@@ -676,6 +836,9 @@ class ChatBotUI(object):
guide_scale=cfg_scale,
guide_rescale=rescale,
seed=seed,
refiner_prompt=refiner_prompt,
refiner_scale=refiner_scale,
use_ace=use_ace
)
img = imgs[0]
@@ -784,21 +947,25 @@ class ChatBotUI(object):
while len(history) >= self.max_msgs:
history.pop(0)
return history, images, history_result, self.get_history(
history), gr.update(value=''), gr.update(
visible=False), retry_msg
return (history, images, gr.Image(value=save_path),
history_result, self.get_history(
history), gr.update(), gr.update(
visible=False), retry_msg)
chat_inputs = [
self.legacy_image_uploader, self.ui_mode, self.use_ace,
self.extend_prompt, self.history, self.images, self.use_history,
self.history_result, self.negative_prompt, self.cfg_scale,
self.rescale, self.step, self.seed, self.output_height,
self.rescale, self.refiner_prompt, self.refiner_scale,
self.step, self.seed, self.output_height,
self.output_width, self.video_auto, self.video_step,
self.video_frames, self.video_cfg_scale, self.video_fps,
self.video_seed
]
chat_outputs = [
self.history, self.images, self.history_result, self.chatbot,
self.history, self.images, self.legacy_image_viewer,
self.history_result, self.chatbot,
self.text, self.gallery, self.retry_msg
]
@@ -832,9 +999,13 @@ class ChatBotUI(object):
w = int(w / ratio)
img = img.resize((w, h))
edit_image.append(img)
if img_mask is not None:
img_mask = img_mask if np.sum(np.array(img_mask)) > 0 else None
edit_image_mask.append(
img_mask if img_mask is not None else None)
edit_task.append(task)
if ref1 is not None:
ref1 = ref1 if np.sum(np.array(ref1)) > 0 else None
if ref1 is not None:
edit_image.append(ref1)
edit_image_mask.append(None)
@@ -859,6 +1030,8 @@ class ChatBotUI(object):
prompt=[prompt] * img_num,
negative_prompt=[''] * img_num,
seed=seed,
refiner_prompt=self.pipe.input.get("refiner_prompt", ""),
refiner_scale=self.pipe.input.get("refiner_scale", 0.0),
)
img = imgs[0]
@@ -868,8 +1041,13 @@ class ChatBotUI(object):
img_str = f'<img src="data:image/png;base64,{img_b64}" style="pointer-events: none;">'
history = [(prompt,
f'{pre_info} The generated image is:\n {img_str}')]
img_id = get_md5(img_b64)[:12]
save_path = os.path.join(self.cache_dir, f'{img_id}.png')
img.convert('RGB').save(save_path)
return self.get_history(history), gr.update(value=''), gr.update(
visible=False), gr.update(value=-1)
visible=False), gr.Image(value=save_path), gr.update(value=-1)
with self.eg:
self.example_task = gr.Text(label='Task Name',
@@ -895,8 +1073,9 @@ class ChatBotUI(object):
self.example_task, self.example_image, self.example_mask,
self.example_ref_im1, self.text, self.seed
],
outputs=[self.chatbot, self.text, self.gallery, self.seed],
outputs=[self.chatbot, self.text, self.gallery, self.legacy_image_viewer, self.seed],
examples_per_page=4,
cache_examples=False,
run_on_click=True)
########################################
@@ -904,14 +1083,16 @@ class ChatBotUI(object):
return (gr.update(visible=True,
scale=1), gr.update(visible=True, scale=1),
gr.update(visible=True), gr.update(visible=False),
gr.update(visible=False), gr.update(visible=False))
gr.update(visible=False), gr.update(visible=False),
gr.update(visible=True))
self.upload_btn.click(upload_image,
inputs=[],
outputs=[
self.chat_page, self.editor_page,
self.upload_tab, self.edit_tab,
self.image_view_tab, self.video_view_tab
self.image_view_tab, self.video_view_tab,
self.upload_tabs
])
########################################
@@ -926,13 +1107,19 @@ class ChatBotUI(object):
]
if len(imgs) > 0:
if len(imgs) == 2:
view_img = copy.deepcopy(imgs)
if self.gradio_version >= '5.0.0':
view_img = copy.deepcopy(imgs[-1])
else:
view_img = copy.deepcopy(imgs)
edit_img = copy.deepcopy(imgs[-1])
else:
view_img = [
copy.deepcopy(imgs[-1]),
copy.deepcopy(imgs[-1])
]
if self.gradio_version >= '5.0.0':
view_img = copy.deepcopy(imgs[-1])
else:
view_img = [
copy.deepcopy(imgs[-1]),
copy.deepcopy(imgs[-1])
]
edit_img = copy.deepcopy(imgs[-1])
return (gr.update(visible=True,
@@ -941,11 +1128,12 @@ class ChatBotUI(object):
gr.update(visible=False), gr.update(visible=True),
gr.update(visible=True), gr.update(visible=False),
gr.update(value=edit_img),
gr.update(value=view_img), gr.update(value=None))
gr.update(value=view_img), gr.update(value=None),
gr.update(visible=True))
else:
return (gr.update(), gr.update(), gr.update(), gr.update(),
gr.update(), gr.update(), gr.update(), gr.update(),
gr.update())
gr.update(), gr.update())
elif isinstance(evt.value, dict) and evt.value.get(
'component', '') == 'video':
value = evt.value['value']['video']['path']
@@ -953,11 +1141,12 @@ class ChatBotUI(object):
scale=1), gr.update(visible=True, scale=1),
gr.update(visible=False), gr.update(visible=False),
gr.update(visible=False), gr.update(visible=True),
gr.update(), gr.update(), gr.update(value=value))
gr.update(), gr.update(), gr.update(value=value),
gr.update())
else:
return (gr.update(), gr.update(), gr.update(), gr.update(),
gr.update(), gr.update(), gr.update(), gr.update(),
gr.update())
gr.update(), gr.update())
self.chatbot.select(edit_image,
outputs=[
@@ -965,16 +1154,17 @@ class ChatBotUI(object):
self.upload_tab, self.edit_tab,
self.image_view_tab, self.video_view_tab,
self.image_editor, self.image_viewer,
self.video_viewer
self.video_viewer, self.edit_tabs
])
self.image_viewer.change(lambda x: x,
inputs=self.image_viewer,
outputs=self.image_viewer)
if self.gradio_version < '5.0.0':
self.image_viewer.change(lambda x: x,
inputs=self.image_viewer,
outputs=self.image_viewer)
########################################
def submit_upload_image(image, history, images):
history, images = self.add_uploaded_image_to_history(
history, images, _ = self.add_uploaded_image_to_history(
image, history, images)
return gr.update(visible=False), gr.update(
visible=True), gr.update(
@@ -1207,13 +1397,13 @@ class ChatBotUI(object):
history.append(
(None,
f'This is uploaded image:\n {img_str} image ID is: {img_id}'))
return history, images
return history, images, img_id
def run_gr(cfg):
with gr.Blocks() as demo:
chatbot = ChatBotUI(cfg)
chatbot.create_bot_ui()
chatbot.create_ui()
chatbot.set_callbacks()
demo.launch(server_name='0.0.0.0',
server_port=cfg.args.server_port,
@@ -1225,6 +1415,7 @@ if __name__ == '__main__':
parser.add_argument('--server_port',
dest='server_port',
help='',
type=int,
default=2345)
parser.add_argument('--root_path', dest='root_path', help='', default='')
cfg = Config(load=True, parser_ins=parser)
+69 -56
View File
@@ -3,6 +3,7 @@
import os
from scepter.modules.utils.file_system import FS
from PIL import Image
def download_image(image, local_path=None):
@@ -10,44 +11,56 @@ def download_image(image, local_path=None):
local_path = FS.get_from(image, local_path=local_path)
return local_path
def blank_image():
return Image.new('RGBA', (128, 128), (0, 0, 0, 0))
def get_examples(cache_dir):
print('Downloading Examples ...')
bl_img = blank_image()
examples = [
[
'Facial Editing',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/e33edc106953.png?raw=true',
os.path.join(cache_dir, 'examples/e33edc106953.png')), None,
None, '{image} let the man smile', 6666
os.path.join(cache_dir, 'examples/e33edc106953.png')), bl_img,
bl_img, '{image} let the man smile', 6666
],
[
'Facial Editing',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/5d2bcc91a3e9.png?raw=true',
os.path.join(cache_dir, 'examples/5d2bcc91a3e9.png')), None,
None, 'let the man in {image} wear sunglasses', 9999
os.path.join(cache_dir, 'examples/5d2bcc91a3e9.png')), bl_img,
bl_img, 'let the man in {image} wear sunglasses', 9999
],
[
'Facial Editing',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/5d2bcc91a3e9.png?raw=true',
os.path.join(cache_dir, 'examples/5d2bcc91a3e9.png')), bl_img,
bl_img, 'let the man in {image} wear sunglasses', 9999
],
[
'Facial Editing',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/3a52eac708bd.png?raw=true',
os.path.join(cache_dir, 'examples/3a52eac708bd.png')), None,
None, '{image} red hair', 9999
os.path.join(cache_dir, 'examples/3a52eac708bd.png')), bl_img,
bl_img, '{image} red hair', 9999
],
[
'Facial Editing',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/3f4dc464a0ea.png?raw=true',
os.path.join(cache_dir, 'examples/3f4dc464a0ea.png')), None,
None, '{image} let the man serious', 99999
os.path.join(cache_dir, 'examples/3f4dc464a0ea.png')), bl_img,
bl_img, '{image} let the man serious', 99999
],
[
'Controllable Generation',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/131ca90fd2a9.png?raw=true',
os.path.join(cache_dir,
'examples/131ca90fd2a9.png')), None, None,
'examples/131ca90fd2a9.png')), bl_img, bl_img,
'"A person sits contemplatively on the ground, surrounded by falling autumn leaves. Dressed in a green sweater and dark blue pants, they rest their chin on their hand, exuding a relaxed demeanor. Their stylish checkered slip-on shoes add a touch of flair, while a black purse lies in their lap. The backdrop of muted brown enhances the warm, cozy atmosphere of the scene." , generate the image that corresponds to the given scribble {image}.',
613725
],
@@ -59,7 +72,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/33e9f27c2c48_mask.png?raw=true',
os.path.join(cache_dir,
'examples/33e9f27c2c48_mask.png')), None,
'examples/33e9f27c2c48_mask.png')), bl_img,
'Put the text "C A T" at the position marked by mask in the {image}',
6666
],
@@ -67,7 +80,7 @@ def get_examples(cache_dir):
'Style Transfer',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/9e73e7eeef55.png?raw=true',
os.path.join(cache_dir, 'examples/9e73e7eeef55.png')), None,
os.path.join(cache_dir, 'examples/9e73e7eeef55.png')), bl_img,
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/2e02975293d6.png?raw=true',
os.path.join(cache_dir, 'examples/2e02975293d6.png')),
@@ -81,7 +94,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/f2b22c08be3f_mask.png?raw=true',
os.path.join(cache_dir,
'examples/f2b22c08be3f_mask.png')), None,
'examples/f2b22c08be3f_mask.png')), bl_img,
'Could the {image} be widened within the space designated by mask, while retaining the original?',
6666
],
@@ -89,57 +102,57 @@ def get_examples(cache_dir):
'Image Segmentation',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/db3ebaa81899.png?raw=true',
os.path.join(cache_dir, 'examples/db3ebaa81899.png')), None,
None, '{image} Segmentation', 6666
os.path.join(cache_dir, 'examples/db3ebaa81899.png')), bl_img,
bl_img, '{image} Segmentation', 6666
],
[
'Depth Estimation',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/f1927c4692ba.png?raw=true',
os.path.join(cache_dir, 'examples/f1927c4692ba.png')), None,
None, '{image} Depth Estimation', 6666
os.path.join(cache_dir, 'examples/f1927c4692ba.png')), bl_img,
bl_img, '{image} Depth Estimation', 6666
],
[
'Pose Estimation',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/014e5bf3b4d1.png?raw=true',
os.path.join(cache_dir, 'examples/014e5bf3b4d1.png')), None,
None, '{image} distinguish the poses of the figures', 999999
os.path.join(cache_dir, 'examples/014e5bf3b4d1.png')), bl_img,
bl_img, '{image} distinguish the poses of the figures', 999999
],
[
'Scribble Extraction',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/5f59a202f8ac.png?raw=true',
os.path.join(cache_dir, 'examples/5f59a202f8ac.png')), None,
None, 'Generate a scribble of {image}, please.', 6666
os.path.join(cache_dir, 'examples/5f59a202f8ac.png')), bl_img,
bl_img, 'Generate a scribble of {image}, please.', 6666
],
[
'Mosaic',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/3a2f52361eea.png?raw=true',
os.path.join(cache_dir, 'examples/3a2f52361eea.png')), None,
None, 'Adapt {image} into a mosaic representation.', 6666
os.path.join(cache_dir, 'examples/3a2f52361eea.png')), bl_img,
bl_img, 'Adapt {image} into a mosaic representation.', 6666
],
[
'Edge map Extraction',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/b9d1e519d6e5.png?raw=true',
os.path.join(cache_dir, 'examples/b9d1e519d6e5.png')), None,
None, 'Get the edge-enhanced result for {image}.', 6666
os.path.join(cache_dir, 'examples/b9d1e519d6e5.png')), bl_img,
bl_img, 'Get the edge-enhanced result for {image}.', 6666
],
[
'Grayscale',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/c4ebbe2ba29b.png?raw=true',
os.path.join(cache_dir, 'examples/c4ebbe2ba29b.png')), None,
None, 'transform {image} into a black and white one', 6666
os.path.join(cache_dir, 'examples/c4ebbe2ba29b.png')), bl_img,
bl_img, 'transform {image} into a black and white one', 6666
],
[
'Contour Extraction',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/19652d0f6c4b.png?raw=true',
os.path.join(cache_dir,
'examples/19652d0f6c4b.png')), None, None,
'examples/19652d0f6c4b.png')), bl_img, bl_img,
'Would you be able to make a contour picture from {image} for me?',
6666
],
@@ -148,7 +161,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/249cda2844b7.png?raw=true',
os.path.join(cache_dir,
'examples/249cda2844b7.png')), None, None,
'examples/249cda2844b7.png')), bl_img, bl_img,
'Following the segmentation outcome in mask of {image}, develop a real-life image using the explanatory note in "a mighty cat lying on the bed”.',
6666
],
@@ -157,7 +170,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/411f6c4b8e6c.png?raw=true',
os.path.join(cache_dir,
'examples/411f6c4b8e6c.png')), None, None,
'examples/411f6c4b8e6c.png')), bl_img, bl_img,
'use the depth map {image} and the text caption "a cut white cat" to create a corresponding graphic image',
999999
],
@@ -166,7 +179,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/a35c96ed137a.png?raw=true',
os.path.join(cache_dir,
'examples/a35c96ed137a.png')), None, None,
'examples/a35c96ed137a.png')), bl_img, bl_img,
'help translate this posture schema {image} into a colored image based on the context I provided "A beautiful woman Climbing the climbing wall, wearing a harness and climbing gear, skillfully maneuvering up the wall with her back to the camera, with a safety rope."',
3599999
],
@@ -175,7 +188,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/dcb2fc86f1ce.png?raw=true',
os.path.join(cache_dir,
'examples/dcb2fc86f1ce.png')), None, None,
'examples/dcb2fc86f1ce.png')), bl_img, bl_img,
'Transform and generate an image using mosaic {image} and "Monarch butterflies gracefully perch on vibrant purple flowers, showcasing their striking orange and black wings in a lush garden setting." description',
6666
],
@@ -184,7 +197,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/4cd4ee494962.png?raw=true',
os.path.join(cache_dir,
'examples/4cd4ee494962.png')), None, None,
'examples/4cd4ee494962.png')), bl_img, bl_img,
'make this {image} colorful as per the "beautiful sunflowers"',
6666
],
@@ -193,7 +206,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/a47e3a9cd166.png?raw=true',
os.path.join(cache_dir,
'examples/a47e3a9cd166.png')), None, None,
'examples/a47e3a9cd166.png')), bl_img, bl_img,
'Take the edge conscious {image} and the written guideline "A whimsical animated character is depicted holding a delectable cake adorned with blue and white frosting and a drizzle of chocolate. The character wears a yellow headband with a bow, matching a cozy yellow sweater. Her dark hair is styled in a braid, tied with a yellow ribbon. With a golden fork in hand, she stands ready to enjoy a slice, exuding an air of joyful anticipation. The scene is creatively rendered with a charming and playful aesthetic." and produce a realistic image.',
613725
],
@@ -202,7 +215,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/d890ed8a3ac2.png?raw=true',
os.path.join(cache_dir,
'examples/d890ed8a3ac2.png')), None, None,
'examples/d890ed8a3ac2.png')), bl_img, bl_img,
'creating a vivid image based on {image} and description "This image features a delicious rectangular tart with a flaky, golden-brown crust. The tart is topped with evenly sliced tomatoes, layered over a creamy cheese filling. Aromatic herbs are sprinkled on top, adding a touch of green and enhancing the visual appeal. The background includes a soft, textured fabric and scattered white flowers, creating an elegant and inviting presentation. Bright red tomatoes in the upper right corner hint at the fresh ingredients used in the dish."',
6666
],
@@ -211,7 +224,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/0844a686a179.png?raw=true',
os.path.join(cache_dir,
'examples/0844a686a179.png')), None, None,
'examples/0844a686a179.png')), bl_img, bl_img,
'Eliminate noise interference in {image} and maximize the crispness to obtain superior high-definition quality',
6666
],
@@ -223,7 +236,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/fa91b6b7e59b_mask.png?raw=true',
os.path.join(cache_dir,
'examples/fa91b6b7e59b_mask.png')), None,
'examples/fa91b6b7e59b_mask.png')), bl_img,
'Ensure to overhaul the parts of the {image} indicated by the mask.',
6666
],
@@ -235,7 +248,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/632899695b26_mask.png?raw=true',
os.path.join(cache_dir,
'examples/632899695b26_mask.png')), None,
'examples/632899695b26_mask.png')), bl_img,
'Refashion the mask portion of {image} in accordance with "A yellow egg with a smiling face painted on it"',
6666
],
@@ -244,7 +257,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/354d17594afe.png?raw=true',
os.path.join(cache_dir,
'examples/354d17594afe.png')), None, None,
'examples/354d17594afe.png')), bl_img, bl_img,
'{image} change the dog\'s posture to walking in the water, and change the background to green plants and a pond.',
6666
],
@@ -253,7 +266,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/38946455752b.png?raw=true',
os.path.join(cache_dir,
'examples/38946455752b.png')), None, None,
'examples/38946455752b.png')), bl_img, bl_img,
'{image} change the color of the dress from white to red and the model\'s hair color red brown to blonde.Other parts remain unchanged',
6669
],
@@ -262,7 +275,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/3ba5202f0cd8.png?raw=true',
os.path.join(cache_dir,
'examples/3ba5202f0cd8.png')), None, None,
'examples/3ba5202f0cd8.png')), bl_img, bl_img,
'Keep the same facial feature in @3ba5202f0cd8, change the woman\'s clothing from a Blue denim jacket to a white turtleneck sweater and adjust her posture so that she is supporting her chin with both hands. Other aspects, such as background, hairstyle, facial expression, etc, remain unchanged.',
99999
],
@@ -270,22 +283,22 @@ def get_examples(cache_dir):
'Facial Editing',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/369365b94725.png?raw=true',
os.path.join(cache_dir, 'examples/369365b94725.png')), None,
None, '{image} Make her looking at the camera', 6666
os.path.join(cache_dir, 'examples/369365b94725.png')), bl_img,
bl_img, '{image} Make her looking at the camera', 6666
],
[
'Facial Editing',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/92751f2e4a0e.png?raw=true',
os.path.join(cache_dir, 'examples/92751f2e4a0e.png')), None,
None, '{image} Remove the smile from his face', 9899999
os.path.join(cache_dir, 'examples/92751f2e4a0e.png')), bl_img,
bl_img, '{image} Remove the smile from his face', 9899999
],
[
'Remove Text',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/8530a6711b2e.png?raw=true',
os.path.join(cache_dir, 'examples/8530a6711b2e.png')), None,
None, 'Aim to remove any textual element in {image}', 6666
os.path.join(cache_dir, 'examples/8530a6711b2e.png')), bl_img,
bl_img, 'Aim to remove any textual element in {image}', 6666
],
[
'Remove Text',
@@ -295,7 +308,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/c4d7fb28f8f6_mask.png?raw=true',
os.path.join(cache_dir,
'examples/c4d7fb28f8f6_mask.png')), None,
'examples/c4d7fb28f8f6_mask.png')), bl_img,
'Rub out any text found in the mask sector of the {image}.', 6666
],
[
@@ -303,7 +316,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/e2f318fa5e5b.png?raw=true',
os.path.join(cache_dir,
'examples/e2f318fa5e5b.png')), None, None,
'examples/e2f318fa5e5b.png')), bl_img, bl_img,
'Remove the unicorn in this {image}, ensuring a smooth edit.',
99999
],
@@ -315,7 +328,7 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/1ae96d8aca00_mask.png?raw=true',
os.path.join(cache_dir, 'examples/1ae96d8aca00_mask.png')),
None, 'Discard the contents of the mask area from {image}.', 99999
bl_img, 'Discard the contents of the mask area from {image}.', 99999
],
[
'Add Object',
@@ -325,22 +338,22 @@ def get_examples(cache_dir):
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/80289f48e511_mask.png?raw=true',
os.path.join(cache_dir,
'examples/80289f48e511_mask.png')), None,
'examples/80289f48e511_mask.png')), bl_img,
'add a Hot Air Balloon into the {image}, per the mask', 613725
],
[
'Style Transfer',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/d725cb2009e8.png?raw=true',
os.path.join(cache_dir, 'examples/d725cb2009e8.png')), None,
None, 'Change the style of {image} to colored pencil style', 99999
os.path.join(cache_dir, 'examples/d725cb2009e8.png')), bl_img,
bl_img, 'Change the style of {image} to colored pencil style', 99999
],
[
'Style Transfer',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/e0f48b3fd010.png?raw=true',
os.path.join(cache_dir, 'examples/e0f48b3fd010.png')), None,
None, 'make {image} to Walt Disney Animation style', 99999
os.path.join(cache_dir, 'examples/e0f48b3fd010.png')), bl_img,
bl_img, 'make {image} to Walt Disney Animation style', 99999
],
[
'Try On',
@@ -359,8 +372,8 @@ def get_examples(cache_dir):
'Workflow',
download_image(
'https://github.com/ali-vilab/ace-page/blob/main/assets/examples/cb85353c004b.png?raw=true',
os.path.join(cache_dir, 'examples/cb85353c004b.png')), None,
None, '<workflow> ice cream {image}', 99999
os.path.join(cache_dir, 'examples/cb85353c004b.png')), bl_img,
bl_img, '<workflow> ice cream {image}', 99999
],
]
print('Finish. Start building UI ...')
@@ -6,6 +6,7 @@ from scepter.modules.inference.largen_inference import LargenInference
from scepter.modules.inference.pixart_inference import PixArtInference
from scepter.modules.inference.sd3_inference import SD3Inference
from scepter.modules.inference.stylebooth_inference import StyleboothInference
from scepter.modules.inference.cogvideox_inference import CogVideoXInference
from scepter.modules.utils.logger import get_logger
@@ -111,6 +112,8 @@ class PipelineManager():
PipelineBuilder = SD3Inference
elif pipeline_name.startswith('FLUX'):
PipelineBuilder = FluxInference
elif pipeline_name.startswith('COGVIDEO'):
PipelineBuilder = CogVideoXInference
else:
PipelineBuilder = DiffusionInference
new_inference = PipelineBuilder(logger=self.logger)
@@ -98,6 +98,8 @@ class DiffusionUIName():
self.sample = 'Sampler'
self.sample_steps = 'Sample Steps'
self.image_number = 'Images Number'
self.num_frames = 'Number of Frames'
self.fps = 'Number of FPS'
self.resolutions_height = 'Output Height'
self.resolutions_width = 'Output Width'
self.negative_prompt = 'Negative Prompt'
@@ -121,6 +123,8 @@ class DiffusionUIName():
self.sample = '采样器'
self.sample_steps = '采样步数'
self.image_number = '图片数量'
self.num_frames = '视频帧数量'
self.fps = '视频帧率'
self.resolutions_height = '输出高度'
self.resolutions_width = '输出宽度'
self.negative_prompt = '负向提示'
@@ -100,6 +100,7 @@ class ControlUI(UIBase):
label=self.component_names.control_model,
choices=self.controller_choices,
value=self.controller_default,
allow_custom_value=True,
interactive=True)
with gr.Column(scale=1, min_width=0):
self.cond_button = gr.Button('Extract')
@@ -26,7 +26,7 @@ class DiffusionUI(UIBase):
self.default_resolutions = pipe_manager.pipeline_level_modules[
now_pipeline].paras.RESOLUTIONS
self.default_input = pipe_manager.pipeline_level_modules[
now_pipeline].input
now_pipeline].input_cfg
self.diffusion_paras = self.load_all_paras()
# deal with resolution
@@ -66,7 +66,6 @@ class DiffusionUI(UIBase):
if value is not None and cur_default.get(
key.lower()) not in value:
value.append(cur_default.get(key.lower()))
return diffusion_paras
def load_all_paras(self):
@@ -115,22 +114,15 @@ class DiffusionUI(UIBase):
label=self.component_names.resolutions_height,
choices=[key for key in self.cur_h_level_dict.keys()],
value=default_res[0],
allow_custom_value=True,
interactive=True)
with gr.Column(scale=1):
self.output_width = gr.Dropdown(
label=self.component_names.resolutions_width,
choices=self.cur_h_level_dict[default_res[0]],
value=default_res[1],
allow_custom_value=True,
interactive=True)
with gr.Row(equal_height=True):
self.image_number = gr.Slider(
label=self.component_names.image_number,
minimum=self.cur_paras.SAMPLES.get('MIN', 1),
maximum=self.cur_paras.SAMPLES.get('MAX', 4),
step=1,
value=self.cur_paras.SAMPLES.get('DEFAULT', 1),
visible=self.cur_paras.SAMPLES.get('VISIBLE', True),
interactive=True)
with gr.Row(equal_height=True):
self.sample_steps = gr.Slider(
label=self.component_names.sample_steps,
@@ -139,7 +131,6 @@ class DiffusionUI(UIBase):
step=1,
value=self.cur_paras.SAMPLE_STEPS.get('DEFAULT', 30),
interactive=True)
self.guide_scale = gr.Slider(
label=self.component_names.guide_scale,
minimum=self.cur_paras.GUIDE_SCALE.get('MIN', 1),
@@ -156,6 +147,32 @@ class DiffusionUI(UIBase):
value=self.cur_paras.GUIDE_RESCALE.get('DEFAULT', 0.5),
visible=self.cur_paras.GUIDE_RESCALE.get('VISIBLE', True),
interactive=True)
with gr.Row(equal_height=True):
self.fps = gr.Slider(
label=self.component_names.fps,
minimum=self.cur_paras.FPS.get('MIN', 1),
maximum=self.cur_paras.FPS.get('MAX', 50),
step=1,
value=self.cur_paras.FPS.get('DEFAULT', 8),
visible=self.cur_paras.FPS.get('VISIBLE', True),
interactive=True)
self.num_frames = gr.Slider(
label=self.component_names.num_frames,
minimum=self.cur_paras.NUM_FRAMES.get('MIN', 1),
maximum=self.cur_paras.NUM_FRAMES.get('MAX', 100),
step=1,
value=self.cur_paras.NUM_FRAMES.get('DEFAULT', 49),
visible=self.cur_paras.NUM_FRAMES.get('VISIBLE', True),
interactive=True)
with gr.Row(equal_height=True):
self.image_number = gr.Slider(
label=self.component_names.image_number,
minimum=self.cur_paras.SAMPLES.get('MIN', 1),
maximum=self.cur_paras.SAMPLES.get('MAX', 4),
step=1,
value=self.cur_paras.SAMPLES.get('DEFAULT', 1),
visible=self.cur_paras.SAMPLES.get('VISIBLE', True),
interactive=True)
with gr.Row(equal_height=True):
with gr.Column(scale=1):
self.seed_random = gr.Checkbox(
@@ -177,6 +194,8 @@ class DiffusionUI(UIBase):
'output_height': self.output_height,
'output_width': self.output_width,
'image_number': self.image_number,
'num_frames': self.num_frames,
'fps': self.fps,
'sample_steps': self.sample_steps,
'guide_scale': self.guide_scale,
'guide_rescale': self.guide_rescale,
@@ -1,6 +1,7 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import os
import imageio
from collections import OrderedDict
import gradio as gr
@@ -183,8 +184,10 @@ class GalleryUI(UIBase):
'guide_scale':
args.pop('guide_scale'),
'guide_rescale':
args.pop('guide_rescale')
args.pop('guide_rescale'),
}
if 'num_frames' in args:
pipeline_input.update({'num_frames': args['num_frames']})
args.update({'input': pipeline_input})
def appedix_init(args):
@@ -217,7 +220,7 @@ class GalleryUI(UIBase):
load_init(args)
results = current_pipeline(**args)
images = []
images, videos = [], []
before_images = []
if 'images' in results:
images_tensor = results['images'] * 255
@@ -226,6 +229,11 @@ class GalleryUI(UIBase):
1, 2, 0).cpu().numpy().astype(np.uint8))
for idx in range(images_tensor.shape[0])
]
if 'videos' in results:
videos = [
(video.permute(1, 2, 3, 0).cpu().numpy() * 255).astype(np.uint8)
for video in results['videos']
]
if 'before_refine_images' in results and results[
'before_refine_images'] is not None:
before_refine_images_tensor = results['before_refine_images'] * 255
@@ -236,7 +244,7 @@ class GalleryUI(UIBase):
]
if 'seed' in results:
print(results['seed'])
print(images, before_images)
# print(images, before_images)
if args['largen_state']:
largen_history.extend(images)
@@ -250,12 +258,27 @@ class GalleryUI(UIBase):
f'cur_gallery_{i}.jpg')
img.save(save_image)
save_list.append(save_image)
images = save_list
ret_data = save_list
else:
ret_data = images
if len(videos) > 0:
save_list = []
fps = args.get('fps', 8)
for i, video in enumerate(videos):
save_video_path = os.path.join(self.local_work_dir,
f'cur_gallery_{i}.mp4')
writer = imageio.get_writer(save_video_path, fps=fps)
for frame in video:
writer.append_data(np.array(frame))
writer.close()
save_list.append(save_video_path)
ret_data = save_list
return (
gr.Column(visible=len(before_images) > 0),
before_images,
images,
ret_data,
largen_history,
gr.update(value=largen_history),
)
@@ -188,10 +188,11 @@ class ModelManageUI(UIBase):
diffusion_ui.cur_h_level_dict = h_level_dict
default_input = self.pipe_manager.pipeline_level_modules[
now_pipeline].input
now_pipeline].input_cfg
cur_paras = diffusion_ui.get_default(diffusion_ui.diffusion_paras,
default_input)
diffusion_ui.cur_paras = cur_paras
return (
diffusion_model,
gr.Dropdown(value=all_module_name['first_stage_model']),
@@ -224,7 +225,11 @@ class ModelManageUI(UIBase):
gr.Slider(value=cur_paras.GUIDE_SCALE.get('DEFAULT', 7.5),
visible=cur_paras.GUIDE_SCALE.get('VISIBLE', True)),
gr.Slider(value=cur_paras.GUIDE_RESCALE.get('DEFAULT', 0.5),
visible=cur_paras.GUIDE_RESCALE.get('VISIBLE', True)))
visible=cur_paras.GUIDE_RESCALE.get('VISIBLE', True)),
gr.Slider(value=cur_paras.NUM_FRAMES.get('DEFAULT', 49),
visible=cur_paras.NUM_FRAMES.get('VISIBLE', False)),
gr.Slider(value=cur_paras.FPS.get('DEFAULT', 8),
visible=cur_paras.FPS.get('VISIBLE', False)))
self.diffusion_model.change(
diffusion_model_change,
@@ -240,6 +245,7 @@ class ModelManageUI(UIBase):
diffusion_ui.prompt_prefix, diffusion_ui.output_height,
diffusion_ui.sampler, diffusion_ui.discretization,
diffusion_ui.sample_steps, diffusion_ui.guide_scale,
diffusion_ui.guide_rescale
diffusion_ui.guide_rescale, diffusion_ui.num_frames,
diffusion_ui.fps
],
queue=True)
@@ -22,7 +22,8 @@ class CreateDatasetUIName():
self.dataset_type = 'Dataset Type'
self.dataset_type_name = {
'scepter_txt2img': 'Text2Image Generation',
'scepter_img2img': 'Image Edit Generation'
'scepter_img2img': 'Image Edit Generation',
'scepter_txt2vid': 'Text2Video Generation'
}
self.user_data_name = (
f'Current Dataset Name. Changes of dataset name take '
@@ -34,30 +35,35 @@ class CreateDatasetUIName():
'scepter_txt2img':
f'{self.default_dataset_repo}repo?Revision=master&FilePath=datasets/3D_example_csv.zip',
'scepter_img2img':
f'{self.default_dataset_repo}repo?Revision=master&FilePath=datasets/hed_pair.zip'
f'{self.default_dataset_repo}repo?Revision=master&FilePath=datasets/hed_pair.zip',
'scepter_txt2vid':
f'{self.default_dataset_repo}repo?Revision=master&FilePath=datasets/video_example.zip'
}
self.default_dataset_zip_str = ' and '.join(
[f'[{k}]({v})' for k, v in self.default_dataset_zip.items()])
self.default_dataset_name = {
'scepter_txt2img': '3D_example',
'scepter_img2img': 'hed_example'
'scepter_img2img': 'hed_example',
'scepter_txt2vid': 'video_example'
}
self.btn_create_datasets_from_file = 'Create Dataset From File'
self.user_direction = (
'### User Guide: \n' +
f'* {self.btn_create_datasets} button is used to create a new dataset '
". Please make sure to modify the dataset's name and version. After creation, "
'you can upload images one by one. \n'
'you can upload images or videos one by one. \n'
f'* The "{self.btn_create_datasets_from_file}" button supports creating a new dataset from '
'a file, currently supporting zip files. For zip files, the format should be consistent'
" with the one used during training, ensuring it contains an 'images/' folder and a '"
"train.csv' (which will use the image paths in this file); "
'The first line is Target:FILE, Prompt, followed by the format of each line: image path, description.'
" with the one used during training, ensuring it contains an 'images/' or 'videos/' folder and a '"
"train.csv' (which will use the image or video paths in this file); "
'The first line is Target:FILE, Prompt, followed by the format of each line: image path or video path, '
'description.'
'we also surpport the zip of '
'one level subfolder of images whose format are in jpg, jpeg, png, webp. '
'one level subfolder of images or videos whose format are in jpg, jpeg, png, mp4, webp. '
f'See the ZIP examples: {self.default_dataset_zip_str}. \n' # noqa
'Addition, txt2vid data also supports batch upload of txt file list, followed by the format of '
'each line: video path#;#video description '
f'* If you have refreshed the page, please click the {self.refresh_list_button} '
'button to ensure all previously created datasets are visible in the dropdown menu.\n'
'* For processing and training with large-scale data(for example more than 10K samples), '
@@ -97,7 +103,8 @@ class CreateDatasetUIName():
self.dataset_type = '数据集类型'
self.dataset_type_name = {
'scepter_txt2img': '文生图数据',
'scepter_img2img': '图像编辑(图生图)数据'
'scepter_img2img': '图像编辑(图生图)数据',
'scepter_txt2vid': '文生视频数据'
}
self.user_data_name = f'当前数据集名称,修改后点{self.modify_data_button}生效'
@@ -108,26 +115,30 @@ class CreateDatasetUIName():
'scepter_txt2img':
f'{self.default_dataset_repo}repo?Revision=master&FilePath=datasets/3D_example_csv.zip',
'scepter_img2img':
f'{self.default_dataset_repo}repo?Revision=master&FilePath=datasets/hed_pair.zip'
f'{self.default_dataset_repo}repo?Revision=master&FilePath=datasets/hed_pair.zip',
'scepter_txt2vid':
f'{self.default_dataset_repo}repo?Revision=master&FilePath=datasets/video_example.zip'
}
self.default_dataset_zip_str = ' 和 '.join(
[f'[{k}]({v})' for k, v in self.default_dataset_zip.items()])
self.default_dataset_name = {
'scepter_txt2img': '3D_example',
'scepter_img2img': 'hed_example'
'scepter_img2img': 'hed_example',
'scepter_txt2vid': 'video_example'
}
self.btn_create_datasets_from_file = '从文件新建'
self.user_direction = (
'### 使用说明 \n' +
f'* {self.btn_create_datasets} 按钮用于从零新建数据集,请注意修改数据集的name和version,'
'新建完成后可以逐个上传图片。\n' +
'新建完成后可以逐个上传图片或视频。\n' +
f'* {self.btn_create_datasets_from_file} 按钮支持从文件中来新建数据集,目前支持zip文件,'
'需要保证在文件夹外进行打包,并包含images/文件夹和train.csv(会使用该文件中的图片路径),首行为Target:FILE,Prompt,'
'其次每行格式为:图片路径,描述;'
f'同时我们也支持图像文件的zip包,格式在jpg、jpeg、png或webp。数据ZIP样例路径:{self.default_dataset_zip_str}. \n'
'需要保证在文件夹外进行打包,并包含 images/ 或 videos/ 文件夹和train.csv(会使用该文件中的图片或视频路径),首行为Target:FILE,Prompt,'
'其次每行格式为:图片 或 视频 路径,描述;'
f'同时我们也支持图像或视频文件的zip包,格式在jpg、jpeg、png、mp4或webp。数据ZIP样例路径:{self.default_dataset_zip_str}; \n'
'另外,文生视频数据还支持txt文件列表批量上传,文件每行格式为:视频路径#;#视频描述;\n '
+
f'* 如果刷新了页面,请点击{self.refresh_list_button} 按钮以确保所有以往创建的数据集在下拉框中可见。\n'
f'如果刷新了页面,请点击 {self.refresh_list_button} 按钮以确保所有以往创建的数据集在下拉框中可见\n'
'* 对于大规模数据的处理和训练(数据规模大于1万),建议使用命令行形式\n'
'* <span style="color: blue;">请注意观察系统日志的输出以帮助改进操作。</span> \n')
# Error or Warning
@@ -154,11 +165,13 @@ class DatasetGalleryUIName():
self.illegal_blank_dataset = 'Illgal or blank dataset is not allowed editing.'
self.delete_blank_dataset = 'Blank dataset is not allowed deleting.'
self.upload_image = 'Upload Target Image'
self.upload_video = 'Upload Video'
self.upload_src_image = 'Upload Source Image'
self.upload_src_mask = 'Mask Image'
self.upload_image_btn = '\U00002714' # ✔️
self.cancel_upload_btn = '\U00002716' # ✖️
self.image_caption = 'Image Caption'
self.video_caption = 'Video Caption'
self.btn_modify = '\U0001F4DD' # 📝
self.btn_delete = '\U0001f5d1' # 🗑️
@@ -196,10 +209,12 @@ class DatasetGalleryUIName():
f'click{self.btn_reset_edit} to reset edited data,'
f'click{self.btn_cancel_edit} to out of editing mode.')
self.preprocess_choices = [
'Image Preprocess', 'Caption Preprocess'
'Image Preprocess', 'Caption Preprocess', 'Caption translation'
]
self.preprocess_choices_video = ['Video caption generation', 'Caption translation']
self.preview_target_image = 'Preview Target Image'
self.preview_target_video = 'Preview Target Video'
self.preview_src_image = 'Preview Source Image'
self.preview_src_mask_image = 'Preview Source Image Mask'
self.preview_caption = 'Preview Caption'
@@ -211,7 +226,7 @@ class DatasetGalleryUIName():
self.caption_preprocess_btn = 'apply'
self.caption_preview_btn = 'preview'
self.caption_update_mode = 'Caption Update Mode'
self.caption_update_choices = ['Append', 'Replace']
self.caption_update_choices = ['Replace', 'Append']
self.used_device = 'Used Device'
self.used_memory = 'Used Memory'
@@ -234,11 +249,13 @@ class DatasetGalleryUIName():
self.illegal_blank_dataset = '不合法或空白数据集不允许编辑。'
self.delete_blank_dataset = '空白数据集不允许删除。'
self.upload_image = '上传目标图片'
self.upload_video = '上传视频'
self.upload_src_image = '上传待编辑图片'
self.upload_src_mask = '蒙版区域'
self.upload_image_btn = '\U00002714' # ✔️
self.cancel_upload_btn = '\U00002716' # ✖️
self.image_caption = '图片描述'
self.video_caption = '视频描述'
self.btn_modify = '\U0001F4DD' # 📝
self.dataset_images = f'图片数据,点击{self.btn_modify}进入编辑模式'
@@ -256,8 +273,8 @@ class DatasetGalleryUIName():
self.edit_caption = '编辑描述'
self.batch_caption_generate = '处理范围'
self.ori_dataset = '原始数据 高({}) * 宽({}) 图像格式({})'
self.edit_dataset = '可编辑数据 高({}) * 宽({}) 图像格式({})'
self.ori_dataset = '原始数据 高({}) * 宽({}) 格式({})'
self.edit_dataset = '可编辑数据 高({}) * 宽({}) 格式({})'
self.upload_image_info = '图像信息 高({}) * 宽({})'
self.upload_src_image_info = '源图像信息 高({}) * 宽({})'
@@ -276,8 +293,10 @@ class DatasetGalleryUIName():
f'点击{self.btn_cancel_edit}取消编辑,'
f'点击{self.btn_reset_edit}重置数据,'
f'修改编辑范围可以批量编辑不同范围的数据。')
self.preprocess_choices = ['图像预处理', '描述生成']
self.preprocess_choices = ['图像预处理', '描述生成', '描述翻译']
self.preprocess_choices_video = ['视频描述生成', '描述翻译']
self.preview_target_image = '预览图片'
self.preview_target_video = '预览视频'
self.preview_src_image = '预览原图'
self.preview_src_mask_image = '预览蒙版'
self.preview_caption = '预览描述'
@@ -288,7 +307,7 @@ class DatasetGalleryUIName():
self.caption_preprocess_btn = '应用'
self.caption_preview_btn = '预览'
self.caption_update_mode = '描述更新方式'
self.caption_update_choices = ['追加', '替换']
self.caption_update_choices = ['替换', '追加']
self.used_device = '使用设备'
self.used_memory = '使用内存'
self.caption_language = '描述语言'
@@ -386,3 +405,33 @@ class Image2ImageDataCardName():
self.illegal_data_err7 = '上传图像失败{}'
self.delete_err1 = '删除失败,数据已经为空了'
self.export_zip_err1 = '压缩文件失败!'
class Text2VideoDataCardName():
def __init__(self, language='en'):
if language == 'en':
self.illegal_data_err1 = (
'The list supports only "," or "#;#" as delimiters. '
'The two columns represent video path and description, '
'respectively.')
self.illegal_data_err2 = 'Illegal file format'
self.illegal_data_err3 = 'File decompression failed, failed to upload to storage!'
self.illegal_data_err4 = 'Illegal width({}),height({})'
self.illegal_data_err5 = (
'The path should not contain "{}". '
'It should be an OSS path (oss://) or the prefix '
'can be omitted (xxx/xxx)."')
self.illegal_data_err6 = 'Video download failed {}'
self.illegal_data_err7 = 'Video upload failed {}'
self.delete_err1 = 'Deletion failed, the data is already empty.'
self.export_zip_err1 = 'Failed to compress the file!'
elif language == 'zh':
self.illegal_data_err1 = '列表只支持,或#;#作为分割符,两列分别为视频路径/描述'
self.illegal_data_err2 = '非法的文件格式'
self.illegal_data_err3 = '文件解压失败,上传存储器失败!'
self.illegal_data_err4 = '不合法的width({}),height({})'
self.illegal_data_err5 = '路径不支持{},应该为oss路径(oss://)或者省略前缀(xxx/xxx)'
self.illegal_data_err6 = '下载视频失败{}'
self.illegal_data_err7 = '上传视频失败{}'
self.delete_err1 = '删除失败,数据已经为空了'
self.export_zip_err1 = '压缩文件失败!'
@@ -15,6 +15,8 @@ from scepter.studio.preprocess.utils.img2img_data_card import \
Image2ImageDataCard
from scepter.studio.preprocess.utils.txt2img_data_card import \
Text2ImageDataCard
from scepter.studio.preprocess.utils.txt2vid_data_card import \
Text2VideoDataCard
from scepter.studio.utils.uibase import UIBase
from tqdm import tqdm
@@ -44,7 +46,9 @@ class CreateDatasetUI(UIBase):
'scepter_txt2img':
Text2ImageDataCard,
'scepter_img2img':
Image2ImageDataCard
Image2ImageDataCard,
'scepter_txt2vid':
Text2VideoDataCard
})
self.components_name = CreateDatasetUIName(language)
self.default_dataset_type = list(self.dataset_type_dict.keys())[0]
@@ -466,8 +470,7 @@ class CreateDatasetUI(UIBase):
], [
self.panel_state, self.dataset_name, self.user_dataset_name,
self.sys_log
],
queue=False)
], queue=False)
def show_edit_panel(panel_state, data_name):
if panel_state:
@@ -599,10 +602,9 @@ class CreateDatasetUI(UIBase):
trans_dataset_type, [])
else:
dataset_list = []
return gr.Dropdown(
value=dataset_list[-1] if len(dataset_list) > 0 else '',
choices=dataset_list), self.components_name.system_log.format(
'')
return (gr.Dropdown(
value=dataset_list[-1] if len(dataset_list) > 0 else '', choices=dataset_list),
self.components_name.system_log.format(''))
manager.user_name.change(dataset_type_change,
inputs=[self.dataset_type, manager.user_name],
File diff suppressed because it is too large Load Diff
@@ -15,8 +15,11 @@ from scepter.modules.utils.file_system import FS
import numpy as np
from scepter.studio.preprocess.processors.base_processor import \
BaseCaptionProcessor
import io
from decord import cpu, VideoReader, bridge
__all__ = ['BlipImageBase', 'QWVL', 'QWVLQuantize', 'InternVL15']
__all__ = ['BlipImageBase', 'QWVL', 'QWVLQuantize', 'InternVL15', 'CogVLM2Llama3Caption',
'OpusMtZhEn', 'OpusMtEnZh']
def get_region(image, mask, mask_id):
locs = np.where(np.array(mask) == mask_id)
@@ -467,3 +470,306 @@ class InternVL15(QWVL):
response = self.model_info['model'].chat(self.model_info['tokenizer'], image, prompt, generation_config=generation_config)
response = response.replace("\n", "").strip()
return response
class CogVLM2Llama3Caption(BaseCaptionProcessor):
def __init__(self, cfg, language='en'):
super().__init__(cfg, language=language)
self.model_path = cfg.MODEL_PATH
self.model_info = {
'device': 'offline',
'model': None,
'tokenizer': None
}
self.prompt = cfg.PROMPT
self.temperature = cfg.TEMPERATURE
self.max_new_tokens = cfg.MAX_NEW_TOKENS
self.pad_token_id = cfg.PAD_TOKEN_ID
self.top_k = cfg.TOP_K
self.top_p = cfg.TOP_P
self.TORCH_TYPE = torch.bfloat16 if (torch.cuda.is_available() and
torch.cuda.get_device_capability()
[0] >= 8) else torch.float16
def load_model(self):
is_flg, msg = super().load_model()
if not is_flg:
return is_flg, msg
if self.model_info['device'] == 'offline':
model = None
try:
from transformers import AutoModelForCausalLM, AutoTokenizer
local_model_dir = FS.get_dir_to_local_dir(self.model_path)
tokenizer = AutoTokenizer.from_pretrained(
local_model_dir,
trust_remote_code=True,
)
model = AutoModelForCausalLM.from_pretrained(
local_model_dir,
device_map='auto',
torch_dtype=self.TORCH_TYPE,
trust_remote_code=True
).eval().to(we.device_id)
except Exception as e:
if model is not None:
del model
return False, f"Load model error '{e}'"
self.model_info['device'] = model.device
self.model_info['model'] = model
self.model_info['tokenizer'] = tokenizer
elif self.model_info['device'] == 'cpu':
try:
self.model_info['model'].to(we.device_id)
self.model_info['device'] = we.device_id
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
except Exception as e:
del self.model_info['model']
self.model_info['model'] = None
self.model_info['device'] = 'offline'
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
return False, f"Load model error '{e}'"
return True, ''
def unload_model(self):
super().unload_model()
if self.delete_instance:
self.model_info['device'] = 'offline'
if self.model_info['model'] is not None:
self.model_info['model'] = self.model_info['model'].to('cpu')
del self.model_info['model']
self.model_info['model'] = None
elif (isinstance(self.model_info['device'], numbers.Number)
or str(self.model_info['device']).startswith('cuda')):
self.model_info['device'] = 'cpu'
self.model_info['model'] = self.model_info['model'].to('cpu')
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
return True, ''
def load_video(self, video_data, strategy='chat'):
bridge.set_bridge('torch')
mp4_stream = video_data
num_frames = 24
decord_vr = VideoReader(io.BytesIO(mp4_stream), ctx=cpu(0))
frame_id_list = None
total_frames = len(decord_vr)
if strategy == 'base':
clip_end_sec = 60
clip_start_sec = 0
start_frame = int(clip_start_sec * decord_vr.get_avg_fps())
end_frame = min(total_frames,
int(clip_end_sec * decord_vr.get_avg_fps())) if clip_end_sec is not None else total_frames
frame_id_list = np.linspace(start_frame, end_frame - 1, num_frames, dtype=int)
elif strategy == 'chat':
timestamps = decord_vr.get_frame_timestamp(np.arange(total_frames))
timestamps = [i[0] for i in timestamps]
max_second = round(max(timestamps)) + 1
frame_id_list = []
for second in range(max_second):
closest_num = min(timestamps, key=lambda x: abs(x - second))
index = timestamps.index(closest_num)
frame_id_list.append(index)
if len(frame_id_list) >= num_frames:
break
video_data = decord_vr.get_batch(frame_id_list)
video_data = video_data.permute(3, 0, 1, 2)
return video_data
def get_caption(self, prompt, video_data, temperature):
strategy = 'chat'
video = self.load_video(video_data, strategy=strategy)
history = []
query = prompt
model = self.model_info['model']
tokenizer = self.model_info['tokenizer']
inputs = model.build_conversation_input_ids(
tokenizer=tokenizer,
query=query,
images=[video],
history=history,
template_version=strategy
)
inputs = {
'input_ids': inputs['input_ids'].unsqueeze(0).to('cuda'),
'token_type_ids': inputs['token_type_ids'].unsqueeze(0).to(we.device_id),
'attention_mask': inputs['attention_mask'].unsqueeze(0).to(we.device_id),
'images': [[inputs['images'][0].to(we.device_id).to(self.TORCH_TYPE)]],
}
gen_kwargs = {
"max_new_tokens": self.max_new_tokens,
"pad_token_id": self.pad_token_id,
"top_k": self.top_k,
"do_sample": True,
"top_p": self.top_p,
"temperature": temperature,
}
with torch.no_grad():
outputs = model.generate(**inputs, **gen_kwargs)
outputs = outputs[:, inputs['input_ids'].shape[1]:]
response = tokenizer.decode(outputs[0], skip_special_tokens=True)
return response
def __call__(self, *args, **kwargs):
video_path = kwargs.pop('video_path', None)
with open(video_path, 'rb') as f:
video_data = f.read()
response = (self.get_caption(self.prompt, video_data, self.temperature))
return response
class OpusMtZhEn(BaseCaptionProcessor):
def __init__(self, cfg, language='en'):
super().__init__(cfg, language=language)
self.model_path = cfg.MODEL_PATH
self.model_info = {
'device': 'offline',
'model': None,
'tokenizer': None
}
def load_model(self):
is_flg, msg = super().load_model()
if not is_flg:
return is_flg, msg
if self.model_info['device'] == 'offline':
model = None
try:
from transformers import MarianMTModel, AutoTokenizer
local_model_dir = FS.get_dir_to_local_dir(self.model_path)
tokenizer = AutoTokenizer.from_pretrained(
local_model_dir
)
model = MarianMTModel.from_pretrained(
local_model_dir
).to(we.device_id)
except Exception as e:
if model is not None:
del model
return False, f"Load model error '{e}'"
self.model_info['device'] = model.device
self.model_info['model'] = model
self.model_info['tokenizer'] = tokenizer
elif self.model_info['device'] == 'cpu':
try:
self.model_info['model'].to(we.device_id)
self.model_info['device'] = we.device_id
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
except Exception as e:
del self.model_info['model']
self.model_info['model'] = None
self.model_info['device'] = 'offline'
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
return False, f"Load model error '{e}'"
return True, ''
def unload_model(self):
super().unload_model()
if self.delete_instance:
self.model_info['device'] = 'offline'
if self.model_info['model'] is not None:
self.model_info['model'] = self.model_info['model'].to('cpu')
del self.model_info['model']
self.model_info['model'] = None
elif (isinstance(self.model_info['device'], numbers.Number)
or str(self.model_info['device']).startswith('cuda')):
self.model_info['device'] = 'cpu'
self.model_info['model'] = self.model_info['model'].to('cpu')
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
return True, ''
def get_caption(self, data):
model = self.model_info['model']
tokenizer = self.model_info['tokenizer']
batch = tokenizer(data, return_tensors="pt").to(we.device_id)
generated_ids = model.generate(**batch)
translated_text = tokenizer.batch_decode(generated_ids, skip_special_tokens=True)[0]
return translated_text
def __call__(self, *args, **kwargs):
caption = kwargs.pop('caption', None)
response = (self.get_caption(caption))
return response
class OpusMtEnZh(BaseCaptionProcessor):
def __init__(self, cfg, language='en'):
super().__init__(cfg, language=language)
self.model_path = cfg.MODEL_PATH
self.model_info = {
'device': 'offline',
'model': None,
'tokenizer': None
}
def load_model(self):
is_flg, msg = super().load_model()
if not is_flg:
return is_flg, msg
if self.model_info['device'] == 'offline':
model = None
try:
from transformers import MarianMTModel, AutoTokenizer
local_model_dir = FS.get_dir_to_local_dir(self.model_path)
tokenizer = AutoTokenizer.from_pretrained(
local_model_dir
)
model = MarianMTModel.from_pretrained(
local_model_dir
).to(we.device_id)
except Exception as e:
if model is not None:
del model
return False, f"Load model error '{e}'"
self.model_info['device'] = model.device
self.model_info['model'] = model
self.model_info['tokenizer'] = tokenizer
elif self.model_info['device'] == 'cpu':
try:
self.model_info['model'].to(we.device_id)
self.model_info['device'] = we.device_id
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
except Exception as e:
del self.model_info['model']
self.model_info['model'] = None
self.model_info['device'] = 'offline'
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
return False, f"Load model error '{e}'"
return True, ''
def unload_model(self):
super().unload_model()
if self.delete_instance:
self.model_info['device'] = 'offline'
if self.model_info['model'] is not None:
self.model_info['model'] = self.model_info['model'].to('cpu')
del self.model_info['model']
self.model_info['model'] = None
elif (isinstance(self.model_info['device'], numbers.Number)
or str(self.model_info['device']).startswith('cuda')):
self.model_info['device'] = 'cpu'
self.model_info['model'] = self.model_info['model'].to('cpu')
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
return True, ''
def get_caption(self, data):
model = self.model_info['model']
tokenizer = self.model_info['tokenizer']
batch = tokenizer(data, return_tensors="pt").to(we.device_id)
generated_ids = model.generate(**batch)
translated_text = tokenizer.batch_decode(generated_ids, skip_special_tokens=True)[0]
return translated_text
def __call__(self, *args, **kwargs):
caption = kwargs.pop('caption', None)
response = (self.get_caption(caption))
return response
@@ -0,0 +1,347 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import csv
import os
import shutil
import time
import decord
from tqdm import tqdm
import gradio as gr
from scepter.modules.utils.directory import get_md5
from scepter.modules.utils.file_system import FS
from scepter.studio.preprocess.caption_editor_ui.component_names import \
Text2VideoDataCardName
from scepter.studio.preprocess.utils.data_card import (BaseDataCard, find_prefix)
class Text2VideoDataCard(BaseDataCard):
def __init__(self,
dataset_folder,
dataset_name=None,
src_file=None,
surfix=None,
user_name='admin',
language='en'
):
super().__init__(dataset_folder,
dataset_name=dataset_name,
user_name=user_name)
self.meta['task_type'] = 'txt2vid'
self.components_name = Text2VideoDataCardName(language)
if self.new_dataset:
if surfix == '.zip':
file_list = self.load_from_zip(src_file, dataset_folder,
self.local_dataset_folder)
elif surfix in ['.txt', '.csv']:
file_list = self.load_from_list(src_file, dataset_folder,
self.local_dataset_folder)
elif surfix is None:
file_list = []
else:
raise gr.Error(
f'{self.components_name.illegal_data_err2} {surfix}')
is_flag = FS.put_dir_from_local_dir(self.local_dataset_folder,
dataset_folder,
multi_thread=True)
if not is_flag:
raise gr.Error(f'{self.components_name.illegal_data_err3}')
self.meta['cursor'] = 0 if len(file_list) > 0 else -1
self.meta['file_list'] = file_list
self.update_dataset()
def apply_changes(self):
edit_index_list = self.edit_list
for index in edit_index_list:
one_data = self.data[index]
self.data[index]['caption'] = one_data['edit_caption']
self.update_dataset()
return True, ''
def load_from_list(self, save_file, dataset_folder, local_dataset_folder):
file_list = []
videos_folder = os.path.join(local_dataset_folder, 'videos')
os.makedirs(videos_folder, exist_ok=True)
with FS.get_from(save_file) as local_path:
with open(local_path, 'r') as f:
for line in tqdm(f):
line = line.strip()
if line == '':
continue
try:
src_video_path, caption = line.split(
'#;#', 1)
except Exception:
try:
src_video_path, caption = line.split(
',', 1)
except Exception:
raise gr.Error(
self.components_name.illegal_data_err1)
relative_path = os.path.join(
'videos', f'{get_md5(src_video_path)[:18]}_{int(time.time())}.mp4')
video_path = os.path.join(dataset_folder, relative_path)
FS.get_from(src_video_path, local_path=video_path)
is_legal, new_path, prefix = find_prefix(src_video_path)
w, h, fps, duration = self.get_video_meta(video_path)
file_list.append({
'video_path':
video_path,
'relative_path':
relative_path,
'width':
w,
'height':
h,
'fps':
fps,
'duration':
duration,
'caption':
caption,
'edit_caption':
caption,
'prefix':
prefix
})
return file_list
def load_from_zip(self, save_file, data_folder, local_dataset_folder):
with FS.get_from(save_file) as local_path:
res = os.popen(
f"unzip -o '{local_path}' -d '{local_dataset_folder}'")
res = res.readlines()
if not os.path.exists(local_dataset_folder):
raise gr.Error(f'Unzip {save_file} failed {str(res)}')
file_folder = None
train_list = None
hit_dir = None
raw_list = {}
mac_osx = os.path.join(local_dataset_folder, '__MACOSX')
if os.path.exists(mac_osx):
res = os.popen(f"rm -rf '{mac_osx}'")
res = res.readlines()
for one_dir in FS.walk_dir(local_dataset_folder, recurse=False):
if one_dir.endswith('__MACOSX'):
res = os.popen(f"rm -rf '{one_dir}'")
res = res.readlines()
continue
if FS.isdir(one_dir):
if one_dir.endswith('videos') or one_dir.endswith('videos/'):
file_folder = one_dir
hit_dir = one_dir
else:
sub_dir = FS.walk_dir(one_dir)
for one_s_dir in sub_dir:
if FS.isdir(one_s_dir) and one_s_dir.split(
one_dir)[1].replace('/', '') == 'videos':
file_folder = one_s_dir
hit_dir = one_dir
if FS.isfile(one_s_dir) and one_s_dir.split(
one_dir)[1].replace('/', '') == 'train.csv':
train_list = one_s_dir
if file_folder is not None and train_list is not None:
break
elif one_dir.endswith('train.csv'):
train_list = one_dir
else:
continue
if file_folder is not None and train_list is not None:
break
if file_folder is None and len(raw_list) < 1:
raise gr.Error(
"video doesn't exist, or nothing exists in your zip")
if train_list is None:
raise gr.Error("pair list doesn't exist")
new_file_folder = f'{local_dataset_folder}/videos'
os.makedirs(new_file_folder, exist_ok=True)
if file_folder is not None:
_ = FS.get_dir_to_local_dir(file_folder, new_file_folder)
if not os.path.exists(new_file_folder):
raise gr.Error(f'{str(res)}')
new_train_list = f'{local_dataset_folder}/train.csv'
res = os.popen(f"mv '{train_list}' '{new_train_list}'")
res = res.readlines()
if not os.path.exists(new_train_list):
raise gr.Error(f'{str(res)}')
if not file_folder == hit_dir:
try:
res = os.popen(f"rm -rf '{hit_dir}/videos/*'")
_ = res.readlines()
res = os.popen(f"rm -rf '{hit_dir}'")
_ = res.readlines()
except Exception:
pass
file_list = self.load_train_file(new_train_list)
# remove unused data
for one_dir in FS.walk_dir(local_dataset_folder):
if 'videos' in one_dir or one_dir.endswith(
'file.csv') or one_dir.endswith('train.csv'):
continue
os.system(f'rm -rf {one_dir}')
return file_list
def load_train_file(self, file_path):
base_folder = os.path.dirname(file_path)
file_list = []
video_set = set()
with open(file_path, 'r') as f:
reader = csv.reader(f)
for row in reader:
if len(row) == 2:
src_video_path, prompt = row[0], row[1]
else:
return gr.Error(self.components_name.illegal_data_err2)
if src_video_path == 'Target:FILE':
continue
local_video_path = os.path.join(base_folder, src_video_path)
w, h, fps, duration = self.get_video_meta(local_video_path)
if src_video_path in video_set:
src_video_path, surfix = os.path.splitext(src_video_path)
src_video_path = f'{src_video_path}_{int(time.time() * 100)}{surfix}'
new_local_video_path = os.path.join(
base_folder, src_video_path)
self.copy_video(src_video_path, new_local_video_path)
video_set.add(src_video_path)
file_list.append({
'video_path':
local_video_path,
'relative_path':
src_video_path,
'width':
w,
'height':
h,
'fps':
fps,
"duration":
duration,
'caption':
prompt,
'edit_caption':
prompt,
'prefix':
''
})
return file_list
def write_train_file(self):
file_list = self.meta['file_list']
with open(self.local_train_file, 'w') as f:
writer = csv.writer(f)
writer.writerow(['Target:FILE', 'Prompt'])
for one_file in file_list:
relative_file = one_file['relative_path']
if relative_file.startswith('/'):
relative_file = relative_file[1:]
writer.writerow([relative_file, one_file['caption'].
strip().replace("\n", "")])
FS.put_object_from_local_file(self.local_train_file, self.train_file)
def write_data_file(self):
file_list = self.meta['file_list']
with open(self.local_save_file_list, 'w') as f:
for one_file in file_list:
f.write('{}#;#{}#;#{}#;#{}\n'.format(one_file['relative_path'],
one_file['width'],
one_file['height'],
one_file['caption'].strip().replace("\n", ""))
)
FS.put_object_from_local_file(self.local_save_file_list,
self.save_file_list)
def add_record(self, video, caption, **kwargs):
local_work_dir = self.meta['local_work_dir']
work_dir = self.meta['work_dir']
save_folder = os.path.join(local_work_dir, 'videos')
os.makedirs(save_folder, exist_ok=True)
w, h, fps, duration = self.get_video_meta(video)
relative_path = os.path.join(
'videos', f'{get_md5(video)[:18]}_{int(time.time())}.mp4')
video_path = os.path.join(work_dir, relative_path)
local_video_path = os.path.join(local_work_dir, relative_path)
self.copy_video(video, local_video_path)
self.data.append({
'video_path': video_path,
'relative_path': relative_path,
'width': w,
'height': h,
'fps': fps,
'duration': duration,
'caption': caption,
'edit_caption': caption,
'prefix': ''
})
self.set_cursor(len(self.meta['file_list']) - 1)
self.update_dataset()
return True
def copy_video(self, source_path, target_path):
if not os.path.isfile(source_path):
raise gr.Error('Video path not exist.')
try:
shutil.copy2(source_path, target_path)
except Exception as e:
raise gr.Error(str(e))
def get_video_meta(self, video):
video_reader = decord.VideoReader(video)
w = video_reader[0].shape[1]
h = video_reader[0].shape[0]
fps = video_reader.get_avg_fps()
video_length = len(video_reader)
duration = video_length / fps
return w, h, fps, duration
def delete_record(self):
if len(self) < 1:
raise gr.Error(self.components_name.delete_err1)
current_file = self.data.pop(self.cursor)
self.set_cursor(self.cursor - 1)
local_file = os.path.join(self.meta['local_work_dir'],
current_file['relative_path'])
try:
os.remove(local_file)
except Exception:
print(f'remove file {local_file} error')
if self.cursor >= len(self.meta['file_list']):
self.set_cursor(0)
if self.cursor < 0:
self.set_cursor(len(self) - 1)
if len(self.meta['file_list']) == 0:
self.set_cursor(-1)
self.update_dataset()
def export_zip(self, export_folder):
self.update_dataset()
zip_path = os.path.join(export_folder, f'{self.dataset_name}.zip')
local_zip, _ = FS.map_to_local(zip_path)
os.makedirs(os.path.dirname(local_zip), exist_ok=True)
res = os.popen(
f"cd '{self.local_work_dir}' && mkdir -p '{self.dataset_name}' "
f"&& cp -rf videos '{self.dataset_name}/videos' "
f"&& cp -rf train.csv '{self.dataset_name}/train.csv' "
f"&& zip -r '{os.path.abspath(local_zip)}' '{self.dataset_name}'/* "
f"&& rm -rf '{self.dataset_name}'")
print(res.readlines())
FS.put_object_from_local_file(local_zip, zip_path)
if not FS.exists(zip_path):
raise gr.Error(self.components_name.export_zip_err1)
return local_zip
+15 -3
View File
@@ -44,18 +44,22 @@ def kill_job(pid):
class Trainer():
def __init__(self, run_script, status_message):
def __init__(self, run_script, status_message, visible_gpus):
self.run_script = run_script
self.status_message = status_message
self.visible_gpus = visible_gpus
self.proc = None
def __call__(self, task_name):
torch.cuda.empty_cache()
error_folder = './error_logs'
os.makedirs(error_folder, exist_ok=True)
self.status_message.error_log = f'{error_folder}/{int(time.time())}.log'
cmd = f'PYTHONPATH=. python {self.run_script} ' \
f'--cfg={task_name}/train.yaml 2> {self.status_message.error_log}'
if self.visible_gpus is not None and len(self.visible_gpus) > 0:
cmd = f'CUDA_VISIBLE_DEVICES={",".join([str(i) for i in self.visible_gpus])} ' + cmd
# cmd = [f"python {self.run_script}"]
print(cmd)
try:
@@ -87,6 +91,7 @@ class TrainManager():
self.runing_tasks = {}
self.run_script = run_script
self.work_dir = work_dir
self.visible_gpus = list(range(torch.cuda.device_count()))
def task_dispatch():
while True:
@@ -143,7 +148,7 @@ class TrainManager():
task_name = self.task_queue.pop(0)
print(f'start task {task_name}')
status_message = TaskStatus()
train_ins = Trainer(self.run_script, status_message)
train_ins = Trainer(self.run_script, status_message, self.visible_gpus)
train_thread = threading.Thread(target=train_ins,
args=(os.path.join(
self.work_dir,
@@ -174,11 +179,18 @@ class TrainManager():
self.task_manage = threading.Thread(target=task_dispatch, daemon=True)
self.task_manage.start()
def set_gpus(self, gpus=None):
if gpus is None:
self.visible_gpus = list(range(torch.cuda.device_count()))
else:
self.visible_gpus = list(gpus)
def check_memory(self):
# Check Cuda Memory
visible_gpus = self.visible_gpus
mem_msg = ''
if torch.cuda.is_available():
for device_id in range(torch.cuda.device_count()):
for device_id in visible_gpus:
free_mem, total_mem = torch.cuda.mem_get_info(device_id)
free_mem = free_mem / (1024**3)
total_mem = total_mem / (1024**3)
@@ -74,12 +74,13 @@ class TrainerUIName():
self.task_choices = ['Text2Image', 'Image Editing']
self.data_task_map = {
'scepter_txt2img': None,
'scepter_img2img': 'edit'
'scepter_img2img': 'edit',
'scepter_txt2vid': 'dit'
}
if language == 'en':
self.user_direction = '''
### User Guide
- Data: Data preparation is done through the Data Manager. (zip example: [3D](https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=datasets/3D_example_csv.zip))
- Data: Data preparation is done through the Data Manager. (zip example: [3D](https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=datasets/3D_example_csv.zip), txt example(Support only video data): [txt](https://modelscope.cn/models/iic/scepter/resolve/master/datasets/video_example.txt))
- Parameters: You can try modifying the related parameters.
- Training: Click on [Start Training].
- Testing: After completing the training, click [Go to inference].
@@ -93,14 +94,20 @@ class TrainerUIName():
self.data_source_choices = [
'Dataset zip', 'MaaS Dataset', 'Dataset Management'
]
self.illegal_data_err = (
'The list supports only "," or "#;#" as delimiters. '
'The two columns represent video path and description, '
'respectively.')
self.data_source_value = 'Dataset zip'
self.data_source_name = 'Data Source'
self.data_type_map = {
'scepter_txt2img': 'Text2Image Generation',
'scepter_img2img': 'Image Edit Generation'
'scepter_img2img': 'Image Edit Generation',
'scepter_txt2vid': 'Text2Video Generation'
}
self.data_type_choices = list(self.data_type_map.keys())
self.data_type_value = 'scepter_txt2img'
self.data_type_value_video = 'scepter_txt2vid'
self.data_type_name = 'Data Type'
self.ori_data_name = 'Data Name'
# Supports MaaS dataset/local/HTTP Zip package
@@ -120,8 +127,8 @@ class TrainerUIName():
self.base_model = 'Base Model'
self.tuner_name = 'Tuner Method'
self.base_model_revision = 'Model Version Number'
self.resolution_height = 'Train Image Height'
self.resolution_width = 'Train Image Width'
self.resolution_height = 'Train Image or Video Height'
self.resolution_width = 'Train Image or Video Width'
self.resolution_height_max = 'Resolution Height Max'
self.resolution_width_max = 'Resolution Width Max'
self.train_epoch = 'Total Training Epochs'
@@ -142,6 +149,8 @@ class TrainerUIName():
self.bucket_resolution_steps = 'Bucket Resolution Steps'
self.bucket_no_upscale = 'Bucket No Upscale'
self.bucket_no_upscale_ins = 'Disable Automatic Image Upscaling'
self.accumulate_step = 'Accumulate Step'
self.gpus = 'Select GPUs'
# Error or Warning
self.training_err1 = 'CUDA is unavailable.'
self.training_err2 = 'Currently insufficient VRAM, training failed!'
@@ -153,7 +162,7 @@ class TrainerUIName():
elif language == 'zh':
self.user_direction = '''
### 使用说明
- 数据: 通过数据管理器进行数据的准备(ZIP样例:[3D](https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=datasets/3D_example_csv.zip))
- 数据: 通过数据管理器进行数据的准备(ZIP样例:[3D](https://www.modelscope.cn/api/v1/models/iic/scepter/repo?Revision=master&FilePath=datasets/3D_example_csv.zip),txt样例(仅支持视频数据):[txt](https://modelscope.cn/models/iic/scepter/resolve/master/datasets/video_example.txt))
- 参数: 可尝试进行相关参数的修改
- 训练: 点击【开始训练】
- 测试: 完成训练后点击【使用模型】
@@ -161,14 +170,17 @@ class TrainerUIName():
- 对于大规模数据的处理和训练,建议使用命令行形式
''' # noqa
self.data_source_choices = ['数据集zip', 'MaaS数据集', '数据管理器']
self.illegal_data_err = '列表只支持,或#;#作为分割符,两列分别为视频路径/描述'
self.data_source_value = '数据集zip'
self.data_source_name = '数据集来源'
self.data_type_map = {
'scepter_txt2img': '文生图数据',
'scepter_img2img': '图像编辑(图生图)数据'
'scepter_img2img': '图像编辑(图生图)数据',
'scepter_txt2vid': '文生视频数据'
}
self.data_type_choices = list(self.data_type_map.keys())
self.data_type_value = 'scepter_txt2img'
self.data_type_value_video = 'scepter_txt2vid'
self.data_type_name = '数据类型'
self.ori_data_name = '数据集名称'
self.ms_data_name_place_hold = '请使用数据管理器导入' # '支持MaaS数据集/本地/Http Zip包'
@@ -187,8 +199,8 @@ class TrainerUIName():
self.base_model = '基础模型'
self.tuner_name = '微调方法'
self.base_model_revision = '模型版本号'
self.resolution_height = '训练图片高度'
self.resolution_width = '训练图片宽度'
self.resolution_height = '训练图片或视频高度'
self.resolution_width = '训练图片或视频宽度'
self.resolution_height_max = '最大训练高度'
self.resolution_width_max = '最大训练宽度'
self.train_epoch = '总训练轮数'
@@ -208,6 +220,8 @@ class TrainerUIName():
self.bucket_resolution_steps = '分桶分辨率步长'
self.bucket_no_upscale = '分桶分辨率不做放大'
self.bucket_no_upscale_ins = '禁止图片分辨率上采样'
self.accumulate_step = '梯度累积数量'
self.gpus = '选择GPU'
# Error or Warning
self.training_err1 = 'CUDA不可用.'
self.training_err2 = '目前显存不足,训练失败!'
@@ -59,13 +59,16 @@ class ModelUI(UIBase):
status_file = os.path.join(self.work_dir, one_dir,
'status.json')
if FS.exists(status_file):
status = json.load(open(status_file, 'r'))
try:
status = json.load(open(status_file, 'r'))
except:
continue
status['model_name'] = one_dir
have_model_list.append(status)
have_model_list.sort(key=lambda x: x['start_time'])
self.user_level_model_list[user_name] = [
v['model_name'] for v in have_model_list
][:100]
]
def get_ckpt_list(self, output_model):
all_ckpt_list = []
@@ -177,11 +180,12 @@ class ModelUI(UIBase):
def model_name_change(model_name):
if model_name is None:
return '', gr.Column(), '', []
return '', gr.Column(), None, []
message = trainer_ui.trainer_ins.get_log(model_name)
status = trainer_ui.trainer_ins.get_status(model_name)
ckpt_list = self.get_ckpt_list(model_name)
ckpt_value = ckpt_list[-1] if len(ckpt_list) > 0 else ''
ckpt_value = ckpt_list[-1] if len(ckpt_list) > 0 else None
ckpt_list = ckpt_list if isinstance(ckpt_list, list) and len(ckpt_list) > 0 else None
if ckpt_value is not None and len(ckpt_value) > 0:
gallery_value = self.get_gallery_list(model_name, ckpt_value)
else:
@@ -264,13 +268,14 @@ class ModelUI(UIBase):
message = trainer_ui.trainer_ins.get_log(model_name)
status = trainer_ui.trainer_ins.get_status(model_name)
ckpt_list = self.get_ckpt_list(model_name)
ckpt_value = ckpt_list[-1] if len(ckpt_list) > 0 else ''
ckpt_value = ckpt_list[-1] if len(ckpt_list) > 0 else None
ckpt_list = ckpt_list if isinstance(ckpt_list, list) and len(ckpt_list) > 0 else None
ret_gallery = ckpt_name_change(model_name, ckpt_value)
model_list = self.user_level_model_list.get(login_user_name, [])
self.load_history(login_user_name)
return (message, gr.Column(visible=status in ('running',
'success')),
gr.Dropdown(choices=self.user_level_model_list.get(
login_user_name, []),
gr.Dropdown(choices=model_list,
value=model_name),
gr.Dropdown(choices=ckpt_list,
value=ckpt_value), ret_gallery)
@@ -330,7 +335,8 @@ class ModelUI(UIBase):
message = trainer_ui.trainer_ins.get_log(model_name)
status = trainer_ui.trainer_ins.get_status(model_name)
ckpt_list = self.get_ckpt_list(model_name)
ckpt_value = ckpt_list[-1] if len(ckpt_list) > 0 else ''
ckpt_value = ckpt_list[-1] if len(ckpt_list) > 0 else None
ckpt_list = ckpt_list if isinstance(ckpt_list, list) and len(ckpt_list) > 0 else None
ret_gallery = ckpt_name_change(model_name, ckpt_value)
self.load_history(login_user_name)
return (message, gr.Column(visible=status in ('running',
@@ -553,7 +559,7 @@ class ModelUI(UIBase):
if len(model_list) > 0:
model_name = model_list[-1]
else:
model_name = ''
model_name = None
return gr.Dropdown(choices=model_list, value=model_name)
manager.user_name.change(model_name_change,
@@ -5,12 +5,17 @@ import datetime
import json
import os
import random
import time
from collections import OrderedDict
from tqdm import tqdm
import decord
import gradio as gr
import scepter
import torch
import yaml
from scepter.modules.utils.directory import get_md5
from scepter.modules.utils.file_system import FS
from scepter.studio.self_train.scripts.trainer import TrainManager
from scepter.studio.self_train.self_train_ui.component_names import \
@@ -84,6 +89,7 @@ class TrainerUI(UIBase):
self.current_train_model = None
self.trainer_ins = TrainManager(self.run_script, self.work_dir_pre)
self.component_names = TrainerUIName(language=language)
self.save_file_local_path = cfg.SAVE_FILE_LOCAL_PATH
self.h_level_dict = {}
for hw_tuple in self.train_para_data.RESOLUTIONS.get('VALUES', []):
@@ -107,7 +113,7 @@ class TrainerUI(UIBase):
choices=self.component_names.data_source_choices,
value=self.component_names.data_source_value,
label=self.component_names.data_source_name,
interactive=False)
interactive=True)
self.data_type = gr.Dropdown(
choices=[
self.component_names.data_type_map[key] for key
@@ -116,20 +122,20 @@ class TrainerUI(UIBase):
value=self.component_names.data_type_map[
self.component_names.data_type_value],
label=self.component_names.data_type_name,
interactive=False)
interactive=True)
self.ori_data_name = gr.Textbox(
label=self.component_names.ori_data_name,
max_lines=1,
placeholder=self.component_names.ori_data_name,
interactive=False)
interactive=True)
self.ms_data_name = gr.Textbox(
label=' or '.join(
self.component_names.data_source_choices),
max_lines=1,
placeholder=self.component_names.
ms_data_name_place_hold,
visible=False,
interactive=False)
visible=True,
interactive=True)
with gr.Group(visible=False) as self.ms_data_box:
with gr.Row():
self.ms_data_space = gr.Textbox(
@@ -199,8 +205,7 @@ class TrainerUI(UIBase):
visible=lora_visible) as self.lora_param:
self.lora_alpha = gr.Number(
label='LoRA Alpha',
value=self.para_data.get(
'lora_alpha', 4),
value=self.para_data.get('lora_alpha', 4),
interactive=True)
self.lora_rank = gr.Number(
label='LoRA Rank',
@@ -310,6 +315,19 @@ class TrainerUI(UIBase):
precision=0,
interactive=True)
with gr.Row():
self.accumulate_step = gr.Number(
label=self.component_names.accumulate_step,
value=self.para_data.get('ACCUMULATE_STEP', 1),
precision=0,
interactive=True)
self.gpus = gr.Dropdown(
choices=list(range(torch.cuda.device_count())),
value=list(range(torch.cuda.device_count())),
label=self.component_names.gpus,
multiselect=True,
interactive=True)
with gr.Row():
self.prompt_prefix = gr.Text(
label=self.component_names.prompt_prefix,
@@ -355,6 +373,15 @@ class TrainerUI(UIBase):
self.component_names.data_type_value], 'damo',
'style_custom_dataset', 'style_custom_dataset',
'3D'
],
[
self.component_names.data_source_choices[2],
self.component_names.data_type_map[
self.component_names.data_type_value_video],
'',
'https://modelscope.cn/models/iic/scepter/resolve/master/datasets/video_example.txt', # noqa
'video_example_txt',
''
]
],
inputs=[
@@ -439,6 +466,7 @@ class TrainerUI(UIBase):
eval_prompts = [] if is_edit else self.train_para_data.get(
'EVAL_PROMPTS', [])
eval_prompts = ret_data.get('EVAL_PROMPTS', eval_prompts)
return ret_data.get('EPOCHS', 10), \
ret_data.get('LEARNING_RATE', 0.0001), \
ret_data.get('SAVE_INTERVAL', 10), \
@@ -613,7 +641,8 @@ class TrainerUI(UIBase):
lora_alpha, lora_rank, text_lora_alpha, text_lora_rank,
sce_ratio, enable_resolution_bucket,
min_bucket_resolution, max_bucket_resolution,
bucket_resolution_steps, bucket_no_upscale, user_name):
bucket_resolution_steps, bucket_no_upscale,
accumulate_step, gpus, user_name):
# Check Cuda
if not torch.cuda.is_available() and not self.is_debug:
raise gr.Error(self.component_names.training_err1)
@@ -621,7 +650,7 @@ class TrainerUI(UIBase):
if work_name == 'custom' or work_name is None or work_name == '':
raise gr.Error(self.component_names.training_err4)
work_dir = os.path.join(self.work_dir_pre, work_name)
login_user_name = user_name
self.current_train_model = work_name
if os.path.exists(work_dir) or os.path.exists(
f'.flag/{work_name}.tmp'):
@@ -655,7 +684,7 @@ class TrainerUI(UIBase):
if ms_data_name is None:
raise gr.Error(self.component_names.training_err3)
def prepare_train_data(data_cfg):
def prepare_train_image_data(data_cfg):
data_cfg['BATCH_SIZE'] = int(train_batch_size)
data_cfg['PROMPT_PREFIX'] = prompt_prefix
data_cfg['REPLACE_KEYWORDS'] = replace_keywords
@@ -737,8 +766,12 @@ class TrainerUI(UIBase):
if os.path.exists(local_data_dir) and os.path.exists(
local_file_list):
data_cfg.update({
'NAME': 'ImageTextPairDataset' if data_cfg['NAME'] == 'ImageTextPairMSDataset' else data_cfg['NAME'],
'ENABLE_RESOLUTION_BUCKET': enable_resolution_bucket,
'NAME':
'ImageTextPairDataset'
if data_cfg['NAME'] == 'ImageTextPairMSDataset'
else data_cfg['NAME'],
'ENABLE_RESOLUTION_BUCKET':
enable_resolution_bucket,
'SAMPLER': {
'NAME':
'ResolutionBatchSampler',
@@ -761,7 +794,8 @@ class TrainerUI(UIBase):
'BUCKET_NO_UPSCALE':
bucket_no_upscale
},
'DATA_NUM': data_num
'DATA_NUM':
data_num
})
if 'TRANSFORMS' in data_cfg:
for trans in data_cfg['TRANSFORMS']:
@@ -776,6 +810,69 @@ class TrainerUI(UIBase):
return data_cfg
def prepare_train_video_data(data_cfg):
if ms_data_name.startswith('http') and (
'.txt' in ms_data_name or '.csv' in ms_data_name):
data_name = get_data_from_list()
else:
data_name = os.path.join(ms_data_name, 'file.txt')
data_cfg['BATCH_SIZE'] = int(train_batch_size)
data_cfg['PROMPT_PREFIX'] = prompt_prefix
if data_cfg['NAME'] in ['VideoGenDataset']:
data_cfg['SAMPLER']['SUB_SAMPLERS'][0][
'PATH_PREFIX'] = os.path.dirname(data_name)
data_cfg['SAMPLER']['SUB_SAMPLERS'][0][
'INDEX_FILE'] = data_name
elif data_cfg['NAME'] in ['VideoGenDatasetOTF']:
data_cfg['PATH_PREFIX'] = os.path.dirname(data_name)
data_cfg['DATA_FILE'] = data_name
else:
raise Exception('Unsupported data type {}'.format(
data_cfg['NAME']))
return data_cfg
def get_data_from_list():
file_list = []
file = FS.get_from(ms_data_name)
with FS.get_from(file) as local_path:
with open(local_path, 'r') as f:
for line in tqdm(f):
line = line.strip()
if line == '':
continue
try:
src_video_path, caption = line.split('#;#', 1)
except Exception:
try:
src_video_path, caption = line.split(
',', 1)
except Exception:
raise gr.Error(
self.component_names.illegal_data_err)
relative_path = os.path.join(
'videos',
f'{get_md5(src_video_path)[:18]}_{int(time.time())}.mp4'
)
video_path = os.path.join(
self.save_file_local_path, ori_data_name,
relative_path)
local_path = FS.get_from(src_video_path,
local_path=video_path)
video_reader = decord.VideoReader(local_path)
w = video_reader[0].shape[1]
h = video_reader[0].shape[0]
file_list.append('{}#;#{}#;#{}#;#{}\n'.format(
relative_path, w, h, caption))
local_save_file_list = os.path.join(self.save_file_local_path,
ori_data_name, 'file.txt')
directory = os.path.dirname(local_save_file_list)
os.makedirs(directory, exist_ok=True)
FS.delete_object(file)
with open(local_save_file_list, 'w') as f:
f.writelines(file_list)
return local_save_file_list
def prepare_eval_data(data_cfg):
data_cfg['PROMPT_PREFIX'] = prompt_prefix
data_cfg['IMAGE_SIZE'] = [
@@ -850,8 +947,14 @@ class TrainerUI(UIBase):
]
cfg['SOLVER']['TUNER'] = tuner_cfg_list
cfg['SOLVER']['TRAIN_DATA'] = prepare_train_data(
cfg['SOLVER']['TRAIN_DATA'])
if cfg['SOLVER']['TRAIN_DATA']['NAME'] in [
'VideoGenDataset', 'VideoGenDatasetOTF'
]:
cfg['SOLVER']['TRAIN_DATA'] = prepare_train_video_data(
cfg['SOLVER']['TRAIN_DATA'])
else:
cfg['SOLVER']['TRAIN_DATA'] = prepare_train_image_data(
cfg['SOLVER']['TRAIN_DATA'])
if eval_prompts is not None and len(eval_prompts) > 0:
cfg['SOLVER']['EVAL_DATA'] = prepare_eval_data(
cfg['SOLVER']['EVAL_DATA'])
@@ -868,6 +971,8 @@ class TrainerUI(UIBase):
hook['INTERVAL'] = save_interval
hook['PUSH_TO_HUB'] = push_to_hub
hook['HUB_MODEL_ID'] = hub_model_id
if hook['NAME'] == 'BackwardHook':
hook['ACCUMULATE_STEP'] = int(accumulate_step)
if 'EVAL_HOOKS' in cfg['SOLVER']:
for hook in cfg['SOLVER']['EVAL_HOOKS']:
if hook['NAME'] == 'ProbeDataHook':
@@ -881,6 +986,7 @@ class TrainerUI(UIBase):
default_flow_style=False)
return cfg_file
self.trainer_ins.set_gpus(gpus)
before_kill_inference = self.trainer_ins.check_memory()
if hasattr(manager, 'inference'):
for k, v in manager.inference.pipe_manager.pipeline_level_modules.items(
@@ -892,7 +998,6 @@ class TrainerUI(UIBase):
'dynamic_unload')):
manager.preprocess.dataset_gallery.processors_manager.dynamic_unload(
)
after_kill_inference = self.trainer_ins.check_memory()
message = f'GPU info: {before_kill_inference}. \n\n'
message += f'After unloading inference models, the GPU info: {after_kill_inference}. \n\n'
@@ -923,7 +1028,8 @@ class TrainerUI(UIBase):
self.text_lora_alpha, self.text_lora_rank, self.sce_ratio,
self.enable_resolution_bucket, self.min_bucket_resolution,
self.max_bucket_resolution, self.bucket_resolution_steps,
self.bucket_no_upscale, manager.user_name
self.bucket_no_upscale, self.accumulate_step, self.gpus,
manager.user_name
],
outputs=[inference_ui.output_model_name],
queue=True)
+17 -7
View File
@@ -65,6 +65,13 @@ if __name__ == '__main__':
choices=['en', 'zh'],
default='en',
help='Now we only support english(en) and chinese(zh)')
parser.add_argument('--tab',
dest='tab',
choices=['all', 'chatbot'],
default='all',
help='The tabs will be launched, '
'set [all] to use all tools and set [chatbot] to use chatbot only.')
args = parser.parse_args()
if not os.path.exists(args.config):
print(
@@ -88,7 +95,7 @@ if __name__ == '__main__':
if not FS.exists(info['CONFIG']):
raise f"{info['CONFIG']} doesn't exist."
interface = None
if ifid == 'home':
if ifid == 'home' and args.tab in ["all", ifid]:
from scepter.studio.home.home import HomeUI
interface = HomeUI(info['CONFIG'],
@@ -96,7 +103,7 @@ if __name__ == '__main__':
language=args.language,
root_work_dir=config.WORK_DIR)
print('init home page success!')
if ifid == 'preprocess':
if ifid == 'preprocess' and args.tab in ["all", ifid]:
from scepter.studio.preprocess.preprocess import PreprocessUI
interface = PreprocessUI(info['CONFIG'],
@@ -104,7 +111,7 @@ if __name__ == '__main__':
language=args.language,
root_work_dir=config.WORK_DIR)
print('init preprocess success!')
if ifid == 'self_train':
if ifid == 'self_train' and args.tab in ["all", ifid]:
from scepter.studio.self_train.self_train import SelfTrainUI
interface = SelfTrainUI(info['CONFIG'],
@@ -112,14 +119,14 @@ if __name__ == '__main__':
language=args.language,
root_work_dir=config.WORK_DIR)
print('init self-train success!')
if ifid == 'tuner_manager':
if ifid == 'tuner_manager' and args.tab in ["all", ifid]:
from scepter.studio.tuner_manager.tuner_manager import TunerManagerUI
interface = TunerManagerUI(info['CONFIG'],
is_debug=args.debug,
language=args.language,
root_work_dir=config.WORK_DIR)
print('init tuner-manager success!')
if ifid == 'inference':
if ifid == 'inference' and args.tab in ["all", ifid]:
from scepter.studio.inference.inference import InferenceUI
interface = InferenceUI(info['CONFIG'],
@@ -127,7 +134,7 @@ if __name__ == '__main__':
language=args.language,
root_work_dir=config.WORK_DIR)
print('init inference success!')
if ifid == 'ChatBot':
if ifid == 'chatbot' and args.tab in ["all", ifid]:
from scepter.studio.chatbot.chatbot import ChatBotUI
interface = ChatBotUI(info['CONFIG'],
@@ -177,10 +184,13 @@ if __name__ == '__main__':
if len(auth_info) > 0:
demo.load(init_value, outputs=[tab_manager.user_name])
allowed_paths = [config['WORK_DIR']]
allowed_paths.extend(list(set([fs_cfg['TEMP_DIR'] for fs_cfg in config['FILE_SYSTEM']])) if 'FILE_SYSTEM' in config else [])
demo.queue(status_update_rate=1).launch(
server_name=args.host if args.host else config['HOST'],
server_port=int(args.port) if args.port else config['PORT'],
root_path=config['ROOT'],
show_error=True,
debug=True,
auth=check_auth if len(auth_info) > 0 else None)
auth=check_auth if len(auth_info) > 0 else None,
allowed_paths=allowed_paths)
+1 -1
View File
@@ -1,7 +1,7 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
__version__ = '1.2.0'
__version__ = '1.3.0'
version_info = tuple(int(x) for x in __version__.split('.')[0:3])
@@ -0,0 +1,296 @@
NAME: ACE_0.6B_1024
IS_DEFAULT: False
DEFAULT_PARAS:
PARAS:
#
INPUT:
INPUT_IMAGE:
INPUT_MASK:
TASK:
PROMPT: ""
NEGATIVE_PROMPT: ""
OUTPUT_HEIGHT: 1024
OUTPUT_WIDTH: 1024
SAMPLER: ddim
SAMPLE_STEPS: 20
GUIDE_SCALE: 4.5
GUIDE_RESCALE: 0.5
SEED: -1
TAR_INDEX: 0
OUTPUT:
LATENT:
IMAGES:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: float16
INPUT: ["IMAGE"]
- NAME: decode
DTYPE: float16
INPUT: ["LATENT"]
#
DIFFUSION_MODEL:
FUNCTION:
- NAME: forward
DTYPE: float16
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE"]
#
COND_STAGE_MODEL:
FUNCTION:
- NAME: encode_list_of_list
DTYPE: bfloat16
INPUT: ["PROMPT"]
#
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
USE_TEXT_POS_EMBEDDINGS: True
#
MODEL:
NAME: LatentDiffusionACE
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
DECODER_BIAS: 0.5
DEFAULT_N_PROMPT: ""
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
USE_TEXT_POS_EMBEDDINGS: True
#
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: eps
MIN_SNR_GAMMA:
NOISE_SCHEDULER:
NAME: LinearScheduler
NUM_TIMESTEPS: 1000
BETA_MIN: 0.0001
BETA_MAX: 0.02
#
DIFFUSION_MODEL:
NAME: ACE
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/dit/ace_0.6b_1024px.pth
IGNORE_KEYS: [ ]
PATCH_SIZE: 2
IN_CHANNELS: 4
HIDDEN_SIZE: 1152
DEPTH: 28
NUM_HEADS: 16
MLP_RATIO: 4.0
PRED_SIGMA: True
DROP_PATH: 0.0
WINDOW_DIZE: 0
Y_CHANNELS: 4096
MAX_SEQ_LEN: 1024
QK_NORM: True
USE_GRAD_CHECKPOINT: True
ATTENTION_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/vae/vae.bin
IGNORE_KEYS: []
#
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: 4
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: 4
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/text_encoder/t5-v1_1-xxl/
TOKENIZER_PATH: ms://iic/ACE-0.6B-1024px@models/tokenizer/t5-v1_1-xxl
LENGTH: 120
T5_DTYPE: bfloat16
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
CLEAN: whitespace
USE_GRAD: False
#
MODEL_LOCAL:
NAME: LatentDiffusionACE
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
DECODER_BIAS: 0.5
DEFAULT_N_PROMPT: ""
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
USE_TEXT_POS_EMBEDDINGS: True
#
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: eps
MIN_SNR_GAMMA:
NOISE_SCHEDULER:
NAME: LinearScheduler
NUM_TIMESTEPS: 1000
BETA_MIN: 0.0001
BETA_MAX: 0.02
#
DIFFUSION_MODEL:
NAME: ACE
PRETRAINED_MODEL: models/scepter/ACE-0.6B-1024px/models/dit/ace_0.6b_1024px.pth
IGNORE_KEYS: [ ]
PATCH_SIZE: 2
IN_CHANNELS: 4
HIDDEN_SIZE: 1152
DEPTH: 28
NUM_HEADS: 16
MLP_RATIO: 4.0
PRED_SIGMA: True
DROP_PATH: 0.0
WINDOW_DIZE: 0
Y_CHANNELS: 4096
MAX_SEQ_LEN: 1024
QK_NORM: True
USE_GRAD_CHECKPOINT: True
ATTENTION_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL: models/scepter/ACE-0.6B-1024px/models/vae/vae.bin
IGNORE_KEYS: []
#
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: 4
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: 4
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: models/scepter/ACE-0.6B-1024px/models/text_encoder/t5-v1_1-xxl/
TOKENIZER_PATH: models/scepter/ACE-0.6B-1024px/models/tokenizer/t5-v1_1-xxl
LENGTH: 120
T5_DTYPE: bfloat16
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
CLEAN: whitespace
USE_GRAD: False
#
MODEL_HF:
NAME: LatentDiffusionACE
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
DECODER_BIAS: 0.5
DEFAULT_N_PROMPT: ""
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
USE_TEXT_POS_EMBEDDINGS: True
#
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: eps
MIN_SNR_GAMMA:
NOISE_SCHEDULER:
NAME: LinearScheduler
NUM_TIMESTEPS: 1000
BETA_MIN: 0.0001
BETA_MAX: 0.02
#
DIFFUSION_MODEL:
NAME: ACE
PRETRAINED_MODEL: hf://scepter-studio/ACE-0.6B-1024px@models/dit/ace_0.6b_1024px.pth
IGNORE_KEYS: [ ]
PATCH_SIZE: 2
IN_CHANNELS: 4
HIDDEN_SIZE: 1152
DEPTH: 28
NUM_HEADS: 16
MLP_RATIO: 4.0
PRED_SIGMA: True
DROP_PATH: 0.0
WINDOW_DIZE: 0
Y_CHANNELS: 4096
MAX_SEQ_LEN: 1024
QK_NORM: True
USE_GRAD_CHECKPOINT: True
ATTENTION_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL: hf://scepter-studio/ACE-0.6B-1024px@models/vae/vae.bin
IGNORE_KEYS: []
#
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: 4
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: 4
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: hf://scepter-studio/ACE-0.6B-1024px@models/text_encoder/t5-v1_1-xxl/
TOKENIZER_PATH: hf://scepter-studio/ACE-0.6B-1024px@models/tokenizer/t5-v1_1-xxl
LENGTH: 120
T5_DTYPE: bfloat16
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
CLEAN: whitespace
USE_GRAD: False
@@ -0,0 +1,729 @@
NAME: ACE_0.6B_1024_REFINER
IS_DEFAULT: False
DEFAULT_PARAS:
PARAS:
#
INPUT:
INPUT_IMAGE:
INPUT_MASK:
TASK:
PROMPT: ""
NEGATIVE_PROMPT: ""
OUTPUT_HEIGHT: 1024
OUTPUT_WIDTH: 1024
SAMPLER: ddim
SAMPLE_STEPS: 50
GUIDE_SCALE: 4.5
GUIDE_RESCALE: 0.5
SEED: -1
TAR_INDEX: 0
OUTPUT:
LATENT:
IMAGES:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: float16
INPUT: ["IMAGE"]
- NAME: decode
DTYPE: float16
INPUT: ["LATENT"]
#
DIFFUSION_MODEL:
FUNCTION:
- NAME: forward
DTYPE: float16
INPUT: ["SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE"]
#
COND_STAGE_MODEL:
FUNCTION:
- NAME: encode_list_of_list
DTYPE: bfloat16
INPUT: ["PROMPT"]
#
MODEL:
NAME: LatentDiffusionACE
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
DECODER_BIAS: 0.5
DEFAULT_N_PROMPT: ""
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
USE_TEXT_POS_EMBEDDINGS: True
#
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: eps
MIN_SNR_GAMMA:
NOISE_SCHEDULER:
NAME: LinearScheduler
NUM_TIMESTEPS: 1000
BETA_MIN: 0.0001
BETA_MAX: 0.02
#
DIFFUSION_MODEL:
NAME: ACE
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/dit/ace_0.6b_1024px.pth
IGNORE_KEYS: [ ]
PATCH_SIZE: 2
IN_CHANNELS: 4
HIDDEN_SIZE: 1152
DEPTH: 28
NUM_HEADS: 16
MLP_RATIO: 4.0
PRED_SIGMA: True
DROP_PATH: 0.0
WINDOW_DIZE: 0
Y_CHANNELS: 4096
MAX_SEQ_LEN: 1024
QK_NORM: True
USE_GRAD_CHECKPOINT: True
ATTENTION_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/vae/vae.bin
IGNORE_KEYS: []
#
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: 4
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: 4
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: ms://iic/ACE-0.6B-1024px@models/text_encoder/t5-v1_1-xxl/
TOKENIZER_PATH: ms://iic/ACE-0.6B-1024px@models/tokenizer/t5-v1_1-xxl
LENGTH: 120
T5_DTYPE: bfloat16
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
CLEAN: whitespace
USE_GRAD: False
#
MODEL_LOCAL:
NAME: LatentDiffusionACE
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
DECODER_BIAS: 0.5
DEFAULT_N_PROMPT: ""
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
USE_TEXT_POS_EMBEDDINGS: True
#
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: eps
MIN_SNR_GAMMA:
NOISE_SCHEDULER:
NAME: LinearScheduler
NUM_TIMESTEPS: 1000
BETA_MIN: 0.0001
BETA_MAX: 0.02
#
DIFFUSION_MODEL:
NAME: ACE
PRETRAINED_MODEL: models/scepter/ACE-0.6B-1024px/models/dit/ace_0.6b_1024px.pth
IGNORE_KEYS: [ ]
PATCH_SIZE: 2
IN_CHANNELS: 4
HIDDEN_SIZE: 1152
DEPTH: 28
NUM_HEADS: 16
MLP_RATIO: 4.0
PRED_SIGMA: True
DROP_PATH: 0.0
WINDOW_DIZE: 0
Y_CHANNELS: 4096
MAX_SEQ_LEN: 1024
QK_NORM: True
USE_GRAD_CHECKPOINT: True
ATTENTION_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL: models/scepter/ACE-0.6B-1024px/models/vae/vae.bin
IGNORE_KEYS: []
#
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: 4
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: 4
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: models/scepter/ACE-0.6B-1024px/models/text_encoder/t5-v1_1-xxl/
TOKENIZER_PATH: models/scepter/ACE-0.6B-1024px/models/tokenizer/t5-v1_1-xxl
LENGTH: 120
T5_DTYPE: bfloat16
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
CLEAN: whitespace
USE_GRAD: False
#
MODEL_HF:
NAME: LatentDiffusionACE
PRETRAINED_MODEL:
IGNORE_KEYS: [ ]
SCALE_FACTOR: 0.18215
SIZE_FACTOR: 8
DECODER_BIAS: 0.5
DEFAULT_N_PROMPT: ""
TEXT_IDENTIFIER: [ '{image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
USE_TEXT_POS_EMBEDDINGS: True
#
DIFFUSION:
NAME: BaseDiffusion
PREDICTION_TYPE: eps
MIN_SNR_GAMMA:
NOISE_SCHEDULER:
NAME: LinearScheduler
NUM_TIMESTEPS: 1000
BETA_MIN: 0.0001
BETA_MAX: 0.02
#
DIFFUSION_MODEL:
NAME: ACE
PRETRAINED_MODEL: hf://scepter-studio/ACE-0.6B-1024px@models/dit/ace_0.6b_1024px.pth
IGNORE_KEYS: [ ]
PATCH_SIZE: 2
IN_CHANNELS: 4
HIDDEN_SIZE: 1152
DEPTH: 28
NUM_HEADS: 16
MLP_RATIO: 4.0
PRED_SIGMA: True
DROP_PATH: 0.0
WINDOW_DIZE: 0
Y_CHANNELS: 4096
MAX_SEQ_LEN: 1024
QK_NORM: True
USE_GRAD_CHECKPOINT: True
ATTENTION_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKL
EMBED_DIM: 4
PRETRAINED_MODEL: hf://scepter-studio/ACE-0.6B-1024px@models/vae/vae.bin
IGNORE_KEYS: []
#
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: 4
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: 4
DROPOUT: 0.0
RESAMP_WITH_CONV: True
GIVE_PRE_END: False
TANH_OUT: False
#
COND_STAGE_MODEL:
NAME: T5EmbedderHF
PRETRAINED_MODEL: hf://scepter-studio/ACE-0.6B-1024px@models/text_encoder/t5-v1_1-xxl/
TOKENIZER_PATH: hf://scepter-studio/ACE-0.6B-1024px@models/tokenizer/t5-v1_1-xxl
LENGTH: 120
T5_DTYPE: bfloat16
ADDED_IDENTIFIER: [ '{image}', '{caption}', '{mask}', '{ref_image}', '{image1}', '{image2}', '{image3}', '{image4}', '{image5}', '{image6}', '{image7}', '{image8}', '{image9}' ]
CLEAN: whitespace
USE_GRAD: False
REFINER_SCALE: 0.4
#REFINER_PROMPT: "High Resolution, Sharpness, Clarity, Detail Enhancement, Noise Reduction, HD, 4k, Image Restoration, HDR"
REFINER_PROMPT: ""
ACE_PROMPT: [
"A cute cartoon rabbit holding a whiteboard that says 'ACE Refiner', standing in a sunny meadow filled with flowers, with a big smile and bright colors.",
"A beautiful young woman with long flowing hair, wearing a summer dress, holding a whiteboard that reads 'ACE Refiner' while sitting on a park bench surrounded by cherry blossoms.",
"An adorable cartoon cat wearing oversized glasses, holding a whiteboard that says 'ACE Refiner', perched on a stack of colorful books in a cozy library setting.",
"A charming girl with pigtails, wearing a cute school uniform, enthusiastically holding a whiteboard that has 'ACE Refiner' written on it, in a bright and cheerful classroom full of educational posters.",
"A friendly cartoon dog with floppy ears, sitting in front of a doghouse, proudly holding a whiteboard that says 'ACE Refiner', with a playful expression and a blue sky in the background.",
"A cute anime girl with big expressive eyes, dressed in a colorful outfit, holding a whiteboard that reads 'ACE Refiner' in a fantastical landscape filled with mythical creatures.",
"A vibrant cartoon fox holding a whiteboard that says 'ACE Refiner', standing on a rock by a sparkling stream, surrounded by lush greenery and butterflies.",
"A stylish young woman in a business outfit, smiling as she holds a whiteboard written with 'ACE Refiner', in a modern office filled with plants and natural light.",
"A cute cartoon unicorn holding a sparkling whiteboard that says 'ACE Refiner', frolicking in a magical forest, with rainbows and stars in the background.",
"A happy family, consisting of a cute little girl and her playful puppy, holding a whiteboard that says 'ACE Refiner', together in their backyard on a sunny day."
]
REFINER_MODEL:
NAME: ""
IS_DEFAULT: False
DEFAULT_PARAS:
PARAS:
RESOLUTIONS: [ [ 1024, 1024 ] ]
INPUT:
INPUT_IMAGE:
INPUT_MASK:
TASK:
PROMPT: ""
NEGATIVE_PROMPT: ""
OUTPUT_HEIGHT: 1024
OUTPUT_WIDTH: 1024
SAMPLER: flow_euler
SAMPLE_STEPS: 50
GUIDE_SCALE: 3.5
GUIDE_RESCALE: 0.5
OUTPUT:
LATENT:
IMAGES:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: bfloat16
INPUT: [ "IMAGE" ]
- NAME: decode
DTYPE: bfloat16
INPUT: [ "LATENT" ]
PARAS:
SCALE_FACTOR: 1.5305
SHIFT_FACTOR: 0.0609
SIZE_FACTOR: 8
DIFFUSION_MODEL:
FUNCTION:
- NAME: forward
DTYPE: bfloat16
INPUT: [ "SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE" ]
COND_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: bfloat16
INPUT: [ "PROMPT" ]
MODEL:
DIFFUSION:
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
NOISE_SCHEDULER:
NAME: FlowMatchSigmaScheduler
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
LOGIT_MEAN: 0.0
LOGIT_STD: 1.0
MODE_SCALE: 1.29
DIFFUSION_MODEL:
NAME: FluxMR
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@flux1-dev.safetensors
IN_CHANNELS: 64
OUT_CHANNELS: 64
HIDDEN_SIZE: 3072
NUM_HEADS: 24
AXES_DIM: [ 16, 56, 56 ]
THETA: 10000
VEC_IN_DIM: 768
GUIDANCE_EMBED: True
CONTEXT_IN_DIM: 4096
MLP_RATIO: 4.0
QKV_BIAS: True
DEPTH: 19
DEPTH_SINGLE_BLOCKS: 38
USE_GRAD_CHECKPOINT: True
ATTN_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLFlux
EMBED_DIM: 16
PRETRAINED_MODEL: ms://AI-ModelScope/FLUX.1-dev@ae.safetensors
IGNORE_KEYS: [ ]
BATCH_SIZE: 8
USE_CONV: False
SCALE_FACTOR: 0.3611
SHIFT_FACTOR: 0.1159
#
ENCODER:
NAME: Encoder
USE_CHECKPOINT: False
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
USE_CHECKPOINT: False
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: T5PlusClipFluxEmbedder
T5_MODEL:
NAME: HFEmbedder
HF_MODEL_CLS: T5EncoderModel
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder_2/
HF_TOKENIZER_CLS: T5Tokenizer
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer_2/
MAX_LENGTH: 512
OUTPUT_KEY: last_hidden_state
D_TYPE: bfloat16
BATCH_INFER: False
CLEAN: whitespace
CLIP_MODEL:
NAME: HFEmbedder
HF_MODEL_CLS: CLIPTextModel
MODEL_PATH: ms://AI-ModelScope/FLUX.1-dev@text_encoder/
HF_TOKENIZER_CLS: CLIPTokenizer
TOKENIZER_PATH: ms://AI-ModelScope/FLUX.1-dev@tokenizer/
MAX_LENGTH: 77
OUTPUT_KEY: pooler_output
D_TYPE: bfloat16
BATCH_INFER: True
CLEAN: whitespace
REFINER_MODEL_LOCAL:
NAME: ""
IS_DEFAULT: False
DEFAULT_PARAS:
PARAS:
RESOLUTIONS: [ [ 1024, 1024 ] ]
INPUT:
INPUT_IMAGE:
INPUT_MASK:
TASK:
PROMPT: ""
NEGATIVE_PROMPT: ""
OUTPUT_HEIGHT: 1024
OUTPUT_WIDTH: 1024
SAMPLER: flow_euler
SAMPLE_STEPS: 50
GUIDE_SCALE: 3.5
GUIDE_RESCALE: 0.5
OUTPUT:
LATENT:
IMAGES:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: bfloat16
INPUT: [ "IMAGE" ]
- NAME: decode
DTYPE: bfloat16
INPUT: [ "LATENT" ]
PARAS:
SCALE_FACTOR: 1.5305
SHIFT_FACTOR: 0.0609
SIZE_FACTOR: 8
DIFFUSION_MODEL:
FUNCTION:
- NAME: forward
DTYPE: bfloat16
INPUT: [ "SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE" ]
COND_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: bfloat16
INPUT: [ "PROMPT" ]
MODEL:
DIFFUSION:
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
NOISE_SCHEDULER:
NAME: FlowMatchSigmaScheduler
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
LOGIT_MEAN: 0.0
LOGIT_STD: 1.0
MODE_SCALE: 1.29
DIFFUSION_MODEL:
NAME: FluxMR
PRETRAINED_MODEL: models/scepter/FLUX.1-dev/flux1-dev.safetensors
IN_CHANNELS: 64
OUT_CHANNELS: 64
HIDDEN_SIZE: 3072
NUM_HEADS: 24
AXES_DIM: [ 16, 56, 56 ]
THETA: 10000
VEC_IN_DIM: 768
GUIDANCE_EMBED: True
CONTEXT_IN_DIM: 4096
MLP_RATIO: 4.0
QKV_BIAS: True
DEPTH: 19
DEPTH_SINGLE_BLOCKS: 38
USE_GRAD_CHECKPOINT: True
ATTN_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLFlux
EMBED_DIM: 16
PRETRAINED_MODEL: models/scepter/FLUX.1-dev/ae.safetensors
IGNORE_KEYS: [ ]
BATCH_SIZE: 8
USE_CONV: False
SCALE_FACTOR: 0.3611
SHIFT_FACTOR: 0.1159
#
ENCODER:
NAME: Encoder
USE_CHECKPOINT: False
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
USE_CHECKPOINT: False
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: T5PlusClipFluxEmbedder
T5_MODEL:
NAME: HFEmbedder
HF_MODEL_CLS: T5EncoderModel
MODEL_PATH: models/scepter/FLUX.1-dev/text_encoder_2/
HF_TOKENIZER_CLS: T5Tokenizer
TOKENIZER_PATH: models/scepter/FLUX.1-dev/FLUX.1-dev/tokenizer_2/
MAX_LENGTH: 512
OUTPUT_KEY: last_hidden_state
D_TYPE: bfloat16
BATCH_INFER: False
CLEAN: whitespace
CLIP_MODEL:
NAME: HFEmbedder
HF_MODEL_CLS: CLIPTextModel
MODEL_PATH: models/scepter/FLUX.1-dev/FLUX.1-dev/text_encoder/
HF_TOKENIZER_CLS: CLIPTokenizer
TOKENIZER_PATH: models/scepter/FLUX.1-dev/FLUX.1-dev/tokenizer/
MAX_LENGTH: 77
OUTPUT_KEY: pooler_output
D_TYPE: bfloat16
BATCH_INFER: True
CLEAN: whitespace
REFINER_MODEL_HF:
NAME: ""
IS_DEFAULT: False
DEFAULT_PARAS:
PARAS:
RESOLUTIONS: [ [ 1024, 1024 ] ]
INPUT:
INPUT_IMAGE:
INPUT_MASK:
TASK:
PROMPT: ""
NEGATIVE_PROMPT: ""
OUTPUT_HEIGHT: 1024
OUTPUT_WIDTH: 1024
SAMPLER: flow_euler
SAMPLE_STEPS: 50
GUIDE_SCALE: 3.5
GUIDE_RESCALE: 0.5
OUTPUT:
LATENT:
IMAGES:
SEED:
MODULES_PARAS:
FIRST_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: bfloat16
INPUT: [ "IMAGE" ]
- NAME: decode
DTYPE: bfloat16
INPUT: [ "LATENT" ]
PARAS:
SCALE_FACTOR: 1.5305
SHIFT_FACTOR: 0.0609
SIZE_FACTOR: 8
DIFFUSION_MODEL:
FUNCTION:
- NAME: forward
DTYPE: bfloat16
INPUT: [ "SAMPLE_STEPS", "SAMPLE", "GUIDE_SCALE" ]
COND_STAGE_MODEL:
FUNCTION:
- NAME: encode
DTYPE: bfloat16
INPUT: [ "PROMPT" ]
MODEL:
DIFFUSION:
NAME: DiffusionFluxRF
PREDICTION_TYPE: raw
NOISE_SCHEDULER:
NAME: FlowMatchSigmaScheduler
WEIGHTING_SCHEME: logit_normal
SHIFT: 3.0
LOGIT_MEAN: 0.0
LOGIT_STD: 1.0
MODE_SCALE: 1.29
DIFFUSION_MODEL:
NAME: FluxMR
PRETRAINED_MODEL: hf://black-forest-labs/FLUX.1-dev@flux1-dev.safetensors
IN_CHANNELS: 64
OUT_CHANNELS: 64
HIDDEN_SIZE: 3072
NUM_HEADS: 24
AXES_DIM: [ 16, 56, 56 ]
THETA: 10000
VEC_IN_DIM: 768
GUIDANCE_EMBED: True
CONTEXT_IN_DIM: 4096
MLP_RATIO: 4.0
QKV_BIAS: True
DEPTH: 19
DEPTH_SINGLE_BLOCKS: 38
USE_GRAD_CHECKPOINT: True
ATTN_BACKEND: flash_attn
#
FIRST_STAGE_MODEL:
NAME: AutoencoderKLFlux
EMBED_DIM: 16
PRETRAINED_MODEL: hf://black-forest-labs/FLUX.1-dev@ae.safetensors
IGNORE_KEYS: [ ]
BATCH_SIZE: 8
USE_CONV: False
SCALE_FACTOR: 0.3611
SHIFT_FACTOR: 0.1159
#
ENCODER:
NAME: Encoder
USE_CHECKPOINT: False
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
USE_CHECKPOINT: False
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: T5PlusClipFluxEmbedder
T5_MODEL:
NAME: HFEmbedder
HF_MODEL_CLS: T5EncoderModel
MODEL_PATH: hf://black-forest-labs/FLUX.1-dev@text_encoder_2/
HF_TOKENIZER_CLS: T5Tokenizer
TOKENIZER_PATH: hf://black-forest-labs/FLUX.1-dev@tokenizer_2/
MAX_LENGTH: 512
OUTPUT_KEY: last_hidden_state
D_TYPE: bfloat16
BATCH_INFER: False
CLEAN: whitespace
CLIP_MODEL:
NAME: HFEmbedder
HF_MODEL_CLS: CLIPTextModel
MODEL_PATH: hf://black-forest-labs/FLUX.1-dev@text_encoder/
HF_TOKENIZER_CLS: CLIPTokenizer
TOKENIZER_PATH: hf://black-forest-labs/FLUX.1-dev@tokenizer/
MAX_LENGTH: 77
OUTPUT_KEY: pooler_output
D_TYPE: bfloat16
BATCH_INFER: True
CLEAN: whitespace
@@ -39,7 +39,7 @@ DEFAULT_PARAS:
#
COND_STAGE_MODEL:
FUNCTION:
- NAME: encode_list
- NAME: encode_list_of_list
DTYPE: bfloat16
INPUT: ["PROMPT"]
#
+2 -2
View File
@@ -13,8 +13,8 @@ DEFAULT_PARAS:
VISIBLE: False
PROMPT_PREFIX: ""
SAMPLE:
VALUES: ["flow_eluer"]
DEFAULT: "flow_eluer"
VALUES: ["flow_euler"]
DEFAULT: "flow_euler"
SAMPLE_STEPS: 50
GUIDE_SCALE: 3.5
GUIDE_RESCALE:
@@ -13,8 +13,8 @@ DEFAULT_PARAS:
VISIBLE: False
PROMPT_PREFIX: ""
SAMPLE:
VALUES: ["flow_eluer"]
DEFAULT: "flow_eluer"
VALUES: ["flow_euler"]
DEFAULT: "flow_euler"
SAMPLE_STEPS: 4
GUIDE_SCALE: 3.5
GUIDE_RESCALE:
+13 -1
View File
@@ -59,6 +59,18 @@ BASE_MODELS:
FIRST_STAGE_MODEL: ACE_0.6B_512_AutoencoderKL
COND_STAGE_MODEL: ACE_0.6B_512_T5EmbedderHF
CONFIG: config/ace_0.6b_512_pro.yaml
-
NAME: ACE_0.6B_1024
DIFFUSION_MODEL: ACE_0.6B_1024_ACE
FIRST_STAGE_MODEL: ACE_0.6B_1024_AutoencoderKL
COND_STAGE_MODEL: ACE_0.6B_1024_T5EmbedderHF
CONFIG: config/ace_0.6b_1024_pro.yaml
-
NAME: ACE_0.6B_1024_REFINER
DIFFUSION_MODEL: ACE_0.6B_1024_REFINER_ACE
FIRST_STAGE_MODEL: ACE_0.6B_1024_REFINER_AutoencoderKL
COND_STAGE_MODEL: ACE_0.6B_1024_REFINER_T5EmbedderHF
CONFIG: config/ace_0.6b_1024_refiner_pro.yaml
MODEL_SOURCE:
- "ModelScope"
@@ -83,7 +95,7 @@ BASE_PARAMETERS:
- "dpmpp_2m_karras"
- "dpmpp_sde_karras"
- "dpmpp_2m_sde_karras"
- "flow_eluer"
- "flow_euler"
DISCRETIZATION:
- "trailing"
+16 -19
View File
@@ -39,8 +39,8 @@ class ModelNode:
'mantras': ('CONDITIONING', ),
'tuners': ('CONDITIONING', ),
'controls': ('CONDITIONING', ),
'image': ('IMAGE',),
'mask': ('MASK',)
'image': ('IMAGE', ),
'mask': ('MASK', )
}
}
@@ -64,15 +64,17 @@ class ModelNode:
image = [TT.ToPILImage()(image.squeeze(0).permute(2, 0, 1))]
if mask is not None:
mask = [TT.ToPILImage()(mask.squeeze(0))]
data = self.format_parameters(model, model_source, prompt, negative_prompt,
parameters, mantras, tuners, controls, image, mask)
data = self.format_parameters(model, model_source, prompt,
negative_prompt, parameters, mantras,
tuners, controls, image, mask)
cfg = self.model_file.get(model)['config']
cfg = self.source_mapping(cfg, model_source)
self.init_infer(model, cfg)
if model.startswith('ACE'):
output = self.diff_infer(**data[0], **data[1])
output_image = torch.stack([ TT.ToTensor()(img) for img in output]).permute(0, 2, 3, 1).unsqueeze(0)
output_image = torch.stack([TT.ToTensor()(img) for img in output
]).permute(0, 2, 3, 1).unsqueeze(0)
else:
output = self.diff_infer(data[0], **data[1])
x = output['images'].permute(0, 2, 3, 1)
@@ -94,10 +96,14 @@ class ModelNode:
elif source == 'Local':
cfg_new = copy.deepcopy(cfg)
cfg_new.MODEL = cfg_new.MODEL_LOCAL
if hasattr(cfg_new, 'REFINER_MODEL_LOCAL'):
cfg_new.REFINER_MODEL = cfg_new.REFINER_MODEL_LOCAL
return cfg_new
elif source == 'HuggingFace':
cfg_new = copy.deepcopy(cfg)
cfg_new.MODEL = cfg_new.MODEL_HF
if hasattr(cfg_new, 'REFINER_MODEL_HF'):
cfg_new.REFINER_MODEL = cfg_new.REFINER_MODEL_HF
return cfg_new
else:
raise NotImplementedError(f'Unknown model source: {source}')
@@ -143,17 +149,8 @@ class ModelNode:
self.pipeline[model_name] = diff_infer
self.diff_infer = diff_infer
def format_parameters(self,
model,
model_source,
prompt,
negative_prompt,
parameters,
mantras,
tuners,
controls,
image,
mask):
def format_parameters(self, model, model_source, prompt, negative_prompt,
parameters, mantras, tuners, controls, image, mask):
input_data = {'prompt': prompt, 'negative_prompt': negative_prompt}
input_params = {
'diffusion_model': self.model_file.get(model)['diffusion_model'],
@@ -163,13 +160,13 @@ class ModelNode:
}
if image is not None:
input_data.update({"image": image})
input_data.update({'image': image})
if mask is not None:
input_data.update({"mask": mask})
input_data.update({'mask': mask})
if parameters:
seed = parameters.pop('seed', -1)
seed = parameters.get('random_seed', -1)
input_params.update({'seed': seed})
input_data.update(parameters)
+1 -1
View File
@@ -59,6 +59,6 @@ class ParameterNode:
'guide_rescale': guide_rescale,
'discretization': discretization,
'target_size_as_tuple': [output_height, output_width],
'seed': random_seed
'random_seed': random_seed
}
return (out, )
+37 -6
View File
@@ -4,6 +4,7 @@
import os
import unittest
import imageio
import numpy as np
import torchvision.transforms as TT
import torchvision.transforms.functional as TF
@@ -13,13 +14,15 @@ from torchvision.utils import save_image
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.inference.ace_inference import ACEInference
from scepter.modules.inference.diffusion_inference import DiffusionInference
from scepter.modules.inference.flux_inference import FluxInference
from scepter.modules.inference.sd3_inference import SD3Inference
from scepter.modules.inference.flux_inference import FluxInference
from scepter.modules.inference.stylebooth_inference import StyleboothInference
from scepter.modules.inference.cogvideox_inference import CogVideoXInference
from scepter.modules.utils.config import Config
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from scepter.modules.utils.logger import get_logger
from torchvision.utils import save_image
class DiffusionInferenceTest(unittest.TestCase):
@@ -235,10 +238,12 @@ class DiffusionInferenceTest(unittest.TestCase):
cfg = Config(cfg_file=config_file)
diff_infer = SD3Inference(logger=self.logger)
diff_infer.init_from_cfg(cfg)
output = diff_infer({
'prompt': 'a cat holds a blackboard that writes "hello world"',
input_params = {
'seed': 2024
})
}
output = diff_infer({
'prompt': 'a cat holds a blackboard that writes "hello world"'
}, **input_params)
save_path = os.path.join(self.tmp_dir, 'sd3_cat.png')
save_image(output['images'], save_path)
print(save_path)
@@ -249,12 +254,17 @@ class DiffusionInferenceTest(unittest.TestCase):
cfg = Config(cfg_file=config_file)
diff_infer = FluxInference(logger=self.logger)
diff_infer.init_from_cfg(cfg)
output = diff_infer({'prompt': '1 girl', 'seed': 2024})
input_params = {
'seed': 2024
}
output = diff_infer({
'prompt': '1 girl'
}, **input_params)
save_path = os.path.join(self.tmp_dir, 'flux_dev_1girl.png')
save_image(output['images'], save_path)
print(save_path)
# @unittest.skip('')
@unittest.skip('')
def test_ace(self):
config_file = 'scepter/methods/studio/chatbot/models/ace_0.6b_512.yaml'
cfg = Config(cfg_file=config_file)
@@ -266,5 +276,26 @@ class DiffusionInferenceTest(unittest.TestCase):
print(save_path)
# @unittest.skip('')
def test_cogvideox_2b(self):
config_file = 'scepter/methods/studio/inference/dit/cogvideox_2b_pro.yaml'
cfg = Config(cfg_file=config_file)
diff_infer = CogVideoXInference(logger=self.logger)
diff_infer.init_from_cfg(cfg)
input_params = {
'seed': 42
}
output = diff_infer({
'prompt': 'A girl riding a bike.'
}, **input_params)
frames = (output['videos'][0].permute(1, 2, 3, 0).cpu().numpy() * 255).astype(np.uint8)
save_path = os.path.join(self.tmp_dir, 'cogvideox_2b_girlbike.mp4')
writer = imageio.get_writer(save_path, fps=8)
for frame in frames:
writer.append_data(np.array(frame))
writer.close()
print(save_path)
if __name__ == '__main__':
unittest.main()