From a683061c6fa7b80ec93649df1ac1e6f6146b7a09 Mon Sep 17 00:00:00 2001 From: maochaojie Date: Tue, 19 Nov 2024 19:20:02 +0800 Subject: [PATCH] upgrade from 1.2.0 to 1.3.0 --- readme.md | 43 +- requirements/framework.txt | 5 +- requirements/recommended.txt | 7 +- requirements/scepter_studio.txt | 2 +- scepter/methods/edit/dit_ace_0.6b_1024.yaml | 161 ++ .../generation/dit_cogvideox_2b_lora.yaml | 235 +++ .../generation/dit_cogvideox_5b_i2v_lora.yaml | 266 +++ .../generation/dit_cogvideox_5b_lora.yaml | 273 +++ .../generation/dit_flux1.0_dev_1024_lora.yaml | 63 +- .../dit_flux1.0_schnell_1024_lora.yaml | 43 +- scepter/methods/studio/chatbot/chatbot.yaml | 2 +- .../studio/chatbot/models/ace_0.6b_1024.yaml | 127 ++ .../chatbot/models/ace_0.6b_1024_refiner.yaml | 283 +++ .../studio/chatbot/models/ace_0.6b_512.yaml | 2 +- .../inference/dit/cogvideox_2b_pro.yaml | 151 ++ .../inference/dit/cogvideox_5b_pro.yaml | 153 ++ .../studio/inference/dit/flux1.0_dev_pro.yaml | 4 +- .../inference/dit/flux1.0_schnell_pro.yaml | 4 +- .../methods/studio/inference/inference.yaml | 13 +- .../methods/studio/preprocess/preprocess.yaml | 25 + scepter/methods/studio/scepter_ui.yaml | 2 +- .../self_train/dit/cogvideox_2b_pro.yaml | 315 ++++ .../self_train/dit/cogvideox_5b_pro.yaml | 317 ++++ .../studio/self_train/dit/flux1.0_dv_pro.yaml | 94 +- .../self_train/dit/flux1.0_schnell_pro.yaml | 94 +- .../self_train/dit/pixart_alpha_pro.yaml | 2 +- .../methods/studio/self_train/self_train.yaml | 4 +- scepter/modules/data/dataset/__init__.py | 4 +- scepter/modules/data/dataset/dataset.py | 9 +- .../modules/data/dataset/video_gen_dataset.py | 184 ++ scepter/modules/inference/ace_inference.py | 376 ++-- .../modules/inference/cogvideox_inference.py | 181 ++ .../modules/inference/diffusion_inference.py | 4 +- scepter/modules/inference/flux_inference.py | 2 +- scepter/modules/inference/tuner_inference.py | 4 +- scepter/modules/model/backbone/__init__.py | 2 +- .../model/backbone/cogvideox/__init__.py | 3 + .../model/backbone/cogvideox/cogvideox.py | 319 ++++ .../model/backbone/cogvideox/layers.py | 554 ++++++ .../modules/model/backbone/cogvideox/utils.py | 544 ++++++ scepter/modules/model/backbone/flux/flux.py | 130 +- scepter/modules/model/backbone/flux/layers.py | 140 +- scepter/modules/model/diffusion/diffusions.py | 56 +- scepter/modules/model/diffusion/samplers.py | 88 +- scepter/modules/model/diffusion/schedules.py | 63 +- scepter/modules/model/embedder/embedder.py | 103 +- .../model/network/autoencoder/__init__.py | 1 + .../network/autoencoder/ae_kl_cogvideox.py | 1650 +++++++++++++++++ scepter/modules/model/network/ldm/__init__.py | 6 +- scepter/modules/model/network/ldm/ldm_ace.py | 263 ++- .../model/network/ldm/ldm_cogvideox.py | 225 +++ scepter/modules/model/network/ldm/ldm_flux.py | 169 +- scepter/modules/model/utils/basic_utils.py | 22 + scepter/modules/solver/__init__.py | 3 +- scepter/modules/solver/diffusion_solver.py | 192 +- .../modules/solver/diffusion_video_solver.py | 190 ++ scepter/modules/solver/hooks/backward.py | 10 +- scepter/modules/solver/hooks/checkpoint.py | 25 +- scepter/modules/solver/hooks/log.py | 14 +- scepter/modules/utils/distribute.py | 12 +- scepter/modules/utils/visualization.py | 82 +- scepter/studio/chatbot/chatbot.py | 311 +++- .../inference_manager/infer_runer.py | 3 + .../inference/inference_ui/component_names.py | 4 + .../inference/inference_ui/control_ui.py | 1 + .../inference/inference_ui/diffusion_ui.py | 43 +- .../inference/inference_ui/gallery_ui.py | 33 +- .../inference/inference_ui/model_manage_ui.py | 12 +- .../caption_editor_ui/component_names.py | 95 +- .../caption_editor_ui/create_dataset_ui.py | 16 +- .../caption_editor_ui/dataset_gallery_ui.py | 688 +++++-- .../processors/caption_processors.py | 308 ++- .../preprocess/utils/txt2vid_data_card.py | 347 ++++ scepter/studio/self_train/scripts/trainer.py | 18 +- .../self_train_ui/component_names.py | 32 +- .../self_train/self_train_ui/model_ui.py | 24 +- .../self_train/self_train_ui/trainer_ui.py | 140 +- scepter/tools/webui.py | 24 +- scepter/version.py | 2 +- .../workflow/config/ace_0.6b_1024_pro.yaml | 296 +++ .../config/ace_0.6b_1024_refiner_pro.yaml | 729 ++++++++ scepter/workflow/config/ace_0.6b_512_pro.yaml | 2 +- scepter/workflow/config/flux1.0_dev_pro.yaml | 4 +- .../workflow/config/flux1.0_schnell_pro.yaml | 4 +- scepter/workflow/config/scepter_workflow.yaml | 14 +- scepter/workflow/model_node.py | 4 + tests/modules/test_diffusion_inference.py | 43 +- 87 files changed, 10400 insertions(+), 1117 deletions(-) create mode 100644 scepter/methods/edit/dit_ace_0.6b_1024.yaml create mode 100644 scepter/methods/examples/generation/dit_cogvideox_2b_lora.yaml create mode 100644 scepter/methods/examples/generation/dit_cogvideox_5b_i2v_lora.yaml create mode 100644 scepter/methods/examples/generation/dit_cogvideox_5b_lora.yaml create mode 100644 scepter/methods/studio/chatbot/models/ace_0.6b_1024.yaml create mode 100644 scepter/methods/studio/chatbot/models/ace_0.6b_1024_refiner.yaml create mode 100644 scepter/methods/studio/inference/dit/cogvideox_2b_pro.yaml create mode 100644 scepter/methods/studio/inference/dit/cogvideox_5b_pro.yaml create mode 100644 scepter/methods/studio/self_train/dit/cogvideox_2b_pro.yaml create mode 100644 scepter/methods/studio/self_train/dit/cogvideox_5b_pro.yaml create mode 100644 scepter/modules/data/dataset/video_gen_dataset.py create mode 100644 scepter/modules/inference/cogvideox_inference.py create mode 100644 scepter/modules/model/backbone/cogvideox/__init__.py create mode 100644 scepter/modules/model/backbone/cogvideox/cogvideox.py create mode 100644 scepter/modules/model/backbone/cogvideox/layers.py create mode 100644 scepter/modules/model/backbone/cogvideox/utils.py create mode 100644 scepter/modules/model/network/autoencoder/ae_kl_cogvideox.py create mode 100644 scepter/modules/model/network/ldm/ldm_cogvideox.py create mode 100644 scepter/modules/solver/diffusion_video_solver.py create mode 100644 scepter/studio/preprocess/utils/txt2vid_data_card.py create mode 100644 scepter/workflow/config/ace_0.6b_1024_pro.yaml create mode 100644 scepter/workflow/config/ace_0.6b_1024_refiner_pro.yaml diff --git a/readme.md b/readme.md index 41ae787..bb8a8cc 100644 --- a/readme.md +++ b/readme.md @@ -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)
[![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)
[![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 +## 🖼 Gallery for Recent Works + ### FLUX Tuners @@ -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 diff --git a/requirements/framework.txt b/requirements/framework.txt index 44c7dd1..9e8f3ce 100644 --- a/requirements/framework.txt +++ b/requirements/framework.txt @@ -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 \ No newline at end of file diff --git a/requirements/recommended.txt b/requirements/recommended.txt index e976f5c..faed02a 100644 --- a/requirements/recommended.txt +++ b/requirements/recommended.txt @@ -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 \ No newline at end of file diff --git a/requirements/scepter_studio.txt b/requirements/scepter_studio.txt index 369d96f..5811fe9 100644 --- a/requirements/scepter_studio.txt +++ b/requirements/scepter_studio.txt @@ -1,5 +1,5 @@ bitsandbytes -gradio==4.44.1 +gradio gradio_imageslider imagehash psutil diff --git a/scepter/methods/edit/dit_ace_0.6b_1024.yaml b/scepter/methods/edit/dit_ace_0.6b_1024.yaml new file mode 100644 index 0000000..bec3d79 --- /dev/null +++ b/scepter/methods/edit/dit_ace_0.6b_1024.yaml @@ -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 diff --git a/scepter/methods/examples/generation/dit_cogvideox_2b_lora.yaml b/scepter/methods/examples/generation/dit_cogvideox_2b_lora.yaml new file mode 100644 index 0000000..3543c34 --- /dev/null +++ b/scepter/methods/examples/generation/dit_cogvideox_2b_lora.yaml @@ -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 \ No newline at end of file diff --git a/scepter/methods/examples/generation/dit_cogvideox_5b_i2v_lora.yaml b/scepter/methods/examples/generation/dit_cogvideox_5b_i2v_lora.yaml new file mode 100644 index 0000000..3015577 --- /dev/null +++ b/scepter/methods/examples/generation/dit_cogvideox_5b_i2v_lora.yaml @@ -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 \ No newline at end of file diff --git a/scepter/methods/examples/generation/dit_cogvideox_5b_lora.yaml b/scepter/methods/examples/generation/dit_cogvideox_5b_lora.yaml new file mode 100644 index 0000000..56956f1 --- /dev/null +++ b/scepter/methods/examples/generation/dit_cogvideox_5b_lora.yaml @@ -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 \ No newline at end of file diff --git a/scepter/methods/examples/generation/dit_flux1.0_dev_1024_lora.yaml b/scepter/methods/examples/generation/dit_flux1.0_dev_1024_lora.yaml index fcd03bd..5aedf53 100644 --- a/scepter/methods/examples/generation/dit_flux1.0_dev_1024_lora.yaml +++ b/scepter/methods/examples/generation/dit_flux1.0_dev_1024_lora.yaml @@ -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 diff --git a/scepter/methods/examples/generation/dit_flux1.0_schnell_1024_lora.yaml b/scepter/methods/examples/generation/dit_flux1.0_schnell_1024_lora.yaml index 09a4d78..c7a72bd 100644 --- a/scepter/methods/examples/generation/dit_flux1.0_schnell_1024_lora.yaml +++ b/scepter/methods/examples/generation/dit_flux1.0_schnell_1024_lora.yaml @@ -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 diff --git a/scepter/methods/studio/chatbot/chatbot.yaml b/scepter/methods/studio/chatbot/chatbot.yaml index a47eaa1..a2d14ca 100644 --- a/scepter/methods/studio/chatbot/chatbot.yaml +++ b/scepter/methods/studio/chatbot/chatbot.yaml @@ -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/ diff --git a/scepter/methods/studio/chatbot/models/ace_0.6b_1024.yaml b/scepter/methods/studio/chatbot/models/ace_0.6b_1024.yaml new file mode 100644 index 0000000..ae590af --- /dev/null +++ b/scepter/methods/studio/chatbot/models/ace_0.6b_1024.yaml @@ -0,0 +1,127 @@ +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: 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 diff --git a/scepter/methods/studio/chatbot/models/ace_0.6b_1024_refiner.yaml b/scepter/methods/studio/chatbot/models/ace_0.6b_1024_refiner.yaml new file mode 100644 index 0000000..7203272 --- /dev/null +++ b/scepter/methods/studio/chatbot/models/ace_0.6b_1024_refiner.yaml @@ -0,0 +1,283 @@ +NAME: ACE_0.6B_1024_REFINER +IS_DEFAULT: 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 diff --git a/scepter/methods/studio/chatbot/models/ace_0.6b_512.yaml b/scepter/methods/studio/chatbot/models/ace_0.6b_512.yaml index 42d4a24..cc787c9 100644 --- a/scepter/methods/studio/chatbot/models/ace_0.6b_512.yaml +++ b/scepter/methods/studio/chatbot/models/ace_0.6b_512.yaml @@ -39,7 +39,7 @@ DEFAULT_PARAS: # COND_STAGE_MODEL: FUNCTION: - - NAME: encode_list + - NAME: encode_list_of_list DTYPE: bfloat16 INPUT: ["PROMPT"] # diff --git a/scepter/methods/studio/inference/dit/cogvideox_2b_pro.yaml b/scepter/methods/studio/inference/dit/cogvideox_2b_pro.yaml new file mode 100644 index 0000000..3e7c7b7 --- /dev/null +++ b/scepter/methods/studio/inference/dit/cogvideox_2b_pro.yaml @@ -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 \ No newline at end of file diff --git a/scepter/methods/studio/inference/dit/cogvideox_5b_pro.yaml b/scepter/methods/studio/inference/dit/cogvideox_5b_pro.yaml new file mode 100644 index 0000000..d7c3a36 --- /dev/null +++ b/scepter/methods/studio/inference/dit/cogvideox_5b_pro.yaml @@ -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 \ No newline at end of file diff --git a/scepter/methods/studio/inference/dit/flux1.0_dev_pro.yaml b/scepter/methods/studio/inference/dit/flux1.0_dev_pro.yaml index cae7cd1..b1a7be3 100644 --- a/scepter/methods/studio/inference/dit/flux1.0_dev_pro.yaml +++ b/scepter/methods/studio/inference/dit/flux1.0_dev_pro.yaml @@ -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: diff --git a/scepter/methods/studio/inference/dit/flux1.0_schnell_pro.yaml b/scepter/methods/studio/inference/dit/flux1.0_schnell_pro.yaml index c67d9d9..450074f 100644 --- a/scepter/methods/studio/inference/dit/flux1.0_schnell_pro.yaml +++ b/scepter/methods/studio/inference/dit/flux1.0_schnell_pro.yaml @@ -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: diff --git a/scepter/methods/studio/inference/inference.yaml b/scepter/methods/studio/inference/inference.yaml index 5d93331..4abfcdf 100644 --- a/scepter/methods/studio/inference/inference.yaml +++ b/scepter/methods/studio/inference/inference.yaml @@ -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 diff --git a/scepter/methods/studio/preprocess/preprocess.yaml b/scepter/methods/studio/preprocess/preprocess.yaml index 68eabe0..7132568 100644 --- a/scepter/methods/studio/preprocess/preprocess.yaml +++ b/scepter/methods/studio/preprocess/preprocess.yaml @@ -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 \ No newline at end of file diff --git a/scepter/methods/studio/scepter_ui.yaml b/scepter/methods/studio/scepter_ui.yaml index d5e6261..8ec1979 100644 --- a/scepter/methods/studio/scepter_ui.yaml +++ b/scepter/methods/studio/scepter_ui.yaml @@ -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 diff --git a/scepter/methods/studio/self_train/dit/cogvideox_2b_pro.yaml b/scepter/methods/studio/self_train/dit/cogvideox_2b_pro.yaml new file mode 100644 index 0000000..aa47d5f --- /dev/null +++ b/scepter/methods/studio/self_train/dit/cogvideox_2b_pro.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' \ No newline at end of file diff --git a/scepter/methods/studio/self_train/dit/cogvideox_5b_pro.yaml b/scepter/methods/studio/self_train/dit/cogvideox_5b_pro.yaml new file mode 100644 index 0000000..35c2c81 --- /dev/null +++ b/scepter/methods/studio/self_train/dit/cogvideox_5b_pro.yaml @@ -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' \ No newline at end of file diff --git a/scepter/methods/studio/self_train/dit/flux1.0_dv_pro.yaml b/scepter/methods/studio/self_train/dit/flux1.0_dv_pro.yaml index 23701d1..06f53cd 100644 --- a/scepter/methods/studio/self_train/dit/flux1.0_dv_pro.yaml +++ b/scepter/methods/studio/self_train/dit/flux1.0_dv_pro.yaml @@ -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 diff --git a/scepter/methods/studio/self_train/dit/flux1.0_schnell_pro.yaml b/scepter/methods/studio/self_train/dit/flux1.0_schnell_pro.yaml index 265285c..c9c82bd 100644 --- a/scepter/methods/studio/self_train/dit/flux1.0_schnell_pro.yaml +++ b/scepter/methods/studio/self_train/dit/flux1.0_schnell_pro.yaml @@ -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' \ No newline at end of file diff --git a/scepter/methods/studio/self_train/dit/pixart_alpha_pro.yaml b/scepter/methods/studio/self_train/dit/pixart_alpha_pro.yaml index db38854..00369f7 100644 --- a/scepter/methods/studio/self_train/dit/pixart_alpha_pro.yaml +++ b/scepter/methods/studio/self_train/dit/pixart_alpha_pro.yaml @@ -35,7 +35,7 @@ META: SAVE_INTERVAL: 25 EPSEC: 0.818 LEARNING_RATE: 0.0001 - IS_DEFAULT: False + IS_DEFAULT: True TUNER: LORA # TUNERS: diff --git a/scepter/methods/studio/self_train/self_train.yaml b/scepter/methods/studio/self_train/self_train.yaml index d624818..45edc09 100644 --- a/scepter/methods/studio/self_train/self_train.yaml +++ b/scepter/methods/studio/self_train/self_train.yaml @@ -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" diff --git a/scepter/modules/data/dataset/__init__.py b/scepter/modules/data/dataset/__init__.py index fe165fa..347f0c2 100644 --- a/scepter/modules/data/dataset/__init__.py +++ b/scepter/modules/data/dataset/__init__.py @@ -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 \ No newline at end of file diff --git a/scepter/modules/data/dataset/dataset.py b/scepter/modules/data/dataset/dataset.py index 22a6980..1d63e97 100644 --- a/scepter/modules/data/dataset/dataset.py +++ b/scepter/modules/data/dataset/dataset.py @@ -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 diff --git a/scepter/modules/data/dataset/video_gen_dataset.py b/scepter/modules/data/dataset/video_gen_dataset.py new file mode 100644 index 0000000..16a5dd1 --- /dev/null +++ b/scepter/modules/data/dataset/video_gen_dataset.py @@ -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] \ No newline at end of file diff --git a/scepter/modules/inference/ace_inference.py b/scepter/modules/inference/ace_inference.py index 11c4350..e0cf84e 100644 --- a/scepter/modules/inference/ace_inference.py +++ b/scepter/modules/inference/ace_inference.py @@ -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,134 @@ class TextEmbedding(nn.Module): super().__init__() self.pos = nn.Parameter(data=torch.zeros(embedding_shape)) +class RefinerInference(DiffusionInference): + def init_from_cfg(self, cfg): + 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 + + @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=True) + 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=True) + + + 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=True) + 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=True) + return x_samples + class ACEInference(DiffusionInference): def __init__(self, logger=None): @@ -116,9 +244,21 @@ 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_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, @@ -163,6 +303,8 @@ class ACEInference(DiffusionInference): ] return x + + @torch.no_grad() def __call__(self, image=None, @@ -184,7 +326,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 +378,141 @@ 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=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 - # 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=False) + 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=True) + 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=False) + + # 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) + 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) 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 diff --git a/scepter/modules/inference/cogvideox_inference.py b/scepter/modules/inference/cogvideox_inference.py new file mode 100644 index 0000000..a740361 --- /dev/null +++ b/scepter/modules/inference/cogvideox_inference.py @@ -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 diff --git a/scepter/modules/inference/diffusion_inference.py b/scepter/modules/inference/diffusion_inference.py index b705834..cc8c0b8 100644 --- a/scepter/modules/inference/diffusion_inference.py +++ b/scepter/modules/inference/diffusion_inference.py @@ -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 diff --git a/scepter/modules/inference/flux_inference.py b/scepter/modules/inference/flux_inference.py index 4aefc43..d2435e1 100644 --- a/scepter/modules/inference/flux_inference.py +++ b/scepter/modules/inference/flux_inference.py @@ -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: diff --git a/scepter/modules/inference/tuner_inference.py b/scepter/modules/inference/tuner_inference.py index e7c512a..db73795 100644 --- a/scepter/modules/inference/tuner_inference.py +++ b/scepter/modules/inference/tuner_inference.py @@ -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') diff --git a/scepter/modules/model/backbone/__init__.py b/scepter/modules/model/backbone/__init__.py index 71cd841..44b6480 100644 --- a/scepter/modules/model/backbone/__init__.py +++ b/scepter/modules/model/backbone/__init__.py @@ -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) diff --git a/scepter/modules/model/backbone/cogvideox/__init__.py b/scepter/modules/model/backbone/cogvideox/__init__.py new file mode 100644 index 0000000..80b98ec --- /dev/null +++ b/scepter/modules/model/backbone/cogvideox/__init__.py @@ -0,0 +1,3 @@ +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +from scepter.modules.model.backbone.cogvideox.cogvideox import CogVideoXTransformer3DModel \ No newline at end of file diff --git a/scepter/modules/model/backbone/cogvideox/cogvideox.py b/scepter/modules/model/backbone/cogvideox/cogvideox.py new file mode 100644 index 0000000..dcb0fdf --- /dev/null +++ b/scepter/modules/model/backbone/cogvideox/cogvideox.py @@ -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)) diff --git a/scepter/modules/model/backbone/cogvideox/layers.py b/scepter/modules/model/backbone/cogvideox/layers.py new file mode 100644 index 0000000..a1a1990 --- /dev/null +++ b/scepter/modules/model/backbone/cogvideox/layers.py @@ -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 diff --git a/scepter/modules/model/backbone/cogvideox/utils.py b/scepter/modules/model/backbone/cogvideox/utils.py new file mode 100644 index 0000000..03fec8f --- /dev/null +++ b/scepter/modules/model/backbone/cogvideox/utils.py @@ -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) diff --git a/scepter/modules/model/backbone/flux/flux.py b/scepter/modules/model/backbone/flux/flux.py index fb6dfb1..97cbf94 100644 --- a/scepter/modules/model/backbone/flux/flux.py +++ b/scepter/modules/model/backbone/flux/flux.py @@ -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) diff --git a/scepter/modules/model/backbone/flux/layers.py b/scepter/modules/model/backbone/flux/layers.py index eefbcca..9a855d3 100644 --- a/scepter/modules/model/backbone/flux/layers.py +++ b/scepter/modules/model/backbone/flux/layers.py @@ -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, diff --git a/scepter/modules/model/diffusion/diffusions.py b/scepter/modules/model/diffusion/diffusions.py index be63ecb..7c7e15f 100644 --- a/scepter/modules/model/diffusion/diffusions.py +++ b/scepter/modules/model/diffusion/diffusions.py @@ -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': diff --git a/scepter/modules/model/diffusion/samplers.py b/scepter/modules/model/diffusion/samplers.py index 19e563a..e6bb1b8 100644 --- a/scepter/modules/model/diffusion/samplers.py +++ b/scepter/modules/model/diffusion/samplers.py @@ -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 diff --git a/scepter/modules/model/diffusion/schedules.py b/scepter/modules/model/diffusion/schedules.py index 51ccee8..eaef19c 100644 --- a/scepter/modules/model/diffusion/schedules.py +++ b/scepter/modules/model/diffusion/schedules.py @@ -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. diff --git a/scepter/modules/model/embedder/embedder.py b/scepter/modules/model/embedder/embedder.py index bdb4663..4d9f563 100644 --- a/scepter/modules/model/embedder/embedder.py +++ b/scepter/modules/model/embedder/embedder.py @@ -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: diff --git a/scepter/modules/model/network/autoencoder/__init__.py b/scepter/modules/model/network/autoencoder/__init__.py index c0708a2..d7a10b5 100644 --- a/scepter/modules/model/network/autoencoder/__init__.py +++ b/scepter/modules/model/network/autoencoder/__init__.py @@ -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 diff --git a/scepter/modules/model/network/autoencoder/ae_kl_cogvideox.py b/scepter/modules/model/network/autoencoder/ae_kl_cogvideox.py new file mode 100644 index 0000000..f947351 --- /dev/null +++ b/scepter/modules/model/network/autoencoder/ae_kl_cogvideox.py @@ -0,0 +1,1650 @@ +# -*- 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 Dict, Optional, Tuple, Union + +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F + +from scepter.modules.model.network.train_module import TrainModule +from scepter.modules.model.registry import MODELS, 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 scepter.modules.model.base_model import BaseModel +from scepter.modules.model.backbone.cogvideox.utils import get_activation, randn_tensor + + +class DiagonalGaussianDistribution(object): + def __init__(self, parameters: torch.Tensor, deterministic: bool = False): + self.parameters = parameters + self.mean, self.logvar = torch.chunk(parameters, 2, dim=1) + self.logvar = torch.clamp(self.logvar, -30.0, 20.0) + self.deterministic = deterministic + self.std = torch.exp(0.5 * self.logvar) + self.var = torch.exp(self.logvar) + if self.deterministic: + self.var = self.std = torch.zeros_like( + self.mean, device=self.parameters.device, dtype=self.parameters.dtype + ) + + def sample(self, generator: Optional[torch.Generator] = None) -> torch.Tensor: + # make sure sample is on the same device as the parameters and has same dtype + sample = randn_tensor( + self.mean.shape, + generator=generator, + device=self.parameters.device, + dtype=self.parameters.dtype, + ) + x = self.mean + self.std * sample + return x + + def kl(self, other: "DiagonalGaussianDistribution" = None) -> torch.Tensor: + if self.deterministic: + return torch.Tensor([0.0]) + else: + if other is None: + return 0.5 * torch.sum( + torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar, + dim=[1, 2, 3], + ) + else: + return 0.5 * torch.sum( + torch.pow(self.mean - other.mean, 2) / other.var + + self.var / other.var + - 1.0 + - self.logvar + + other.logvar, + dim=[1, 2, 3], + ) + + def nll(self, sample: torch.Tensor, dims: Tuple[int, ...] = [1, 2, 3]) -> torch.Tensor: + if self.deterministic: + return torch.Tensor([0.0]) + logtwopi = np.log(2.0 * np.pi) + return 0.5 * torch.sum( + logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var, + dim=dims, + ) + + def mode(self) -> torch.Tensor: + return self.mean + + +class CogVideoXDownsample3D(nn.Module): + r""" + A 3D Downsampling layer using in [CogVideoX]() by Tsinghua University & ZhipuAI + + Args: + in_channels (`int`): + Number of channels in the input image. + out_channels (`int`): + Number of channels produced by the convolution. + kernel_size (`int`, defaults to `3`): + Size of the convolving kernel. + stride (`int`, defaults to `2`): + Stride of the convolution. + padding (`int`, defaults to `0`): + Padding added to all four sides of the input. + compress_time (`bool`, defaults to `False`): + Whether or not to compress the time dimension. + """ + + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size: int = 3, + stride: int = 2, + padding: int = 0, + compress_time: bool = False, + ): + super().__init__() + + self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding) + self.compress_time = compress_time + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if self.compress_time: + batch_size, channels, frames, height, width = x.shape + + # (batch_size, channels, frames, height, width) -> (batch_size, height, width, channels, frames) -> (batch_size * height * width, channels, frames) + x = x.permute(0, 3, 4, 1, 2).reshape(batch_size * height * width, channels, frames) + + if x.shape[-1] % 2 == 1: + x_first, x_rest = x[..., 0], x[..., 1:] + if x_rest.shape[-1] > 0: + # (batch_size * height * width, channels, frames - 1) -> (batch_size * height * width, channels, (frames - 1) // 2) + x_rest = F.avg_pool1d(x_rest, kernel_size=2, stride=2) + + x = torch.cat([x_first[..., None], x_rest], dim=-1) + # (batch_size * height * width, channels, (frames // 2) + 1) -> (batch_size, height, width, channels, (frames // 2) + 1) -> (batch_size, channels, (frames // 2) + 1, height, width) + x = x.reshape(batch_size, height, width, channels, x.shape[-1]).permute(0, 3, 4, 1, 2) + else: + # (batch_size * height * width, channels, frames) -> (batch_size * height * width, channels, frames // 2) + x = F.avg_pool1d(x, kernel_size=2, stride=2) + # (batch_size * height * width, channels, frames // 2) -> (batch_size, height, width, channels, frames // 2) -> (batch_size, channels, frames // 2, height, width) + x = x.reshape(batch_size, height, width, channels, x.shape[-1]).permute(0, 3, 4, 1, 2) + + # Pad the tensor + pad = (0, 1, 0, 1) + x = F.pad(x, pad, mode="constant", value=0) + batch_size, channels, frames, height, width = x.shape + # (batch_size, channels, frames, height, width) -> (batch_size, frames, channels, height, width) -> (batch_size * frames, channels, height, width) + x = x.permute(0, 2, 1, 3, 4).reshape(batch_size * frames, channels, height, width) + x = self.conv(x) + # (batch_size * frames, channels, height, width) -> (batch_size, frames, channels, height, width) -> (batch_size, channels, frames, height, width) + x = x.reshape(batch_size, frames, x.shape[1], x.shape[2], x.shape[3]).permute(0, 2, 1, 3, 4) + return x + + + +class CogVideoXUpsample3D(nn.Module): + r""" + A 3D Upsample layer using in CogVideoX by Tsinghua University & ZhipuAI # Todo: Wait for paper relase. + + Args: + in_channels (`int`): + Number of channels in the input image. + out_channels (`int`): + Number of channels produced by the convolution. + kernel_size (`int`, defaults to `3`): + Size of the convolving kernel. + stride (`int`, defaults to `1`): + Stride of the convolution. + padding (`int`, defaults to `1`): + Padding added to all four sides of the input. + compress_time (`bool`, defaults to `False`): + Whether or not to compress the time dimension. + """ + + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size: int = 3, + stride: int = 1, + padding: int = 1, + compress_time: bool = False, + ) -> None: + super().__init__() + + self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding) + self.compress_time = compress_time + + def forward(self, inputs: torch.Tensor) -> torch.Tensor: + if self.compress_time: + if inputs.shape[2] > 1 and inputs.shape[2] % 2 == 1: + # split first frame + x_first, x_rest = inputs[:, :, 0], inputs[:, :, 1:] + + x_first = F.interpolate(x_first, scale_factor=2.0) + x_rest = F.interpolate(x_rest, scale_factor=2.0) + x_first = x_first[:, :, None, :, :] + inputs = torch.cat([x_first, x_rest], dim=2) + elif inputs.shape[2] > 1: + inputs = F.interpolate(inputs, scale_factor=2.0) + else: + inputs = inputs.squeeze(2) + inputs = F.interpolate(inputs, scale_factor=2.0) + inputs = inputs[:, :, None, :, :] + else: + # only interpolate 2D + b, c, t, h, w = inputs.shape + inputs = inputs.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w) + inputs = F.interpolate(inputs, scale_factor=2.0) + inputs = inputs.reshape(b, t, c, *inputs.shape[2:]).permute(0, 2, 1, 3, 4) + + b, c, t, h, w = inputs.shape + inputs = inputs.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w) + inputs = self.conv(inputs) + inputs = inputs.reshape(b, t, *inputs.shape[1:]).permute(0, 2, 1, 3, 4) + + return inputs + + +class CogVideoXSafeConv3d(nn.Conv3d): + r""" + A 3D convolution layer that splits the input tensor into smaller parts to avoid OOM in CogVideoX Model. + """ + + def forward(self, input: torch.Tensor) -> torch.Tensor: + memory_count = ( + (input.shape[0] * input.shape[1] * input.shape[2] * input.shape[3] * input.shape[4]) * 2 / 1024**3 + ) + + # Set to 2GB, suitable for CuDNN + if memory_count > 2: + kernel_size = self.kernel_size[0] + part_num = int(memory_count / 2) + 1 + input_chunks = torch.chunk(input, part_num, dim=2) + + if kernel_size > 1: + input_chunks = [input_chunks[0]] + [ + torch.cat((input_chunks[i - 1][:, :, -kernel_size + 1 :], input_chunks[i]), dim=2) + for i in range(1, len(input_chunks)) + ] + + output_chunks = [] + for input_chunk in input_chunks: + output_chunks.append(super().forward(input_chunk)) + output = torch.cat(output_chunks, dim=2) + return output + else: + return super().forward(input) + + +class CogVideoXCausalConv3d(nn.Module): + r"""A 3D causal convolution layer that pads the input tensor to ensure causality in CogVideoX Model. + + Args: + in_channels (`int`): Number of channels in the input tensor. + out_channels (`int`): Number of output channels produced by the convolution. + kernel_size (`int` or `Tuple[int, int, int]`): Kernel size of the convolutional kernel. + stride (`int`, defaults to `1`): Stride of the convolution. + dilation (`int`, defaults to `1`): Dilation rate of the convolution. + pad_mode (`str`, defaults to `"constant"`): Padding mode. + """ + + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size: Union[int, Tuple[int, int, int]], + stride: int = 1, + dilation: int = 1, + pad_mode: str = "constant", + ): + super().__init__() + + if isinstance(kernel_size, int): + kernel_size = (kernel_size,) * 3 + + time_kernel_size, height_kernel_size, width_kernel_size = kernel_size + + self.pad_mode = pad_mode + time_pad = dilation * (time_kernel_size - 1) + (1 - stride) + height_pad = height_kernel_size // 2 + width_pad = width_kernel_size // 2 + + self.height_pad = height_pad + self.width_pad = width_pad + self.time_pad = time_pad + self.time_causal_padding = (width_pad, width_pad, height_pad, height_pad, time_pad, 0) + + self.temporal_dim = 2 + self.time_kernel_size = time_kernel_size + + stride = (stride, 1, 1) + dilation = (dilation, 1, 1) + self.conv = CogVideoXSafeConv3d( + in_channels=in_channels, + out_channels=out_channels, + kernel_size=kernel_size, + stride=stride, + dilation=dilation, + ) + + def fake_context_parallel_forward( + self, inputs: torch.Tensor, conv_cache: Optional[torch.Tensor] = None + ) -> torch.Tensor: + kernel_size = self.time_kernel_size + if kernel_size > 1: + cached_inputs = [conv_cache] if conv_cache is not None else [inputs[:, :, :1]] * (kernel_size - 1) + inputs = torch.cat(cached_inputs + [inputs], dim=2) + return inputs + + def forward(self, inputs: torch.Tensor, conv_cache: Optional[torch.Tensor] = None) -> torch.Tensor: + inputs = self.fake_context_parallel_forward(inputs, conv_cache) + conv_cache = inputs[:, :, -self.time_kernel_size + 1 :].clone() + + padding_2d = (self.width_pad, self.width_pad, self.height_pad, self.height_pad) + inputs = F.pad(inputs, padding_2d, mode="constant", value=0) + + output = self.conv(inputs) + return output, conv_cache + + +class CogVideoXSpatialNorm3D(nn.Module): + r""" + Spatially conditioned normalization as defined in https://arxiv.org/abs/2209.09002. This implementation is specific + to 3D-video like data. + + CogVideoXSafeConv3d is used instead of nn.Conv3d to avoid OOM in CogVideoX Model. + + Args: + f_channels (`int`): + The number of channels for input to group normalization layer, and output of the spatial norm layer. + zq_channels (`int`): + The number of channels for the quantized vector as described in the paper. + groups (`int`): + Number of groups to separate the channels into for group normalization. + """ + + def __init__( + self, + f_channels: int, + zq_channels: int, + groups: int = 32, + ): + super().__init__() + self.norm_layer = nn.GroupNorm(num_channels=f_channels, num_groups=groups, eps=1e-6, affine=True) + self.conv_y = CogVideoXCausalConv3d(zq_channels, f_channels, kernel_size=1, stride=1) + self.conv_b = CogVideoXCausalConv3d(zq_channels, f_channels, kernel_size=1, stride=1) + + def forward( + self, f: torch.Tensor, zq: torch.Tensor, conv_cache: Optional[Dict[str, torch.Tensor]] = None + ) -> torch.Tensor: + new_conv_cache = {} + conv_cache = conv_cache or {} + + if f.shape[2] > 1 and f.shape[2] % 2 == 1: + f_first, f_rest = f[:, :, :1], f[:, :, 1:] + f_first_size, f_rest_size = f_first.shape[-3:], f_rest.shape[-3:] + z_first, z_rest = zq[:, :, :1], zq[:, :, 1:] + z_first = F.interpolate(z_first, size=f_first_size) + z_rest = F.interpolate(z_rest, size=f_rest_size) + zq = torch.cat([z_first, z_rest], dim=2) + else: + zq = F.interpolate(zq, size=f.shape[-3:]) + + conv_y, new_conv_cache["conv_y"] = self.conv_y(zq, conv_cache=conv_cache.get("conv_y")) + conv_b, new_conv_cache["conv_b"] = self.conv_b(zq, conv_cache=conv_cache.get("conv_b")) + + norm_f = self.norm_layer(f) + new_f = norm_f * conv_y + conv_b + return new_f, new_conv_cache + + +class CogVideoXResnetBlock3D(nn.Module): + r""" + A 3D ResNet block used in the CogVideoX model. + + Args: + in_channels (`int`): + Number of input channels. + out_channels (`int`, *optional*): + Number of output channels. If None, defaults to `in_channels`. + dropout (`float`, defaults to `0.0`): + Dropout rate. + temb_channels (`int`, defaults to `512`): + Number of time embedding channels. + groups (`int`, defaults to `32`): + Number of groups to separate the channels into for group normalization. + eps (`float`, defaults to `1e-6`): + Epsilon value for normalization layers. + non_linearity (`str`, defaults to `"swish"`): + Activation function to use. + conv_shortcut (bool, defaults to `False`): + Whether or not to use a convolution shortcut. + spatial_norm_dim (`int`, *optional*): + The dimension to use for spatial norm if it is to be used instead of group norm. + pad_mode (str, defaults to `"first"`): + Padding mode. + """ + + def __init__( + self, + in_channels: int, + out_channels: Optional[int] = None, + dropout: float = 0.0, + temb_channels: int = 512, + groups: int = 32, + eps: float = 1e-6, + non_linearity: str = "swish", + conv_shortcut: bool = False, + spatial_norm_dim: Optional[int] = None, + pad_mode: str = "first", + ): + super().__init__() + + out_channels = out_channels or in_channels + + self.in_channels = in_channels + self.out_channels = out_channels + self.nonlinearity = get_activation(non_linearity) + self.use_conv_shortcut = conv_shortcut + self.spatial_norm_dim = spatial_norm_dim + + if spatial_norm_dim is None: + self.norm1 = nn.GroupNorm(num_channels=in_channels, num_groups=groups, eps=eps) + self.norm2 = nn.GroupNorm(num_channels=out_channels, num_groups=groups, eps=eps) + else: + self.norm1 = CogVideoXSpatialNorm3D( + f_channels=in_channels, + zq_channels=spatial_norm_dim, + groups=groups, + ) + self.norm2 = CogVideoXSpatialNorm3D( + f_channels=out_channels, + zq_channels=spatial_norm_dim, + groups=groups, + ) + + self.conv1 = CogVideoXCausalConv3d( + in_channels=in_channels, out_channels=out_channels, kernel_size=3, pad_mode=pad_mode + ) + + if temb_channels > 0: + self.temb_proj = nn.Linear(in_features=temb_channels, out_features=out_channels) + + self.dropout = nn.Dropout(dropout) + self.conv2 = CogVideoXCausalConv3d( + in_channels=out_channels, out_channels=out_channels, kernel_size=3, pad_mode=pad_mode + ) + + if self.in_channels != self.out_channels: + if self.use_conv_shortcut: + self.conv_shortcut = CogVideoXCausalConv3d( + in_channels=in_channels, out_channels=out_channels, kernel_size=3, pad_mode=pad_mode + ) + else: + self.conv_shortcut = CogVideoXSafeConv3d( + in_channels=in_channels, out_channels=out_channels, kernel_size=1, stride=1, padding=0 + ) + + def forward( + self, + inputs: torch.Tensor, + temb: Optional[torch.Tensor] = None, + zq: Optional[torch.Tensor] = None, + conv_cache: Optional[Dict[str, torch.Tensor]] = None, + ) -> torch.Tensor: + new_conv_cache = {} + conv_cache = conv_cache or {} + + hidden_states = inputs + + if zq is not None: + hidden_states, new_conv_cache["norm1"] = self.norm1(hidden_states, zq, conv_cache=conv_cache.get("norm1")) + else: + hidden_states = self.norm1(hidden_states) + + hidden_states = self.nonlinearity(hidden_states) + hidden_states, new_conv_cache["conv1"] = self.conv1(hidden_states, conv_cache=conv_cache.get("conv1")) + + if temb is not None: + hidden_states = hidden_states + self.temb_proj(self.nonlinearity(temb))[:, :, None, None, None] + + if zq is not None: + hidden_states, new_conv_cache["norm2"] = self.norm2(hidden_states, zq, conv_cache=conv_cache.get("norm2")) + else: + hidden_states = self.norm2(hidden_states) + + hidden_states = self.nonlinearity(hidden_states) + hidden_states = self.dropout(hidden_states) + hidden_states, new_conv_cache["conv2"] = self.conv2(hidden_states, conv_cache=conv_cache.get("conv2")) + + if self.in_channels != self.out_channels: + if self.use_conv_shortcut: + inputs, new_conv_cache["conv_shortcut"] = self.conv_shortcut( + inputs, conv_cache=conv_cache.get("conv_shortcut") + ) + else: + inputs = self.conv_shortcut(inputs) + + hidden_states = hidden_states + inputs + return hidden_states, new_conv_cache + + +class CogVideoXDownBlock3D(nn.Module): + r""" + A downsampling block used in the CogVideoX model. + + Args: + in_channels (`int`): + Number of input channels. + out_channels (`int`, *optional*): + Number of output channels. If None, defaults to `in_channels`. + temb_channels (`int`, defaults to `512`): + Number of time embedding channels. + num_layers (`int`, defaults to `1`): + Number of resnet layers. + dropout (`float`, defaults to `0.0`): + Dropout rate. + resnet_eps (`float`, defaults to `1e-6`): + Epsilon value for normalization layers. + resnet_act_fn (`str`, defaults to `"swish"`): + Activation function to use. + resnet_groups (`int`, defaults to `32`): + Number of groups to separate the channels into for group normalization. + add_downsample (`bool`, defaults to `True`): + Whether or not to use a downsampling layer. If not used, output dimension would be same as input dimension. + compress_time (`bool`, defaults to `False`): + Whether or not to downsample across temporal dimension. + pad_mode (str, defaults to `"first"`): + Padding mode. + """ + + def __init__( + self, + in_channels: int, + out_channels: int, + temb_channels: int, + dropout: float = 0.0, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_act_fn: str = "swish", + resnet_groups: int = 32, + add_downsample: bool = True, + downsample_padding: int = 0, + compress_time: bool = False, + pad_mode: str = "first", + gradient_checkpointing: bool = False + ): + super().__init__() + self.gradient_checkpointing = gradient_checkpointing + resnets = [] + for i in range(num_layers): + in_channel = in_channels if i == 0 else out_channels + resnets.append( + CogVideoXResnetBlock3D( + in_channels=in_channel, + out_channels=out_channels, + dropout=dropout, + temb_channels=temb_channels, + groups=resnet_groups, + eps=resnet_eps, + non_linearity=resnet_act_fn, + pad_mode=pad_mode, + ) + ) + + self.resnets = nn.ModuleList(resnets) + self.downsamplers = None + + if add_downsample: + self.downsamplers = nn.ModuleList( + [ + CogVideoXDownsample3D( + out_channels, out_channels, padding=downsample_padding, compress_time=compress_time + ) + ] + ) + + def forward( + self, + hidden_states: torch.Tensor, + temb: Optional[torch.Tensor] = None, + zq: Optional[torch.Tensor] = None, + conv_cache: Optional[Dict[str, torch.Tensor]] = None, + ) -> torch.Tensor: + r"""Forward method of the `CogVideoXDownBlock3D` class.""" + + new_conv_cache = {} + conv_cache = conv_cache or {} + + for i, resnet in enumerate(self.resnets): + conv_cache_key = f"resnet_{i}" + + if self.training and self.gradient_checkpointing: + + def create_custom_forward(module): + def create_forward(*inputs): + return module(*inputs) + + return create_forward + + hidden_states, new_conv_cache[conv_cache_key] = torch.utils.checkpoint.checkpoint( + create_custom_forward(resnet), + hidden_states, + temb, + zq, + conv_cache=conv_cache.get(conv_cache_key), + ) + else: + hidden_states, new_conv_cache[conv_cache_key] = resnet( + hidden_states, temb, zq, conv_cache=conv_cache.get(conv_cache_key) + ) + + if self.downsamplers is not None: + for downsampler in self.downsamplers: + hidden_states = downsampler(hidden_states) + + return hidden_states, new_conv_cache + + +class CogVideoXMidBlock3D(nn.Module): + r""" + A middle block used in the CogVideoX model. + + Args: + in_channels (`int`): + Number of input channels. + temb_channels (`int`, defaults to `512`): + Number of time embedding channels. + dropout (`float`, defaults to `0.0`): + Dropout rate. + num_layers (`int`, defaults to `1`): + Number of resnet layers. + resnet_eps (`float`, defaults to `1e-6`): + Epsilon value for normalization layers. + resnet_act_fn (`str`, defaults to `"swish"`): + Activation function to use. + resnet_groups (`int`, defaults to `32`): + Number of groups to separate the channels into for group normalization. + spatial_norm_dim (`int`, *optional*): + The dimension to use for spatial norm if it is to be used instead of group norm. + pad_mode (str, defaults to `"first"`): + Padding mode. + """ + + def __init__( + self, + in_channels: int, + temb_channels: int, + dropout: float = 0.0, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_act_fn: str = "swish", + resnet_groups: int = 32, + spatial_norm_dim: Optional[int] = None, + pad_mode: str = "first", + gradient_checkpointing: bool = False + ): + super().__init__() + self.gradient_checkpointing = gradient_checkpointing + resnets = [] + for _ in range(num_layers): + resnets.append( + CogVideoXResnetBlock3D( + in_channels=in_channels, + out_channels=in_channels, + dropout=dropout, + temb_channels=temb_channels, + groups=resnet_groups, + eps=resnet_eps, + spatial_norm_dim=spatial_norm_dim, + non_linearity=resnet_act_fn, + pad_mode=pad_mode, + ) + ) + self.resnets = nn.ModuleList(resnets) + + def forward( + self, + hidden_states: torch.Tensor, + temb: Optional[torch.Tensor] = None, + zq: Optional[torch.Tensor] = None, + conv_cache: Optional[Dict[str, torch.Tensor]] = None, + ) -> torch.Tensor: + r"""Forward method of the `CogVideoXMidBlock3D` class.""" + + new_conv_cache = {} + conv_cache = conv_cache or {} + + for i, resnet in enumerate(self.resnets): + conv_cache_key = f"resnet_{i}" + + if self.training and self.gradient_checkpointing: + + def create_custom_forward(module): + def create_forward(*inputs): + return module(*inputs) + + return create_forward + + hidden_states, new_conv_cache[conv_cache_key] = torch.utils.checkpoint.checkpoint( + create_custom_forward(resnet), hidden_states, temb, zq, conv_cache=conv_cache.get(conv_cache_key) + ) + else: + hidden_states, new_conv_cache[conv_cache_key] = resnet( + hidden_states, temb, zq, conv_cache=conv_cache.get(conv_cache_key) + ) + + return hidden_states, new_conv_cache + + +class CogVideoXUpBlock3D(nn.Module): + r""" + An upsampling block used in the CogVideoX model. + + Args: + in_channels (`int`): + Number of input channels. + out_channels (`int`, *optional*): + Number of output channels. If None, defaults to `in_channels`. + temb_channels (`int`, defaults to `512`): + Number of time embedding channels. + dropout (`float`, defaults to `0.0`): + Dropout rate. + num_layers (`int`, defaults to `1`): + Number of resnet layers. + resnet_eps (`float`, defaults to `1e-6`): + Epsilon value for normalization layers. + resnet_act_fn (`str`, defaults to `"swish"`): + Activation function to use. + resnet_groups (`int`, defaults to `32`): + Number of groups to separate the channels into for group normalization. + spatial_norm_dim (`int`, defaults to `16`): + The dimension to use for spatial norm if it is to be used instead of group norm. + add_upsample (`bool`, defaults to `True`): + Whether or not to use a upsampling layer. If not used, output dimension would be same as input dimension. + compress_time (`bool`, defaults to `False`): + Whether or not to downsample across temporal dimension. + pad_mode (str, defaults to `"first"`): + Padding mode. + """ + + def __init__( + self, + in_channels: int, + out_channels: int, + temb_channels: int, + dropout: float = 0.0, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_act_fn: str = "swish", + resnet_groups: int = 32, + spatial_norm_dim: int = 16, + add_upsample: bool = True, + upsample_padding: int = 1, + compress_time: bool = False, + pad_mode: str = "first", + gradient_checkpointing: bool = False + ): + super().__init__() + self.gradient_checkpointing = gradient_checkpointing + resnets = [] + for i in range(num_layers): + in_channel = in_channels if i == 0 else out_channels + resnets.append( + CogVideoXResnetBlock3D( + in_channels=in_channel, + out_channels=out_channels, + dropout=dropout, + temb_channels=temb_channels, + groups=resnet_groups, + eps=resnet_eps, + non_linearity=resnet_act_fn, + spatial_norm_dim=spatial_norm_dim, + pad_mode=pad_mode, + ) + ) + + self.resnets = nn.ModuleList(resnets) + self.upsamplers = None + + if add_upsample: + self.upsamplers = nn.ModuleList( + [ + CogVideoXUpsample3D( + out_channels, out_channels, padding=upsample_padding, compress_time=compress_time + ) + ] + ) + + def forward( + self, + hidden_states: torch.Tensor, + temb: Optional[torch.Tensor] = None, + zq: Optional[torch.Tensor] = None, + conv_cache: Optional[Dict[str, torch.Tensor]] = None, + ) -> torch.Tensor: + r"""Forward method of the `CogVideoXUpBlock3D` class.""" + + new_conv_cache = {} + conv_cache = conv_cache or {} + + for i, resnet in enumerate(self.resnets): + conv_cache_key = f"resnet_{i}" + + if self.training and self.gradient_checkpointing: + + def create_custom_forward(module): + def create_forward(*inputs): + return module(*inputs) + + return create_forward + + hidden_states, new_conv_cache[conv_cache_key] = torch.utils.checkpoint.checkpoint( + create_custom_forward(resnet), + hidden_states, + temb, + zq, + conv_cache=conv_cache.get(conv_cache_key), + ) + else: + hidden_states, new_conv_cache[conv_cache_key] = resnet( + hidden_states, temb, zq, conv_cache=conv_cache.get(conv_cache_key) + ) + + if self.upsamplers is not None: + for upsampler in self.upsamplers: + hidden_states = upsampler(hidden_states) + + return hidden_states, new_conv_cache + + +@BACKBONES.register_class() +class CogVideoXEncoder3D(BaseModel): + r""" + The `CogVideoXEncoder3D` layer of a variational autoencoder that encodes its input into a latent representation. + + Args: + in_channels (`int`, *optional*, defaults to 3): + The number of input channels. + out_channels (`int`, *optional*, defaults to 3): + The number of output channels. + down_block_types (`Tuple[str, ...]`, *optional*, defaults to `("DownEncoderBlock2D",)`): + The types of down blocks to use. See `~diffusers.models.unet_2d_blocks.get_down_block` for available + options. + block_out_channels (`Tuple[int, ...]`, *optional*, defaults to `(64,)`): + The number of output channels for each block. + act_fn (`str`, *optional*, defaults to `"silu"`): + The activation function to use. See `~diffusers.models.activations.get_activation` for available options. + layers_per_block (`int`, *optional*, defaults to 2): + The number of layers per block. + norm_num_groups (`int`, *optional*, defaults to 32): + The number of groups for normalization. + """ + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + in_channels = cfg.get('IN_CHANNELS', 3) + out_channels = cfg.get('OUT_CHANNELS', 3) + down_block_types = cfg.get('DOWN_BLOCK_TYPES', ["CogVideoXDownBlock3D", + "CogVideoXDownBlock3D", + "CogVideoXDownBlock3D", + "CogVideoXDownBlock3D",]) + block_out_channels = cfg.get('BLOCK_OUT_CHANNELS', [128, 256, 256, 512]) + layers_per_block = cfg.get('LAYERS_PER_BLOCK', 3) + act_fn = cfg.get('ACT_FN', "silu") + norm_eps = cfg.get('NORM_EPS', 1e-6) + norm_num_groups = cfg.get('NORM_NUM_GROUPS', 32) + dropout = cfg.get('DROPOUT', 0.0) + pad_mode = cfg.get('PAD_MODE', "first") + temporal_compression_ratio = cfg.get('TEMPORAL_COMPRESSION_RATIO', 4) + self.gradient_checkpointing = cfg.get('GRADIENT_CHECKPOINTING', False) + + # log2 of temporal_compress_times + temporal_compress_level = int(np.log2(temporal_compression_ratio)) + + self.conv_in = CogVideoXCausalConv3d(in_channels, block_out_channels[0], kernel_size=3, pad_mode=pad_mode) + self.down_blocks = nn.ModuleList([]) + + # down blocks + output_channel = block_out_channels[0] + for i, down_block_type in enumerate(down_block_types): + input_channel = output_channel + output_channel = block_out_channels[i] + is_final_block = i == len(block_out_channels) - 1 + compress_time = i < temporal_compress_level + + if down_block_type == "CogVideoXDownBlock3D": + down_block = CogVideoXDownBlock3D( + in_channels=input_channel, + out_channels=output_channel, + temb_channels=0, + dropout=dropout, + num_layers=layers_per_block, + resnet_eps=norm_eps, + resnet_act_fn=act_fn, + resnet_groups=norm_num_groups, + add_downsample=not is_final_block, + compress_time=compress_time, + gradient_checkpointing=self.gradient_checkpointing + ) + else: + raise ValueError("Invalid `down_block_type` encountered. Must be `CogVideoXDownBlock3D`") + + self.down_blocks.append(down_block) + + # mid block + self.mid_block = CogVideoXMidBlock3D( + in_channels=block_out_channels[-1], + temb_channels=0, + dropout=dropout, + num_layers=2, + resnet_eps=norm_eps, + resnet_act_fn=act_fn, + resnet_groups=norm_num_groups, + pad_mode=pad_mode, + gradient_checkpointing=self.gradient_checkpointing + ) + self.norm_out = nn.GroupNorm(norm_num_groups, block_out_channels[-1], eps=1e-6) + self.conv_act = nn.SiLU() + self.conv_out = CogVideoXCausalConv3d( + block_out_channels[-1], 2 * out_channels, kernel_size=3, pad_mode=pad_mode + ) + + def forward( + self, + sample: torch.Tensor, + temb: Optional[torch.Tensor] = None, + conv_cache: Optional[Dict[str, torch.Tensor]] = None, + ) -> torch.Tensor: + r"""The forward method of the `CogVideoXEncoder3D` class.""" + + new_conv_cache = {} + conv_cache = conv_cache or {} + + hidden_states, new_conv_cache["conv_in"] = self.conv_in(sample, conv_cache=conv_cache.get("conv_in")) + + if self.training and self.gradient_checkpointing: + + def create_custom_forward(module): + def custom_forward(*inputs): + return module(*inputs) + + return custom_forward + + # 1. Down + for i, down_block in enumerate(self.down_blocks): + conv_cache_key = f"down_block_{i}" + hidden_states, new_conv_cache[conv_cache_key] = torch.utils.checkpoint.checkpoint( + create_custom_forward(down_block), + hidden_states, + temb, + None, + conv_cache=conv_cache.get(conv_cache_key), + ) + + # 2. Mid + hidden_states, new_conv_cache["mid_block"] = torch.utils.checkpoint.checkpoint( + create_custom_forward(self.mid_block), + hidden_states, + temb, + None, + conv_cache=conv_cache.get("mid_block"), + ) + else: + # 1. Down + for i, down_block in enumerate(self.down_blocks): + conv_cache_key = f"down_block_{i}" + hidden_states, new_conv_cache[conv_cache_key] = down_block( + hidden_states, temb, None, conv_cache=conv_cache.get(conv_cache_key) + ) + + # 2. Mid + hidden_states, new_conv_cache["mid_block"] = self.mid_block( + hidden_states, temb, None, conv_cache=conv_cache.get("mid_block") + ) + + # 3. Post-process + hidden_states = self.norm_out(hidden_states) + hidden_states = self.conv_act(hidden_states) + + hidden_states, new_conv_cache["conv_out"] = self.conv_out(hidden_states, conv_cache=conv_cache.get("conv_out")) + + return hidden_states, new_conv_cache + +@BACKBONES.register_class() +class CogVideoXDecoder3D(BaseModel): + r""" + The `CogVideoXDecoder3D` layer of a variational autoencoder that decodes its latent representation into an output + sample. + + Args: + in_channels (`int`, *optional*, defaults to 3): + The number of input channels. + out_channels (`int`, *optional*, defaults to 3): + The number of output channels. + up_block_types (`Tuple[str, ...]`, *optional*, defaults to `("UpDecoderBlock2D",)`): + The types of up blocks to use. See `~diffusers.models.unet_2d_blocks.get_up_block` for available options. + block_out_channels (`Tuple[int, ...]`, *optional*, defaults to `(64,)`): + The number of output channels for each block. + act_fn (`str`, *optional*, defaults to `"silu"`): + The activation function to use. See `~diffusers.models.activations.get_activation` for available options. + layers_per_block (`int`, *optional*, defaults to 2): + The number of layers per block. + norm_num_groups (`int`, *optional*, defaults to 32): + The number of groups for normalization. + """ + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + in_channels = cfg.get('IN_CHANNELS', 16) + out_channels = cfg.get('OUT_CHANNELS', 3) + up_block_types = cfg.get('UP_BLOCK_TYPES', ["CogVideoXUpBlock3D", + "CogVideoXUpBlock3D", + "CogVideoXUpBlock3D", + "CogVideoXUpBlock3D",]) + block_out_channels = cfg.get('BLOCK_OUT_CHANNELS', [128, 256, 256, 512]) + layers_per_block = cfg.get('LAYERS_PER_BLOCK', 3) + act_fn = cfg.get('ACT_FN', "silu") + norm_eps = cfg.get('NORM_EPS', 1e-6) + norm_num_groups = cfg.get('NORM_NUM_GROUPS', 32) + dropout = cfg.get('DROPOUT', 0.0) + pad_mode = cfg.get('PAD_MODE', "first") + temporal_compression_ratio = cfg.get('TEMPORAL_COMPRESSION_RATIO', 4) + self.gradient_checkpointing = cfg.get('GRADIENT_CHECKPOINTING', False) + + reversed_block_out_channels = list(reversed(block_out_channels)) + + self.conv_in = CogVideoXCausalConv3d( + in_channels, reversed_block_out_channels[0], kernel_size=3, pad_mode=pad_mode + ) + + # mid block + self.mid_block = CogVideoXMidBlock3D( + in_channels=reversed_block_out_channels[0], + temb_channels=0, + num_layers=2, + resnet_eps=norm_eps, + resnet_act_fn=act_fn, + resnet_groups=norm_num_groups, + spatial_norm_dim=in_channels, + pad_mode=pad_mode, + gradient_checkpointing=self.gradient_checkpointing + ) + + # up blocks + self.up_blocks = nn.ModuleList([]) + + output_channel = reversed_block_out_channels[0] + temporal_compress_level = int(np.log2(temporal_compression_ratio)) + + for i, up_block_type in enumerate(up_block_types): + prev_output_channel = output_channel + output_channel = reversed_block_out_channels[i] + is_final_block = i == len(block_out_channels) - 1 + compress_time = i < temporal_compress_level + + if up_block_type == "CogVideoXUpBlock3D": + up_block = CogVideoXUpBlock3D( + in_channels=prev_output_channel, + out_channels=output_channel, + temb_channels=0, + dropout=dropout, + num_layers=layers_per_block + 1, + resnet_eps=norm_eps, + resnet_act_fn=act_fn, + resnet_groups=norm_num_groups, + spatial_norm_dim=in_channels, + add_upsample=not is_final_block, + compress_time=compress_time, + pad_mode=pad_mode, + gradient_checkpointing=self.gradient_checkpointing + ) + prev_output_channel = output_channel + else: + raise ValueError("Invalid `up_block_type` encountered. Must be `CogVideoXUpBlock3D`") + + self.up_blocks.append(up_block) + + self.norm_out = CogVideoXSpatialNorm3D(reversed_block_out_channels[-1], in_channels, groups=norm_num_groups) + self.conv_act = nn.SiLU() + self.conv_out = CogVideoXCausalConv3d( + reversed_block_out_channels[-1], out_channels, kernel_size=3, pad_mode=pad_mode + ) + + def forward( + self, + sample: torch.Tensor, + temb: Optional[torch.Tensor] = None, + conv_cache: Optional[Dict[str, torch.Tensor]] = None, + ) -> torch.Tensor: + r"""The forward method of the `CogVideoXDecoder3D` class.""" + + new_conv_cache = {} + conv_cache = conv_cache or {} + + hidden_states, new_conv_cache["conv_in"] = self.conv_in(sample, conv_cache=conv_cache.get("conv_in")) + + if self.training and self.gradient_checkpointing: + + def create_custom_forward(module): + def custom_forward(*inputs): + return module(*inputs) + + return custom_forward + + # 1. Mid + hidden_states, new_conv_cache["mid_block"] = torch.utils.checkpoint.checkpoint( + create_custom_forward(self.mid_block), + hidden_states, + temb, + sample, + conv_cache=conv_cache.get("mid_block"), + ) + + # 2. Up + for i, up_block in enumerate(self.up_blocks): + conv_cache_key = f"up_block_{i}" + hidden_states, new_conv_cache[conv_cache_key] = torch.utils.checkpoint.checkpoint( + create_custom_forward(up_block), + hidden_states, + temb, + sample, + conv_cache=conv_cache.get(conv_cache_key), + ) + else: + # 1. Mid + hidden_states, new_conv_cache["mid_block"] = self.mid_block( + hidden_states, temb, sample, conv_cache=conv_cache.get("mid_block") + ) + + # 2. Up + for i, up_block in enumerate(self.up_blocks): + conv_cache_key = f"up_block_{i}" + hidden_states, new_conv_cache[conv_cache_key] = up_block( + hidden_states, temb, sample, conv_cache=conv_cache.get(conv_cache_key) + ) + + # 3. Post-process + hidden_states, new_conv_cache["norm_out"] = self.norm_out( + hidden_states, sample, conv_cache=conv_cache.get("norm_out") + ) + hidden_states = self.conv_act(hidden_states) + hidden_states, new_conv_cache["conv_out"] = self.conv_out(hidden_states, conv_cache=conv_cache.get("conv_out")) + + return hidden_states, new_conv_cache + + +@MODELS.register_class() +class AutoencoderKLCogVideoX(TrainModule): + r""" + A VAE model with KL loss for encoding images into latents and decoding latent representations into images. Used in + [CogVideoX](https://github.com/THUDM/CogVideo). + + This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented + for all models (such as downloading or saving). + + Parameters: + in_channels (int, *optional*, defaults to 3): Number of channels in the input image. + out_channels (int, *optional*, defaults to 3): Number of channels in the output. + down_block_types (`Tuple[str]`, *optional*, defaults to `("DownEncoderBlock2D",)`): + Tuple of downsample block types. + up_block_types (`Tuple[str]`, *optional*, defaults to `("UpDecoderBlock2D",)`): + Tuple of upsample block types. + block_out_channels (`Tuple[int]`, *optional*, defaults to `(64,)`): + Tuple of block output channels. + act_fn (`str`, *optional*, defaults to `"silu"`): The activation function to use. + sample_size (`int`, *optional*, defaults to `32`): Sample input size. + scaling_factor (`float`, *optional*, defaults to `1.15258426`): + The component-wise standard deviation of the trained latent space computed using the first batch of the + training set. This is used to scale the latent space to have unit variance when training the diffusion + model. The latents are scaled with the formula `z = z * scaling_factor` before being passed to the + diffusion model. When decoding, the latents are scaled back to the original scale with the formula: `z = 1 + / scaling_factor * z`. For more details, refer to sections 4.3.2 and D.1 of the [High-Resolution Image + Synthesis with Latent Diffusion Models](https://arxiv.org/abs/2112.10752) paper. + force_upcast (`bool`, *optional*, default to `True`): + If enabled it will force the VAE to run in float32 for high image resolution pipelines, such as SD-XL. VAE + can be fine-tuned / trained to a lower range without loosing too much precision in which case + `force_upcast` can be set to `False` - see: https://huggingface.co/madebyollin/sdxl-vae-fp16-fix + """ + + def __init__(self, cfg, logger=None): + super().__init__(cfg, logger=logger) + self.encoder_cfg = self.cfg.ENCODER + self.decoder_cfg = self.cfg.DECODER + self.encoder = BACKBONES.build(self.encoder_cfg, logger=self.logger) + self.decoder = BACKBONES.build(self.decoder_cfg, logger=self.logger) + + self.out_channels = self.decoder_cfg.OUT_CHANNELS + self.block_out_channels = self.decoder_cfg.BLOCK_OUT_CHANNELS + self.dtype = getattr(torch, cfg.get("DTYPE", "bfloat16")) + sample_height = cfg.get("SAMPLE_HEIGHT", 480) + sample_width = cfg.get("SAMPLE_WIDTH", 720) + use_quant_conv = cfg.get("USE_QUANT_CONV", False) + use_post_quant_conv = cfg.get("USE_POST_QUANT_CONV", False) + self.use_slicing = cfg.get("USE_SLICING", False) + self.use_tiling = cfg.get("USE_TILING", False) + self.scaling_factor_image = cfg.get('SCALING_FACTOR_IMAGE', 1.15258426) + self.gradient_checkpointing = cfg.get('GRADIENT_CHECKPOINTING', False) + + self.quant_conv = CogVideoXSafeConv3d(2 * self.out_channels, 2 * self.out_channels, 1) if use_quant_conv else None + self.post_quant_conv = CogVideoXSafeConv3d(self.out_channels, self.out_channels, 1) if use_post_quant_conv else None + + # Can be increased to decode more latent frames at once, but comes at a reasonable memory cost and it is not + # recommended because the temporal parts of the VAE, here, are tricky to understand. + # If you decode X latent frames together, the number of output frames is: + # (X + (2 conv cache) + (2 time upscale_1) + (4 time upscale_2) - (2 causal conv downscale)) => X + 6 frames + # + # Example with num_latent_frames_batch_size = 2: + # - 12 latent frames: (0, 1), (2, 3), (4, 5), (6, 7), (8, 9), (10, 11) are processed together + # => (12 // 2 frame slices) * ((2 num_latent_frames_batch_size) + (2 conv cache) + (2 time upscale_1) + (4 time upscale_2) - (2 causal conv downscale)) + # => 6 * 8 = 48 frames + # - 13 latent frames: (0, 1, 2) (special case), (3, 4), (5, 6), (7, 8), (9, 10), (11, 12) are processed together + # => (1 frame slice) * ((3 num_latent_frames_batch_size) + (2 conv cache) + (2 time upscale_1) + (4 time upscale_2) - (2 causal conv downscale)) + + # ((13 - 3) // 2) * ((2 num_latent_frames_batch_size) + (2 conv cache) + (2 time upscale_1) + (4 time upscale_2) - (2 causal conv downscale)) + # => 1 * 9 + 5 * 8 = 49 frames + # It has been implemented this way so as to not have "magic values" in the code base that would be hard to explain. Note that + # setting it to anything other than 2 would give poor results because the VAE hasn't been trained to be adaptive with different + # number of temporal frames. + self.num_latent_frames_batch_size = 2 + self.num_sample_frames_batch_size = 8 + + # We make the minimum height and width of sample for tiling half that of the generally supported + self.tile_sample_min_height = sample_height // 2 + self.tile_sample_min_width = sample_width // 2 + self.tile_latent_min_height = int( + self.tile_sample_min_height / (2 ** (len(self.block_out_channels) - 1)) + ) + self.tile_latent_min_width = int(self.tile_sample_min_width / (2 ** (len(self.block_out_channels) - 1))) + + # These are experimental overlap factors that were chosen based on experimentation and seem to work best for + # 720x480 (WxH) resolution. The above resolution is the strongly recommended generation resolution in CogVideoX + # and so the tiling implementation has only been tested on those specific resolutions. + self.tile_overlap_factor_height = 1 / 6 + self.tile_overlap_factor_width = 1 / 5 + + self.enable_slicing() if self.use_slicing else self.disable_slicing() + self.enable_tiling() if self.use_tiling else self.disable_tiling() + + + def enable_tiling( + self, + tile_sample_min_height: Optional[int] = None, + tile_sample_min_width: Optional[int] = None, + tile_overlap_factor_height: Optional[float] = None, + tile_overlap_factor_width: Optional[float] = None, + ) -> None: + r""" + Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to + compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow + processing larger images. + + Args: + tile_sample_min_height (`int`, *optional*): + The minimum height required for a sample to be separated into tiles across the height dimension. + tile_sample_min_width (`int`, *optional*): + The minimum width required for a sample to be separated into tiles across the width dimension. + tile_overlap_factor_height (`int`, *optional*): + The minimum amount of overlap between two consecutive vertical tiles. This is to ensure that there are + no tiling artifacts produced across the height dimension. Must be between 0 and 1. Setting a higher + value might cause more tiles to be processed leading to slow down of the decoding process. + tile_overlap_factor_width (`int`, *optional*): + The minimum amount of overlap between two consecutive horizontal tiles. This is to ensure that there + are no tiling artifacts produced across the width dimension. Must be between 0 and 1. Setting a higher + value might cause more tiles to be processed leading to slow down of the decoding process. + """ + self.use_tiling = True + self.tile_sample_min_height = tile_sample_min_height or self.tile_sample_min_height + self.tile_sample_min_width = tile_sample_min_width or self.tile_sample_min_width + self.tile_latent_min_height = int( + self.tile_sample_min_height / (2 ** (len(self.block_out_channels) - 1)) + ) + self.tile_latent_min_width = int(self.tile_sample_min_width / (2 ** (len(self.block_out_channels) - 1))) + self.tile_overlap_factor_height = tile_overlap_factor_height or self.tile_overlap_factor_height + self.tile_overlap_factor_width = tile_overlap_factor_width or self.tile_overlap_factor_width + + def disable_tiling(self) -> None: + r""" + Disable tiled VAE decoding. If `enable_tiling` was previously enabled, this method will go back to computing + decoding in one step. + """ + self.use_tiling = False + + def enable_slicing(self) -> None: + r""" + Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to + compute decoding in several steps. This is useful to save some memory and allow larger batch sizes. + """ + self.use_slicing = True + + def disable_slicing(self) -> None: + r""" + Disable sliced VAE decoding. If `enable_slicing` was previously enabled, this method will go back to computing + decoding in one step. + """ + self.use_slicing = False + + def _encode(self, x: torch.Tensor) -> torch.Tensor: + batch_size, num_channels, num_frames, height, width = x.shape + + if self.use_tiling and (width > self.tile_sample_min_width or height > self.tile_sample_min_height): + return self.tiled_encode(x) + + frame_batch_size = self.num_sample_frames_batch_size + # Note: We expect the number of frames to be either `1` or `frame_batch_size * k` or `frame_batch_size * k + 1` for some k. + num_batches = num_frames // frame_batch_size if num_frames > 1 else 1 + conv_cache = None + enc = [] + + for i in range(num_batches): + remaining_frames = num_frames % frame_batch_size + start_frame = frame_batch_size * i + (0 if i == 0 else remaining_frames) + end_frame = frame_batch_size * (i + 1) + remaining_frames + x_intermediate = x[:, :, start_frame:end_frame] + x_intermediate, conv_cache = self.encoder(x_intermediate, conv_cache=conv_cache) + if self.quant_conv is not None: + x_intermediate = self.quant_conv(x_intermediate) + enc.append(x_intermediate) + + enc = torch.cat(enc, dim=2) + return enc + + def encode(self, x: torch.Tensor): + """ + Encode a batch of images into latents. + + Args: + x (`torch.Tensor`): Input batch of images. + + Returns: + The latent representations of the encoded videos. + """ + if self.use_slicing and x.shape[0] > 1: + encoded_slices = [self._encode(x_slice) for x_slice in x.split(1)] + h = torch.cat(encoded_slices) + else: + h = self._encode(x) + + posterior = DiagonalGaussianDistribution(h) + return posterior + + def _decode(self, z: torch.Tensor): + batch_size, num_channels, num_frames, height, width = z.shape + + if self.use_tiling and (width > self.tile_latent_min_width or height > self.tile_latent_min_height): + return self.tiled_decode(z) + + frame_batch_size = self.num_latent_frames_batch_size + num_batches = max(num_frames // frame_batch_size, 1) + conv_cache = None + dec = [] + + for i in range(num_batches): + remaining_frames = num_frames % frame_batch_size + start_frame = frame_batch_size * i + (0 if i == 0 else remaining_frames) + end_frame = frame_batch_size * (i + 1) + remaining_frames + z_intermediate = z[:, :, start_frame:end_frame] + if self.post_quant_conv is not None: + z_intermediate = self.post_quant_conv(z_intermediate) + z_intermediate, conv_cache = self.decoder(z_intermediate, conv_cache=conv_cache) + dec.append(z_intermediate) + + dec = torch.cat(dec, dim=2) + return dec + + + def decode(self, z: torch.Tensor): + """ + Decode a batch of images. + + Args: + z (`torch.Tensor`): Input batch of latent vectors. + + Returns: + [`~models.vae.DecoderOutput`] + """ + if self.use_slicing and z.shape[0] > 1: + decoded_slices = [self._decode(z_slice) for z_slice in z.split(1)] + decoded = torch.cat(decoded_slices) + else: + decoded = self._decode(z) + return decoded + + def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: + blend_extent = min(a.shape[3], b.shape[3], blend_extent) + for y in range(blend_extent): + b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * ( + y / blend_extent + ) + return b + + def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor: + blend_extent = min(a.shape[4], b.shape[4], blend_extent) + for x in range(blend_extent): + b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * ( + x / blend_extent + ) + return b + + def tiled_encode(self, x: torch.Tensor) -> torch.Tensor: + r"""Encode a batch of images using a tiled encoder. + + When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several + steps. This is useful to keep memory use constant regardless of image size. The end result of tiled encoding is + different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the + tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the + output, but they should be much less noticeable. + + Args: + x (`torch.Tensor`): Input batch of videos. + + Returns: + `torch.Tensor`: + The latent representation of the encoded videos. + """ + # For a rough memory estimate, take a look at the `tiled_decode` method. + batch_size, num_channels, num_frames, height, width = x.shape + + overlap_height = int(self.tile_sample_min_height * (1 - self.tile_overlap_factor_height)) + overlap_width = int(self.tile_sample_min_width * (1 - self.tile_overlap_factor_width)) + blend_extent_height = int(self.tile_latent_min_height * self.tile_overlap_factor_height) + blend_extent_width = int(self.tile_latent_min_width * self.tile_overlap_factor_width) + row_limit_height = self.tile_latent_min_height - blend_extent_height + row_limit_width = self.tile_latent_min_width - blend_extent_width + frame_batch_size = self.num_sample_frames_batch_size + + # Split x into overlapping tiles and encode them separately. + # The tiles have an overlap to avoid seams between tiles. + rows = [] + for i in range(0, height, overlap_height): + row = [] + for j in range(0, width, overlap_width): + # Note: We expect the number of frames to be either `1` or `frame_batch_size * k` or `frame_batch_size * k + 1` for some k. + num_batches = num_frames // frame_batch_size if num_frames > 1 else 1 + conv_cache = None + time = [] + + for k in range(num_batches): + remaining_frames = num_frames % frame_batch_size + start_frame = frame_batch_size * k + (0 if k == 0 else remaining_frames) + end_frame = frame_batch_size * (k + 1) + remaining_frames + tile = x[ + :, + :, + start_frame:end_frame, + i : i + self.tile_sample_min_height, + j : j + self.tile_sample_min_width, + ] + tile, conv_cache = self.encoder(tile, conv_cache=conv_cache) + if self.quant_conv is not None: + tile = self.quant_conv(tile) + time.append(tile) + + row.append(torch.cat(time, dim=2)) + rows.append(row) + + result_rows = [] + for i, row in enumerate(rows): + result_row = [] + for j, tile in enumerate(row): + # blend the above tile and the left tile + # to the current tile and add the current tile to the result row + if i > 0: + tile = self.blend_v(rows[i - 1][j], tile, blend_extent_height) + if j > 0: + tile = self.blend_h(row[j - 1], tile, blend_extent_width) + result_row.append(tile[:, :, :, :row_limit_height, :row_limit_width]) + result_rows.append(torch.cat(result_row, dim=4)) + + enc = torch.cat(result_rows, dim=3) + return enc + + def tiled_decode(self, z: torch.Tensor): + r""" + Decode a batch of images using a tiled decoder. + + Args: + z (`torch.Tensor`): Input batch of latent vectors. + + Returns: + [`~models.vae.DecoderOutput`] + """ + # Rough memory assessment: + # - In CogVideoX-2B, there are a total of 24 CausalConv3d layers. + # - The biggest intermediate dimensions are: [1, 128, 9, 480, 720]. + # - Assume fp16 (2 bytes per value). + # Memory required: 1 * 128 * 9 * 480 * 720 * 24 * 2 / 1024**3 = 17.8 GB + # + # Memory assessment when using tiling: + # - Assume everything as above but now HxW is 240x360 by tiling in half + # Memory required: 1 * 128 * 9 * 240 * 360 * 24 * 2 / 1024**3 = 4.5 GB + + batch_size, num_channels, num_frames, height, width = z.shape + + overlap_height = int(self.tile_latent_min_height * (1 - self.tile_overlap_factor_height)) + overlap_width = int(self.tile_latent_min_width * (1 - self.tile_overlap_factor_width)) + blend_extent_height = int(self.tile_sample_min_height * self.tile_overlap_factor_height) + blend_extent_width = int(self.tile_sample_min_width * self.tile_overlap_factor_width) + row_limit_height = self.tile_sample_min_height - blend_extent_height + row_limit_width = self.tile_sample_min_width - blend_extent_width + frame_batch_size = self.num_latent_frames_batch_size + + # Split z into overlapping tiles and decode them separately. + # The tiles have an overlap to avoid seams between tiles. + rows = [] + for i in range(0, height, overlap_height): + row = [] + for j in range(0, width, overlap_width): + num_batches = num_frames // frame_batch_size + conv_cache = None + time = [] + + for k in range(num_batches): + remaining_frames = num_frames % frame_batch_size + start_frame = frame_batch_size * k + (0 if k == 0 else remaining_frames) + end_frame = frame_batch_size * (k + 1) + remaining_frames + tile = z[ + :, + :, + start_frame:end_frame, + i : i + self.tile_latent_min_height, + j : j + self.tile_latent_min_width, + ] + if self.post_quant_conv is not None: + tile = self.post_quant_conv(tile) + tile, conv_cache = self.decoder(tile, conv_cache=conv_cache) + time.append(tile) + + row.append(torch.cat(time, dim=2)) + rows.append(row) + + result_rows = [] + for i, row in enumerate(rows): + result_row = [] + for j, tile in enumerate(row): + # blend the above tile and the left tile + # to the current tile and add the current tile to the result row + if i > 0: + tile = self.blend_v(rows[i - 1][j], tile, blend_extent_height) + if j > 0: + tile = self.blend_h(row[j - 1], tile, blend_extent_width) + result_row.append(tile[:, :, :, :row_limit_height, :row_limit_width]) + result_rows.append(torch.cat(result_row, dim=4)) + + dec = torch.cat(result_rows, dim=3) + return dec + + def forward( + self, + sample: torch.Tensor, + sample_posterior: bool = False, + generator: Optional[torch.Generator] = None, + ) -> Union[torch.Tensor, torch.Tensor]: + x = sample + posterior = self.encode(x) + if sample_posterior: + z = posterior.sample(generator=generator) + else: + z = posterior.mode() + dec = self.decode(z) + return dec + + def forward_train(self, sample, sample_posterior=False, generator=None): + return self.forward(sample, sample_posterior, generator) + + def forward_test(self, sample, sample_posterior=False, generator=None): + return self.forward(sample, sample_posterior, generator) + + + @torch.no_grad() + def encode_first_stage(self, x): + if isinstance(x, list): + x = torch.stack(x, dim=0) + latents = self.scaling_factor_image * self.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.decode(latents) + return frames + + def load_pretrained_model(self, pretrained_model): + if pretrained_model is not None: + 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') + missing, unexpected = self.load_state_dict(ckpt, strict=False) + if we.rank == 0: + self.logger.info( + f'Restored from {pretrained_model} with {len(missing)} missing and {len(unexpected)} unexpected keys' + ) + if len(missing) > 0: + self.logger.info(f'Missing Keys:\n {missing}') + if len(unexpected) > 0: + self.logger.info(f'\nUnexpected Keys:\n {unexpected}') + + @staticmethod + def get_config_template(): + return dict_to_yaml('MODEL', + __class__.__name__, + AutoencoderKLCogVideoX.para_dict, + set_name=True) + +def encode_decode_video(model, video_input_path, video_output_path, fps=8, device='cuda'): + import imageio + from torchvision import transforms + + with FS.get_from(video_input_path) as local_read_path: + video_reader = imageio.get_reader(local_read_path, "ffmpeg") + frames = [transforms.ToTensor()(frame) for frame in video_reader] + video_reader.close() + + frames_tensor = torch.stack(frames).to(device).permute(1, 0, 2, 3).unsqueeze(0) + + with torch.no_grad(): + encoded_frames = model.encode(frames_tensor).sample() + decoded_frames = model.decode(encoded_frames) + + frames = decoded_frames.to(dtype=torch.float32) + frames = frames[0].squeeze(0).permute(1, 2, 3, 0).cpu().numpy() + frames = np.clip(frames, 0, 1) * 255 + frames = frames.astype(np.uint8) + + with FS.put_to(video_output_path) as local_save_path: + writer = imageio.get_writer(local_save_path, fps=fps) + for frame in frames: + writer.append_data(frame) + writer.close() + + +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) + ae_model = MODELS.build(cfg.FIRST_STAGE_MODEL, logger=get_logger()).to('cuda') + encode_decode_video(ae_model, cfg.INPUT_PATH, cfg.OUTPUT_PATH, cfg.FPS) diff --git a/scepter/modules/model/network/ldm/__init__.py b/scepter/modules/model/network/ldm/__init__.py index f6198b7..9ebba17 100644 --- a/scepter/modules/model/network/ldm/__init__.py +++ b/scepter/modules/model/network/ldm/__init__.py @@ -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) diff --git a/scepter/modules/model/network/ldm/ldm_ace.py b/scepter/modules/model/network/ldm/ldm_ace.py index 09c32ef..215e156 100644 --- a/scepter/modules/model/network/ldm/ldm_ace.py +++ b/scepter/modules/model/network/ldm/ldm_ace.py @@ -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) diff --git a/scepter/modules/model/network/ldm/ldm_cogvideox.py b/scepter/modules/model/network/ldm/ldm_cogvideox.py new file mode 100644 index 0000000..256e24c --- /dev/null +++ b/scepter/modules/model/network/ldm/ldm_cogvideox.py @@ -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) \ No newline at end of file diff --git a/scepter/modules/model/network/ldm/ldm_flux.py b/scepter/modules/model/network/ldm/ldm_flux.py index a004edb..18455e9 100644 --- a/scepter/modules/model/network/ldm/ldm_flux.py +++ b/scepter/modules/model/network/ldm/ldm_flux.py @@ -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] \ No newline at end of file diff --git a/scepter/modules/model/utils/basic_utils.py b/scepter/modules/model/utils/basic_utils.py index dbf7b5a..bc2a005 100644 --- a/scepter/modules/model/utils/basic_utils.py +++ b/scepter/modules/model/utils/basic_utils.py @@ -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 \ No newline at end of file diff --git a/scepter/modules/solver/__init__.py b/scepter/modules/solver/__init__.py index e404f56..7b0a9a7 100644 --- a/scepter/modules/solver/__init__.py +++ b/scepter/modules/solver/__init__.py @@ -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 \ No newline at end of file diff --git a/scepter/modules/solver/diffusion_solver.py b/scepter/modules/solver/diffusion_solver.py index 63d87d1..3f5c14a 100644 --- a/scepter/modules/solver/diffusion_solver.py +++ b/scepter/modules/solver/diffusion_solver.py @@ -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', True) + 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,21 @@ 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() + self.scaler = amp.GradScaler(enabled=self.enable_gradscaler) else: self.scaler = amp.GradScaler() else: self.scaler = None - self.logger.info(self.model) def load_checkpoint(self, checkpoint: dict): """ @@ -450,10 +480,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 +552,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 +617,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 +659,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'] + " |NegPrompt| " + result['n_prompt']) @@ -678,15 +706,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'] + - " |NegPrompt| " + - result['n_prompt']) + " |NegPrompt| " + + result['n_prompt']) log_data.append(ret_images) log_label.append(ret_labels) ori_label.append(result['prompt']) @@ -715,11 +743,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 +806,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 +849,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'] + - " |NegPrompt| " + - result['n_prompt']) + ret_images.append((result['image'].permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8)) + ret_labels.append(result['prompt'] + + " |NegPrompt| " + + result['n_prompt']) log_data.append(ret_images) log_label.append(ret_labels) self.register_probe({ @@ -947,4 +965,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}.') \ No newline at end of file diff --git a/scepter/modules/solver/diffusion_video_solver.py b/scepter/modules/solver/diffusion_video_solver.py new file mode 100644 index 0000000..3892d92 --- /dev/null +++ b/scepter/modules/solver/diffusion_video_solver.py @@ -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) \ No newline at end of file diff --git a/scepter/modules/solver/hooks/backward.py b/scepter/modules/solver/hooks/backward.py index 9c7d081..5dfb38d 100644 --- a/scepter/modules/solver/hooks/backward.py +++ b/scepter/modules/solver/hooks/backward.py @@ -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() diff --git a/scepter/modules/solver/hooks/checkpoint.py b/scepter/modules/solver/hooks/checkpoint.py index a47de45..61cb071 100644 --- a/scepter/modules/solver/hooks/checkpoint.py +++ b/scepter/modules/solver/hooks/checkpoint.py @@ -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( diff --git a/scepter/modules/solver/hooks/log.py b/scepter/modules/solver/hooks/log.py index 95d3d1e..93e032a 100644 --- a/scepter/modules/solver/hooks/log.py +++ b/scepter/modules/solver/hooks/log.py @@ -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) diff --git a/scepter/modules/utils/distribute.py b/scepter/modules/utils/distribute.py index 6b05f7a..53d5a14 100644 --- a/scepter/modules/utils/distribute.py +++ b/scepter/modules/utils/distribute.py @@ -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 = {} diff --git a/scepter/modules/utils/visualization.py b/scepter/modules/utils/visualization.py index f63ea89..8f2810d 100644 --- a/scepter/modules/utils/visualization.py +++ b/scepter/modules/utils/visualization.py @@ -9,6 +9,7 @@ class Media(Enum): VIDEO = 3 AUDIO = 4 IMAGE_PAIR = 5 + VIDEO_PAIR = 6 class HtmlVisualization(object): @@ -52,11 +53,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 +75,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 @@ -89,13 +93,12 @@ 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 });\n - container.addEventListener('mouseup', () => {\n isDragging = true;\n });\n @@ -108,52 +111,18 @@ class HtmlVisualization(object): let percentage = (clientX - left) / width * 100;\n - - // 限制百分比在0到100之间\n - 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 console.info(slider.style.left);\n });\n - // 初始化滑块位置\n slider.style.left = '50%';\n });\n \n - \n - ''' self.html_body = '{BODY}\n' + self.html_body_script + '\n' @@ -196,7 +165,7 @@ class HtmlVisualization(object): if type == Media.TEXT: ret_str = '' sec_ret_str = f'{label}' if show_label else '' elif type == Media.AUDIO: @@ -229,19 +198,30 @@ class HtmlVisualization(object): sec_ret_str = f'{label}' if show_label else '' elif type == Media.IMAGE_PAIR: assert isinstance(content, (list, tuple)) and len(content) == 2 - ret_str = '\n' - ret_str += '
\n' - f'
' # noqa + f'
' # noqa f' before\n' # noqa f'
\n' # noqa - f'
\n' # noqa + f'
\n' # noqa f' after\n' # noqa f'
\n' # noqa f'
\n' # noqa f'') sec_ret_str = f'{label}' if show_label else '' + elif type == Media.VIDEO_PAIR: + assert isinstance(content, (list, tuple)) and len(content) == 2 + ret_str = f'\n' # noqa + ret_str += f'
\n' + f' \n' # noqa + f' \n' # noqa + f'
\n' # noqa + f'
') + sec_ret_str = f'{label}' if show_label else '' else: raise NotImplementedError if cols_span > 1: diff --git a/scepter/studio/chatbot/chatbot.py b/scepter/studio/chatbot/chatbot.py index af0377f..7bf29bc 100644 --- a/scepter/studio/chatbot/chatbot.py +++ b/scepter/studio/chatbot/chatbot.py @@ -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,6 +50,11 @@ class ChatBotUI(object): is_debug=False, language='en', root_work_dir='./'): + 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}") cfg = Config(cfg_file=cfg_general_file) cfg.WORK_DIR = os.path.join(root_work_dir, cfg.WORK_DIR) @@ -58,25 +62,27 @@ class ChatBotUI(object): 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 = 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 +173,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 +184,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 +200,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 +209,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 +217,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 +282,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 +317,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 +351,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", 0.5), + 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 +444,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 +458,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 +486,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 +530,77 @@ class ChatBotUI(object): lock.acquire() del self.pipe torch.cuda.empty_cache() - model_cfg = Config(load=True, - cfg_file=self.model_choices[model_name]) 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", 0.5), + 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 +657,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 +730,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 +743,8 @@ class ChatBotUI(object): negative_prompt, cfg_scale, rescale, + refiner_prompt, + refiner_scale, step, seed, output_h, @@ -612,12 +756,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 +833,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 +944,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 ] @@ -859,6 +1023,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] @@ -904,14 +1070,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 +1094,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 +1115,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 +1128,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 +1141,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,7 +1384,7 @@ 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): diff --git a/scepter/studio/inference/inference_manager/infer_runer.py b/scepter/studio/inference/inference_manager/infer_runer.py index 288ae9c..e9f6fc7 100644 --- a/scepter/studio/inference/inference_manager/infer_runer.py +++ b/scepter/studio/inference/inference_manager/infer_runer.py @@ -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) diff --git a/scepter/studio/inference/inference_ui/component_names.py b/scepter/studio/inference/inference_ui/component_names.py index c119167..d0b2aec 100644 --- a/scepter/studio/inference/inference_ui/component_names.py +++ b/scepter/studio/inference/inference_ui/component_names.py @@ -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 = '负向提示' diff --git a/scepter/studio/inference/inference_ui/control_ui.py b/scepter/studio/inference/inference_ui/control_ui.py index 81c7884..4936ab5 100644 --- a/scepter/studio/inference/inference_ui/control_ui.py +++ b/scepter/studio/inference/inference_ui/control_ui.py @@ -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') diff --git a/scepter/studio/inference/inference_ui/diffusion_ui.py b/scepter/studio/inference/inference_ui/diffusion_ui.py index 691239d..ca78256 100644 --- a/scepter/studio/inference/inference_ui/diffusion_ui.py +++ b/scepter/studio/inference/inference_ui/diffusion_ui.py @@ -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, diff --git a/scepter/studio/inference/inference_ui/gallery_ui.py b/scepter/studio/inference/inference_ui/gallery_ui.py index 7d2cc6d..571e1f2 100644 --- a/scepter/studio/inference/inference_ui/gallery_ui.py +++ b/scepter/studio/inference/inference_ui/gallery_ui.py @@ -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), ) diff --git a/scepter/studio/inference/inference_ui/model_manage_ui.py b/scepter/studio/inference/inference_ui/model_manage_ui.py index b1e8319..55b9135 100644 --- a/scepter/studio/inference/inference_ui/model_manage_ui.py +++ b/scepter/studio/inference/inference_ui/model_manage_ui.py @@ -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) diff --git a/scepter/studio/preprocess/caption_editor_ui/component_names.py b/scepter/studio/preprocess/caption_editor_ui/component_names.py index 651d72d..2242e3f 100644 --- a/scepter/studio/preprocess/caption_editor_ui/component_names.py +++ b/scepter/studio/preprocess/caption_editor_ui/component_names.py @@ -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' '* 请注意观察系统日志的输出以帮助改进操作。 \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 = '压缩文件失败!' diff --git a/scepter/studio/preprocess/caption_editor_ui/create_dataset_ui.py b/scepter/studio/preprocess/caption_editor_ui/create_dataset_ui.py index 26674a4..0b51320 100644 --- a/scepter/studio/preprocess/caption_editor_ui/create_dataset_ui.py +++ b/scepter/studio/preprocess/caption_editor_ui/create_dataset_ui.py @@ -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], diff --git a/scepter/studio/preprocess/caption_editor_ui/dataset_gallery_ui.py b/scepter/studio/preprocess/caption_editor_ui/dataset_gallery_ui.py index 1fefd8a..6c621a8 100644 --- a/scepter/studio/preprocess/caption_editor_ui/dataset_gallery_ui.py +++ b/scepter/studio/preprocess/caption_editor_ui/dataset_gallery_ui.py @@ -5,6 +5,7 @@ from __future__ import annotations import os.path import time +import cv2 import gradio as gr import imagehash # from gradio.processing_utils import encode_pil_to_base64 @@ -34,7 +35,14 @@ class DatasetGalleryUI(UIBase): self.processors_manager = ProcessorsManager(cfg.PROCESSORS, language=language) + self.video_processors_manager = ProcessorsManager(cfg.VIDEO_PROCESSORS, + language=language) + + self.trans_processors_manager = ProcessorsManager(cfg.TRANSLATION_PROCESSORS, + language=language) + self.component_names = DatasetGalleryUIName(language) + self.translation_name = self.component_names.preprocess_choices[2] if create_ins is not None: self.default_dataset = create_ins.default_dataset current_info = self.default_dataset.current_record @@ -236,11 +244,16 @@ class DatasetGalleryUI(UIBase): with gr.Column(scale=1, min_width=0): self.btn_reset_edit = gr.Button( value=self.component_names.btn_reset_edit) - with gr.Row(): + with gr.Row() as self.mode_select_edit: self.preprocess_checkbox = gr.CheckboxGroup( show_label=False, choices=self.component_names.preprocess_choices, value=None) + with gr.Row(visible=False) as self.mode_select_video: + self.upload_preprocess_video = gr.CheckboxGroup( + show_label=False, + choices=self.component_names.preprocess_choices_video, + value=None) with gr.Column(variant='panel', visible=False, scale=1, @@ -269,7 +282,8 @@ class DatasetGalleryUI(UIBase): type='pil', height=240, interactive=False) - with gr.Row(): + # upload image + with gr.Row() as self.upload_image_row: with gr.Column(scale=1): self.upload_image = gr.Image( label=self.component_names.upload_image, @@ -282,6 +296,23 @@ class DatasetGalleryUI(UIBase): placeholder='', value='', lines=8) + # upload video + with gr.Row(): + with gr.Row(variant='panel', + visible=self.default_dataset_type == + 'scepter_img2video') as self.upload_src_video: + with gr.Column(scale=1): + self.upload_video = gr.Video( + label=self.component_names.upload_video, + sources=['upload'] + ) + with gr.Column(scale=1): + self.upload_caption = gr.Textbox( + label=self.component_names.video_caption, + autoscroll=True, + placeholder='', + value='', + lines=8) with gr.Row(): with gr.Column(min_width=0): self.upload_button = gr.Button( @@ -289,7 +320,8 @@ class DatasetGalleryUI(UIBase): with gr.Column(min_width=0): self.cancel_button = gr.Button( value=self.component_names.cancel_upload_btn) - with gr.Row(): + # mode_select + with gr.Row() as self.mode_select_add: self.upload_preprocess_checkbox = gr.CheckboxGroup( show_label=False, choices=self.component_names.preprocess_choices, @@ -321,7 +353,7 @@ class DatasetGalleryUI(UIBase): interactive=True) self.preview_src_mask_image_tool = gr.State( value='sketch') - with gr.Column(scale=1): + with gr.Column(scale=1) as self.preview_target_image_panel: self.preview_taget_image = gr.ImageMask( label=self.component_names.preview_target_image, sources=[], @@ -331,6 +363,12 @@ class DatasetGalleryUI(UIBase): interactive=True) self.preview_taget_image_tool = gr.State( value='sketch') + with gr.Column(scale=1, visible=False) as self.preview_target_video_panel: + self.preview_target_video = gr.Video( + label=self.component_names.preview_target_video, + sources=['upload'], + interactive=True) + self.preview_target_video_tool = gr.State(value='sketch') with gr.Column(scale=1): self.preview_caption = gr.Textbox( label=self.component_names.preview_caption, @@ -431,7 +469,7 @@ class DatasetGalleryUI(UIBase): 'caption') default_processor_ins = self.processors_manager.get_processor( 'caption', default_processor_method) - with gr.Column(scale=1, min_width=0): + with gr.Column(scale=1, min_width=0) as self.caption_language_panel: self.caption_language = gr.Dropdown( label=self.component_names. caption_language, @@ -450,7 +488,7 @@ class DatasetGalleryUI(UIBase): if len(self.component_names. caption_update_choices) > 0 else None) - with gr.Row(): + with gr.Row() as self.default_use_local_panel: default_use_local = default_processor_ins.get_para_by_language( default_processor_ins.get_language_default ).get('USE_LOCAL', False) @@ -461,7 +499,7 @@ class DatasetGalleryUI(UIBase): interactive=True) with gr.Accordion( label=self.component_names.advance_setting, - open=False): + open=False) as self.advance_setting_panel: with gr.Row(): self.sys_prompt = gr.Text( label=self.component_names. @@ -933,11 +971,28 @@ class DatasetGalleryUI(UIBase): queue=False) def view_mode(): - return gr.Text(value='view') + return (gr.Text(value='view'), + gr.Column(visible=False), + gr.Column(visible=True), + gr.Dropdown(choices=self.processors_manager. + get_choices('caption'), + value=self.processors_manager. + get_default('caption')), + gr.Column(visible=True), + gr.Row(visible=True), + gr.Accordion(visible=True), + gr.CheckboxGroup(value=None)) self.btn_cancel_edit.click(view_mode, inputs=[], - outputs=[self.mode_state], + outputs=[self.mode_state, + self.preview_target_video_panel, + self.preview_target_image_panel, + self.caption_preprocess_method, + self.caption_language_panel, + self.default_use_local_panel, + self.advance_setting_panel, + self.upload_preprocess_video], queue=False) # reset edit information to clean the status of editing @@ -1056,12 +1111,31 @@ class DatasetGalleryUI(UIBase): return (gr.Gallery(), gr.Gallery(), gr.Image(), gr.Textbox(), gr.Text(), gr.Text(), gr.Text(), self.component_names.system_log.format('None'), - gr.Text()) + gr.Text(), gr.Column(visible=False), + gr.Column(visible=True), + gr.Dropdown(choices=self.processors_manager. + get_choices('caption'), + value=self.processors_manager. + get_default('caption')), + gr.Column(visible=True), + gr.Row(visible=True), + gr.Accordion(visible=True), + gr.CheckboxGroup()) is_flg, msg = dataset_ins.apply_changes() if not is_flg: return (gr.Gallery(), gr.Gallery(), gr.Image(), gr.Textbox(), gr.Text(), gr.Text(), gr.Text(), - self.component_names.system_log.format(msg), gr.Text()) + self.component_names.system_log.format(msg), gr.Text(), + gr.Column(visible=False), + gr.Column(visible=True), + gr.Dropdown(choices=self.processors_manager. + get_choices('caption'), + value=self.processors_manager. + get_default('caption')), + gr.Column(visible=True), + gr.Row(visible=True), + gr.Accordion(visible=True), + gr.Checkboxgroup()) image_list = [ os.path.join(dataset_ins.local_work_dir, v['relative_path']) for v in dataset_ins.data @@ -1102,7 +1176,16 @@ class DatasetGalleryUI(UIBase): ret_src_image_gl, ret_mask, gr.Textbox(value=current_record.get('caption', '')), image_info, self.component_names.system_log.format(''), - gr.Text(value='view')) + gr.Text(value='view'), gr.Column(visible=False), + gr.Column(visible=True), + gr.Dropdown(choices=self.processors_manager. + get_choices('caption'), + value=self.processors_manager. + get_default('caption')), + gr.Column(visible=True), + gr.Row(visible=True), + gr.Accordion(visible=True), + gr.CheckboxGroup(value=None)) self.btn_confirm_edit.click( confirm_edit, @@ -1110,15 +1193,19 @@ class DatasetGalleryUI(UIBase): outputs=[ self.gl_dataset_images, self.src_gl_dataset_images, self.src_mask, self.ori_caption, self.image_info, self.sys_log, - self.mode_state + self.mode_state, self.preview_target_video_panel, + self.preview_target_image_panel, + self.caption_preprocess_method, + self.caption_language_panel, + self.default_use_local_panel, + self.advance_setting_panel, + self.upload_preprocess_video ], queue=False) def preprocess_box_change(preprocess_checkbox, dataset_name, - dataset_type, preview_src_image_tool, - preview_src_mask_image_tool, - preview_taget_image_tool): - image_proc_status, caption_proc_status = False, False + dataset_type): + image_proc_status, caption_proc_status, translation_proc_status = False, False, False reverse_status = { v: id for id, v in enumerate(self.component_names.preprocess_choices) @@ -1129,7 +1216,9 @@ class DatasetGalleryUI(UIBase): image_proc_status = True elif hit_status == 1: caption_proc_status = True - if image_proc_status or caption_proc_status: + elif hit_status == 2: + translation_proc_status = True + if image_proc_status or caption_proc_status or translation_proc_status: dataset_type = create_dataset.get_trans_dataset_type( dataset_type) dataset_ins = create_dataset.dataset_dict.get( @@ -1143,27 +1232,11 @@ class DatasetGalleryUI(UIBase): src_image_path = os.path.join( dataset_ins.local_work_dir, one_data['edit_src_relative_path']) - # if preview_src_image_tool == 'sketch': - # image = Image.open(src_image_path) - # w, h = image.size - # ret_src_image = gr.Image(value={ - # 'image': encode_pil_to_base64(image), - # 'mask': default_mask(w, h) - # }, visible=True) - # else: ret_src_image = gr.Image(value=src_image_path, visible=True) src_mask_path = os.path.join( dataset_ins.local_work_dir, one_data['edit_src_mask_relative_path']) - # if preview_src_mask_image_tool == 'sketch': - # image = Image.open(src_mask_path) - # w, h = image.size - # ret_src_mask = gr.Image(value={ - # 'image': encode_pil_to_base64(image), - # 'mask': default_mask(w, h) - # }, visible=True) - # else: ret_src_mask = gr.Image(value=src_mask_path, visible=True) ret_src_panel = gr.Column(visible=True) ret_src_mask_panel = gr.Column(visible=True) @@ -1174,17 +1247,34 @@ class DatasetGalleryUI(UIBase): ret_src_mask_panel = gr.Column(visible=False) image_path = os.path.join(dataset_ins.local_work_dir, one_data['edit_relative_path']) - # if preview_taget_image_tool == 'sketch': - # image = Image.open(image_path) - # w, h = image.size - # ret_target_image = gr.Image(value={ - # 'image': encode_pil_to_base64(image), - # 'mask': default_mask(w, h) - # }) - # else: ret_target_image = gr.Image(value=image_path) ret_caption = gr.Textbox(value=one_data['edit_caption']) ret_panel = gr.Row(visible=True) + + if translation_proc_status and not caption_proc_status: + ret_preprocess_method = gr.Dropdown(choices=self.trans_processors_manager. + get_choices('caption'), + value=self.trans_processors_manager.get_default + ('caption')) + ret_use_device = gr.Text(value=self.trans_processors_manager. + get_default_device('caption')) + ret_use_memory = gr.Text(value=self.trans_processors_manager. + get_default_memory('caption')) + ret_caption_language = gr.Column(visible=False) + ret_use_local = gr.Row(visible=False) + ret_advance_setting = gr.Accordion(visible=False) + else: + ret_caption_language = gr.Column(visible=True) + ret_preprocess_method = gr.Dropdown(choices=self.processors_manager. + get_choices('caption'), + value=self.processors_manager.get_default + ('caption')) + ret_use_device = gr.Text(value=self.processors_manager. + get_default_device('caption')) + ret_use_memory = gr.Text(value=self.processors_manager. + get_default_memory('caption')) + ret_use_local = gr.Row(visible=True) + ret_advance_setting = gr.Accordion(visible=True) else: ret_src_image = gr.Image() ret_src_mask = gr.Image() @@ -1193,22 +1283,29 @@ class DatasetGalleryUI(UIBase): ret_panel = gr.Row(visible=False) ret_src_panel = gr.Column(visible=False) ret_src_mask_panel = gr.Column(visible=False) + ret_caption_language = gr.Column() + ret_preprocess_method = gr.Dropdown() + ret_use_device = gr.Text() + ret_use_memory = gr.Text() + ret_use_local = gr.Row() + ret_advance_setting = gr.Accordion() ret_image_preprocess_method = gr.Dropdown( visible=image_proc_status) return (gr.Column( - visible=image_proc_status or caption_proc_status), + visible=image_proc_status or caption_proc_status or translation_proc_status), gr.Row(visible=image_proc_status), - gr.Row(visible=caption_proc_status), ret_panel, + gr.Row(visible=caption_proc_status or translation_proc_status), ret_panel, ret_src_image, ret_src_panel, ret_src_mask, ret_src_mask_panel, ret_target_image, ret_caption, + ret_caption_language, ret_preprocess_method, + ret_use_device, ret_use_memory, ret_use_local, ret_advance_setting, ret_image_preprocess_method) self.preprocess_checkbox.change( preprocess_box_change, inputs=[ self.preprocess_checkbox, create_dataset.dataset_name, - create_dataset.dataset_type, self.preview_src_image_tool, - self.preview_src_mask_image_tool, self.preview_taget_image_tool + create_dataset.dataset_type ], outputs=[ self.preprocess_panel, @@ -1228,6 +1325,15 @@ class DatasetGalleryUI(UIBase): self.preview_taget_image, # gr.TextBox() self.preview_caption, + # gr.Column() + self.caption_language_panel, + self.caption_preprocess_method, + self.caption_use_device, + self.caption_use_memory, + # gr.Row() + self.default_use_local_panel, + # gr.Accordion() + self.advance_setting_panel, self.image_preprocess_method ], queue=False) @@ -1236,8 +1342,7 @@ class DatasetGalleryUI(UIBase): preprocess_box_change, inputs=[ self.preprocess_checkbox, create_dataset.dataset_name, - create_dataset.dataset_type, self.preview_src_image_tool, - self.preview_src_mask_image_tool, self.preview_taget_image_tool + create_dataset.dataset_type ], outputs=[ self.preprocess_panel, @@ -1257,10 +1362,143 @@ class DatasetGalleryUI(UIBase): self.preview_taget_image, # gr.TextBox() self.preview_caption, + # gr.Column() + self.caption_language_panel, + self.caption_preprocess_method, + self.caption_use_device, + self.caption_use_memory, + # gr.Row() + self.default_use_local_panel, + # gr.Accordion() + self.advance_setting_panel, self.image_preprocess_method ], ) + def preprocess_box_change_video(preprocess_checkbox, dataset_name, dataset_type): + translation_proc_status, caption_proc_status = False, False + reverse_status = { + v: id + for id, v in enumerate(self.component_names.preprocess_choices_video) + } + for value in preprocess_checkbox: + hit_status = reverse_status[value] + if hit_status == 0: + caption_proc_status = True + elif hit_status == 1: + translation_proc_status = True + if translation_proc_status or caption_proc_status: + dataset_type = create_dataset.get_trans_dataset_type( + dataset_type) + dataset_ins = create_dataset.dataset_dict.get( + dataset_type, {}).get(dataset_name, None) + edit_index_list = dataset_ins.edit_list + if len(edit_index_list) > 0: + one_data = dataset_ins.data[edit_index_list[0]] + else: + one_data = {} + + ret_src_image = gr.Image(visible=False) + ret_src_mask = gr.Image(visible=False) + ret_src_panel = gr.Column(visible=False) + ret_src_mask_panel = gr.Column(visible=False) + + video_path = os.path.join(dataset_ins.local_work_dir, + one_data['relative_path']) + ret_target_video = gr.Video(value=video_path) + ret_target_video_panel = gr.Column(visible=True) + ret_target_image = gr.Column(visible=False) + ret_caption_language = gr.Column(visible=False) + ret_use_local = gr.Row(visible=False) + ret_advance_setting = gr.Accordion(visible=False) + ret_caption = gr.Textbox(value=one_data['caption']) + ret_panel = gr.Row(visible=True) + if translation_proc_status: + ret_preprocess_method = gr.Dropdown(choices=self.trans_processors_manager. + get_choices('caption'), + value=self.trans_processors_manager.get_default + ('caption')) + ret_use_device = gr.Text(value=self.trans_processors_manager. + get_default_device('caption')) + ret_use_memory = gr.Text(value=self.trans_processors_manager. + get_default_memory('caption')) + elif caption_proc_status: + ret_preprocess_method = gr.Dropdown(choices=self.video_processors_manager. + get_choices('caption'), + value=self.video_processors_manager.get_default + ('caption')) + ret_use_device = gr.Text(value=self.video_processors_manager. + get_default_device('caption')) + ret_use_memory = gr.Text(value=self.video_processors_manager. + get_default_memory('caption')) + else: + ret_src_image = gr.Image() + ret_src_mask = gr.Image() + ret_target_video = gr.Video() + ret_target_video_panel = gr.Column() + ret_target_image = gr.Column() + ret_caption = gr.Textbox() + ret_panel = gr.Row(visible=False) + ret_caption_language = gr.Column() + ret_use_local = gr.Row() + ret_advance_setting = gr.Accordion() + ret_src_panel = gr.Column(visible=False) + ret_preprocess_method = gr.Dropdown() + ret_src_mask_panel = gr.Column(visible=False) + ret_use_device = gr.Text() + ret_use_memory = gr.Text() + + ret_image_preprocess_method = gr.Dropdown( + visible=caption_proc_status) + return (gr.Column( + visible=caption_proc_status or translation_proc_status), + gr.Row(visible=False), + gr.Row(visible=caption_proc_status or translation_proc_status), + ret_panel, ret_src_image, ret_src_panel, ret_src_mask, + ret_src_mask_panel, ret_target_video, ret_target_video_panel, + ret_target_image, ret_caption, ret_caption_language, + ret_use_local, ret_advance_setting, ret_preprocess_method, + ret_use_device, ret_use_memory, ret_image_preprocess_method) + + self.upload_preprocess_video.change( + preprocess_box_change_video, + inputs=[self.upload_preprocess_video, + create_dataset.dataset_name, + create_dataset.dataset_type], + outputs=[self.preprocess_panel, + self.image_preprocess_panel, + self.caption_preprocess_panel, + # gr.Row() + self.preview_panel, + # gr.Image() + self.preview_src_image, + # gr.Column() + self.preview_src_panel, + # gr.Image() + self.preview_src_mask_image, + # gr.Column() + self.preview_src_mask_panel, + # gr.Video() + self.preview_target_video, + # gr.Column() + self.preview_target_video_panel, + # gr.Column() + self.preview_target_image_panel, + # gr.TextBox() + self.preview_caption, + # gr.Column() + self.caption_language_panel, + # gr.Row() + self.default_use_local_panel, + # gr.Accordion() + self.advance_setting_panel, + self.caption_preprocess_method, + self.caption_use_device, + self.caption_use_memory, + self.image_preprocess_method + ] + ) + def upload_preprocess_box_change(preprocess_checkbox, dataset_type, upload_src_image, upload_src_mask, upload_image, upload_caption): @@ -1749,24 +1987,28 @@ class DatasetGalleryUI(UIBase): ], queue=False) - def caption_preprocess_method_change(caption_preprocess_method): - processor_ins = self.processors_manager.get_processor( - 'caption', caption_preprocess_method) + def caption_preprocess_method_change(caption_preprocess_method, dataset_type): + dataset_type = create_dataset.get_trans_dataset_type(dataset_type) + if dataset_type == 'scepter_txt2vid': + processor_ins = self.video_processors_manager.get_processor( + 'caption', caption_preprocess_method) + else: + processor_ins = self.processors_manager.get_processor( + 'caption', caption_preprocess_method) if processor_ins is None: - return (gr.Text(), gr.Text(), gr.Dropdown(), - self.component_names.system_log.format( - 'Load processor failed, processor is None.')) + return (gr.Text(), + gr.Text(), + gr.Dropdown()) language_choice = processor_ins.get_language_choice language_default = processor_ins.get_language_default return (gr.Text(value=processor_ins.use_device), gr.Text(value=f'{processor_ins.use_memory}M'), gr.Dropdown(choices=language_choice, - value=language_default), - self.component_names.system_log.format('')) + value=language_default)) self.caption_preprocess_method.change( caption_preprocess_method_change, - inputs=[self.caption_preprocess_method], + inputs=[self.caption_preprocess_method, create_dataset.dataset_type], outputs=[ self.caption_use_device, self.caption_use_memory, self.caption_language @@ -1774,9 +2016,15 @@ class DatasetGalleryUI(UIBase): queue=False) def caption_language_change(caption_language, - caption_preprocess_method): - processor_ins = self.processors_manager.get_processor( - 'caption', caption_preprocess_method) + caption_preprocess_method, + dataset_type): + dataset_type = create_dataset.get_trans_dataset_type(dataset_type) + if dataset_type == 'scepter_txt2vid': + processor_ins = self.video_processors_manager.get_processor( + 'caption', caption_preprocess_method) + else: + processor_ins = self.processors_manager.get_processor( + 'caption', caption_preprocess_method) para = processor_ins.get_para_by_language(caption_language) system_prompt = para.get('PROMPT', '') ret_system_prompt = gr.Text(value=system_prompt, @@ -1824,7 +2072,7 @@ class DatasetGalleryUI(UIBase): self.caption_language.change( caption_language_change, - inputs=[self.caption_language, self.caption_preprocess_method], + inputs=[self.caption_language, self.caption_preprocess_method, create_dataset.dataset_type], outputs=[ self.sys_prompt, self.max_new_tokens, self.min_new_tokens, self.num_beams, self.repetition_penalty, self.temperature @@ -1836,18 +2084,25 @@ class DatasetGalleryUI(UIBase): repetition_penalty, temperature, use_local, caption_update_mode, upload_image, upload_src_image, upload_src_mask, - upload_caption, dataset_type, dataset_name): - + upload_caption, dataset_type, dataset_name, + preprocess_name, preprocess_name_txt2vid): reverse_update_mode = { v: idx for idx, v in enumerate( self.component_names.caption_update_choices) } - + dataset_type = create_dataset.get_trans_dataset_type(dataset_type) + trans_status = get_trans_status(dataset_type, preprocess_name, preprocess_name_txt2vid) update_mode = reverse_update_mode.get(caption_update_mode, -1) - - processor_ins = self.processors_manager.get_processor( - 'caption', preprocess_method) + if dataset_type == 'scepter_txt2vid' and not trans_status: + processor_ins = self.video_processors_manager.get_processor( + 'caption', preprocess_method) + elif trans_status: + processor_ins = self.trans_processors_manager.get_processor( + 'caption', preprocess_method) + else: + processor_ins = self.processors_manager.get_processor( + 'caption', preprocess_method) if processor_ins is None: sys_log = 'Current processor is illegal' return gr.Textbox(), gr.Textbox( @@ -1858,58 +2113,70 @@ class DatasetGalleryUI(UIBase): sys_log = f'Load processor failed: {msg}' return gr.Textbox(), gr.Textbox( ), self.component_names.system_log.format(sys_log) - dataset_type = create_dataset.get_trans_dataset_type(dataset_type) dataset_ins = create_dataset.dataset_dict.get( dataset_type, {}).get(dataset_name, None) if dataset_ins is None: return gr.Textbox(), gr.Textbox( ), self.component_names.system_log.format('None') + if mode_state == 'edit': edit_index_list = dataset_ins.edit_list for index in edit_index_list: one_data = dataset_ins.data[index] - relative_image_path = one_data.get( - 'edit_relative_path', one_data['relative_path']) - target_image = Image.open( - os.path.join(dataset_ins.meta['local_work_dir'], - relative_image_path)) - - if dataset_type == 'scepter_img2img': + if (dataset_type == 'scepter_txt2vid' and + not trans_status): + video_path = one_data.get('video_path', None) + kwargs = { + 'video_path': video_path + } + elif trans_status: + caption = one_data.get('caption', None) + kwargs = { + 'caption': caption + } + else: relative_image_path = one_data.get( - 'edit_relative_path', - one_data['src_relative_path']) - src_image = Image.open( + 'edit_relative_path', one_data['relative_path']) + target_image = Image.open( os.path.join(dataset_ins.meta['local_work_dir'], relative_image_path)) - relative_src_mask_path = one_data.get( - 'edit_relative_path', - one_data['src_mask_relative_path']) - src_mask_image = Image.open( - os.path.join(dataset_ins.meta['local_work_dir'], - relative_src_mask_path)) - else: - src_image = None - src_mask_image = None - kwargs = { - 'src_image': src_image, - 'src_mask': src_mask_image, - 'target_image': target_image, - 'caption': one_data.get('edit_caption', ''), - 'preview_src_image': None, - 'preview_src_mask': None, - 'preview_target_image': None, - 'preview_caption': None, - 'use_preview': False, - 'use_local': use_local, - 'sys_prompt': sys_prompt, - 'max_new_tokens': max_new_tokens, - 'min_new_tokens': min_new_tokens, - 'num_beams': num_beams, - 'repetition_penalty': repetition_penalty, - 'temperature': temperature, - 'cache': self.cache - } + if dataset_type == 'scepter_img2img': + relative_image_path = one_data.get( + 'edit_relative_path', + one_data['src_relative_path']) + src_image = Image.open( + os.path.join(dataset_ins.meta['local_work_dir'], + relative_image_path)) + relative_src_mask_path = one_data.get( + 'edit_relative_path', + one_data['src_mask_relative_path']) + src_mask_image = Image.open( + os.path.join(dataset_ins.meta['local_work_dir'], + relative_src_mask_path)) + else: + src_image = None + src_mask_image = None + + kwargs = { + 'src_image': src_image, + 'src_mask': src_mask_image, + 'target_image': target_image, + 'caption': one_data.get('edit_caption', ''), + 'preview_src_image': None, + 'preview_src_mask': None, + 'preview_target_image': None, + 'preview_caption': None, + 'use_preview': False, + 'use_local': use_local, + 'sys_prompt': sys_prompt, + 'max_new_tokens': max_new_tokens, + 'min_new_tokens': min_new_tokens, + 'num_beams': num_beams, + 'repetition_penalty': repetition_penalty, + 'temperature': temperature, + 'cache': self.cache + } response = processor_ins(**kwargs) if update_mode == 0: @@ -1999,7 +2266,8 @@ class DatasetGalleryUI(UIBase): self.use_local, self.caption_update_mode, self.upload_image, self.upload_src_image, self.upload_src_mask, self.upload_caption, create_dataset.dataset_type, - create_dataset.dataset_name + create_dataset.dataset_name, self.preprocess_checkbox, + self.upload_preprocess_video ], outputs=[self.edit_caption, self.upload_caption, self.sys_log], queue=False) @@ -2070,7 +2338,8 @@ class DatasetGalleryUI(UIBase): preprocess_method, dataset_type, sys_prompt, max_new_tokens, min_new_tokens, num_beams, repetition_penalty, temperature, use_local, caption_update_mode, preview_src_image, - preview_src_mask_image, preview_target_image, preview_caption): + preview_src_mask_image, preview_target_image, preview_caption, + video_path, preprocess_name, preprocess_name_txt2vid): reverse_update_mode = { v: idx for idx, v in enumerate( @@ -2078,60 +2347,68 @@ class DatasetGalleryUI(UIBase): } update_mode = reverse_update_mode.get(caption_update_mode, -1) + dataset_type = create_dataset.get_trans_dataset_type(dataset_type) + trans_status = get_trans_status(dataset_type, preprocess_name, preprocess_name_txt2vid) + if dataset_type == 'scepter_txt2vid' and not trans_status: + processor_ins = self.video_processors_manager.get_processor( + 'caption', preprocess_method) + kwargs = { + 'video_path': video_path + } + elif trans_status: + processor_ins = self.trans_processors_manager.get_processor( + 'caption', preprocess_method) + kwargs = { + 'caption': preview_caption + } + else: + processor_ins = self.processors_manager.get_processor( + 'caption', preprocess_method) + if isinstance(preview_target_image, dict): + prev_target_image = preview_target_image['background'] + else: + prev_target_image = preview_target_image - processor_ins = self.processors_manager.get_processor( - 'caption', preprocess_method) - if processor_ins is None: - sys_log = 'Current processor is illegal' - return gr.Textbox(), gr.Textbox( - ), self.component_names.system_log.format(sys_log) + if dataset_type == 'scepter_img2img': + if isinstance(preview_src_image, dict): + prev_src_image = preview_src_image['background'] + else: + prev_src_image = preview_src_image + + if isinstance(preview_src_mask_image, dict): + prev_src_mask = preview_src_mask_image['layers'][0] + else: + prev_src_mask = preview_src_mask_image + + else: + prev_src_image = None + prev_src_mask = None + + kwargs = { + 'src_image': None, + 'src_mask': None, + 'target_image': None, + 'caption': None, + 'preview_src_image': prev_src_image, + 'preview_src_mask': prev_src_mask, + 'preview_target_image': prev_target_image, + 'preview_caption': preview_caption, + 'use_preview': True, + 'use_local': use_local, + 'sys_prompt': sys_prompt, + 'max_new_tokens': max_new_tokens, + 'min_new_tokens': min_new_tokens, + 'num_beams': num_beams, + 'repetition_penalty': repetition_penalty, + 'temperature': temperature, + 'cache': self.cache + } is_flag, msg = processor_ins.load_model() if not is_flag: sys_log = f'Load processor failed: {msg}' return gr.Textbox(), self.component_names.system_log.format( sys_log) - dataset_type = create_dataset.get_trans_dataset_type(dataset_type) - - if isinstance(preview_target_image, dict): - prev_target_image = preview_target_image['background'] - else: - prev_target_image = preview_target_image - - if dataset_type == 'scepter_img2img': - if isinstance(preview_src_image, dict): - prev_src_image = preview_src_image['background'] - else: - prev_src_image = preview_src_image - - if isinstance(preview_src_mask_image, dict): - prev_src_mask = preview_src_mask_image['layers'][0] - else: - prev_src_mask = preview_src_mask_image - - else: - prev_src_image = None - prev_src_mask = None - - kwargs = { - 'src_image': None, - 'src_mask': None, - 'target_image': None, - 'caption': None, - 'preview_src_image': prev_src_image, - 'preview_src_mask': prev_src_mask, - 'preview_target_image': prev_target_image, - 'preview_caption': preview_caption, - 'use_preview': True, - 'use_local': use_local, - 'sys_prompt': sys_prompt, - 'max_new_tokens': max_new_tokens, - 'min_new_tokens': min_new_tokens, - 'num_beams': num_beams, - 'repetition_penalty': repetition_penalty, - 'temperature': temperature, - 'cache': self.cache - } response = processor_ins(**kwargs) if update_mode == 0: @@ -2156,7 +2433,9 @@ class DatasetGalleryUI(UIBase): self.num_beams, self.repetition_penalty, self.temperature, self.use_local, self.caption_update_mode, self.preview_src_image, self.preview_src_mask_image, - self.preview_taget_image, self.preview_caption + self.preview_taget_image, self.preview_caption, + self.preview_target_video, self.preprocess_checkbox, + self.upload_preprocess_video ], outputs=[self.preview_caption, self.sys_log]) @@ -2169,9 +2448,10 @@ class DatasetGalleryUI(UIBase): if dataset_ins is None or len(dataset_ins) < 1: return (gr.Gallery(), gr.Gallery(), gr.Image(), gr.Column(), gr.Row(visible=True), - gr.Column(), gr.Column(), gr.Image(), gr.Row(), '', - gr.Markdown(), gr.Textbox(), gr.CheckboxGroup(), - gr.Column(), gr.Row(), gr.Row(), gr.Column(), + gr.Column(), gr.Column(), gr.Image(), gr.Row(), + gr.Row(), '', gr.Markdown(), gr.Textbox(), + gr.CheckboxGroup(), gr.Column(), gr.Row(), + gr.Row(), gr.Column(), gr.Gallery(), gr.Textbox(), gr.Gallery(), gr.Image(), gr.Column(), gr.Dropdown(), gr.Dropdown(value=[]), @@ -2229,7 +2509,12 @@ class DatasetGalleryUI(UIBase): ret_edit_mask = gr.Image(visible=False) return (gr.Gallery(), gr.Gallery(), gr.Image(), gr.Column(), gr.Row(visible=True), gr.Column(visible=True), - gr.Column(visible=False), gr.Image(), gr.Row(), '', + gr.Column(visible=False), gr.Image(), gr.Row(), gr.Row(), + gr.Row(), gr.Row(), + gr.Row(visible=False) if dataset_type + == 'scepter_txt2vid' else gr.Row(visible=True), + gr.Row(visible=True) if dataset_type + == 'scepter_txt2vid' else gr.Row(visible=False), '', gr.Markdown(visible=False), gr.Textbox(), gr.CheckboxGroup(value=None), gr.Column(visible=False), gr.Row(visible=False), gr.Row(visible=False), @@ -2246,7 +2531,17 @@ class DatasetGalleryUI(UIBase): gr.Row(visible=False), gr.Column(visible=False), gr.Column(visible=True), gr.Image(), gr.Row(visible=True) if dataset_type - == 'scepter_img2img' else gr.Column(visible=False), '', + == 'scepter_img2img' else gr.Column(visible=False), + gr.Row(visible=True) if dataset_type + == 'scepter_txt2vid' else gr.Column(visible=False), + gr.Row(visible=False) if dataset_type + == 'scepter_txt2vid' else gr.Column(visible=True), + gr.Row(visible=False) if dataset_type + == 'scepter_txt2vid' else gr.Row(visible=True), + gr.Row(visible=False) if dataset_type + == 'scepter_txt2vid' else gr.Row(visible=True), + gr.Row(visible=True) if dataset_type + == 'scepter_txt2vid' else gr.Row(visible=False), '', gr.Markdown(visible=True) if dataset_type == 'scepter_img2img' else gr.Markdown(visible=False), gr.Textbox(value=''), gr.CheckboxGroup(), @@ -2263,11 +2558,12 @@ class DatasetGalleryUI(UIBase): if dataset_ins is None: return (gr.Gallery(), gr.Gallery(), gr.Image(), gr.Column(), gr.Row(visible=True), - gr.Column(), gr.Column(), gr.Image(), gr.Row(), '', + gr.Column(), gr.Column(), gr.Image(), gr.Row(), + gr.Row(), gr.Row(), gr.Row(), gr.Row(), gr.Row(), '', gr.Markdown(), gr.Textbox(), gr.CheckboxGroup(), - gr.Column(), gr.Row(), gr.Row(), gr.Column(), - gr.Gallery(), gr.Textbox(), gr.Gallery(), - gr.Image(), gr.Column(), gr.Dropdown(), + gr.Column(), gr.Row(), + gr.Row(), gr.Column(), gr.Gallery(), gr.Textbox(), + gr.Gallery(), gr.Image(), gr.Column(), gr.Dropdown(), gr.Dropdown(value=[]), self.component_names.system_log.format( self.component_names.illegal_blank_dataset)) @@ -2328,7 +2624,8 @@ class DatasetGalleryUI(UIBase): return (ret_gl_gallery, ret_src_gallery, ret_mask, ret_src_panel, gr.Row(visible=False), gr.Column(visible=False), gr.Column(visible=False), - gr.Image(), gr.Row(), '', gr.Markdown(visible=False), + gr.Image(), gr.Row(), gr.Row(), gr.Row(), gr.Row(), + gr.Row(), gr.Row(), '', gr.Markdown(visible=False), gr.Textbox(value=''), gr.CheckboxGroup(value=None), gr.Column(visible=False), gr.Row(visible=False), gr.Row(visible=False), gr.Column(visible=False), @@ -2362,6 +2659,16 @@ class DatasetGalleryUI(UIBase): self.upload_image, # gr.Row self.upload_src_image_panel, + # gr.Row + self.upload_src_video, + # gr.Row + self.upload_image_row, + # gr.Row + self.mode_select_add, + # gr.Row + self.mode_select_edit, + # gr.Row + self.mode_select_video, # gr.Markdown self.upload_image_info, # gr.Markdown @@ -2698,6 +3005,22 @@ class DatasetGalleryUI(UIBase): outputs=[self.upload_image_info], queue=False) + def video_upload(upload_video): + if isinstance(upload_video, dict): + video = upload_video['video'] + else: + video = upload_video + cap = cv2.VideoCapture(video) + w = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) + h = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) + cap.release() + return self.component_names.upload_image_info.format(h, w) + + self.upload_video.upload(video_upload, + inputs=[self.upload_video], + outputs=[self.upload_image_info], + queue=False) + def image_src_upload(upload_image): if isinstance(upload_image, dict): image = upload_image['image'] @@ -2744,7 +3067,7 @@ class DatasetGalleryUI(UIBase): ) def add_file(dataset_type, dataset_name, upload_image, - upload_src_image, upload_src_mask, caption): + upload_src_image, upload_src_mask, upload_video, caption): dataset_type = create_dataset.get_trans_dataset_type(dataset_type) dataset_ins = create_dataset.dataset_dict.get( dataset_type, {}).get(dataset_name, None) @@ -2770,7 +3093,7 @@ class DatasetGalleryUI(UIBase): if dataset_type == 'scepter_img2img' and src_image is not None and image is None: w, h = src_image.size image = Image.new('RGB', (w, h), (0, 0, 0)) - if image is None: + if image is None and dataset_type != 'scepter_txt2vid': return (gr.Image(value=None), gr.Image(value=None), gr.Image(value=None), gr.Text(value=''), '', '', gr.Text(value='view')) @@ -2783,10 +3106,15 @@ class DatasetGalleryUI(UIBase): return (gr.Image(value=None), gr.Image(value=None), gr.Image(value=None), gr.Text(value=''), '', '', gr.Text(value='view')) - dataset_ins.add_record(image, - caption, - src_image=src_image, - src_mask=src_mask) + + if dataset_type == 'scepter_txt2vid': + dataset_ins.add_record(upload_video, + caption) + else: + dataset_ins.add_record(image, + caption, + src_image=src_image, + src_mask=src_mask) return (gr.Image(value=None), gr.Image(value=None), gr.Image(value=None), gr.Text(value=''), '', '', gr.Text(value='view')) @@ -2796,7 +3124,8 @@ class DatasetGalleryUI(UIBase): create_dataset.dataset_type, create_dataset.dataset_name, self.upload_image, self.upload_src_image, - self.upload_src_mask, self.upload_caption + self.upload_src_mask, self.upload_video, + self.upload_caption ], outputs=[ self.upload_image, self.upload_src_image, @@ -2830,3 +3159,16 @@ class DatasetGalleryUI(UIBase): self.edit_caption ], queue=False) + + def get_trans_status(dataset_type, preprocess_name, preprocess_name_txt2vid): + if dataset_type == 'scepter_txt2vid': + choices = self.component_names.preprocess_choices_video + operation = preprocess_name_txt2vid[0] + hit_index = 1 + else: + choices = self.component_names.preprocess_choices + operation = preprocess_name[0] + hit_index = 2 + + reverse_status = {v: id for id, v in enumerate(choices)} + return reverse_status.get(operation, -1) == hit_index diff --git a/scepter/studio/preprocess/processors/caption_processors.py b/scepter/studio/preprocess/processors/caption_processors.py index b63f9e0..bd719e8 100644 --- a/scepter/studio/preprocess/processors/caption_processors.py +++ b/scepter/studio/preprocess/processors/caption_processors.py @@ -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 \ No newline at end of file diff --git a/scepter/studio/preprocess/utils/txt2vid_data_card.py b/scepter/studio/preprocess/utils/txt2vid_data_card.py new file mode 100644 index 0000000..9c3af36 --- /dev/null +++ b/scepter/studio/preprocess/utils/txt2vid_data_card.py @@ -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 diff --git a/scepter/studio/self_train/scripts/trainer.py b/scepter/studio/self_train/scripts/trainer.py index 4ca7924..d7f0a85 100644 --- a/scepter/studio/self_train/scripts/trainer.py +++ b/scepter/studio/self_train/scripts/trainer.py @@ -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) diff --git a/scepter/studio/self_train/self_train_ui/component_names.py b/scepter/studio/self_train/self_train_ui/component_names.py index d809d4b..84757f4 100644 --- a/scepter/studio/self_train/self_train_ui/component_names.py +++ b/scepter/studio/self_train/self_train_ui/component_names.py @@ -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 = '目前显存不足,训练失败!' diff --git a/scepter/studio/self_train/self_train_ui/model_ui.py b/scepter/studio/self_train/self_train_ui/model_ui.py index 7cac932..1b87e08 100644 --- a/scepter/studio/self_train/self_train_ui/model_ui.py +++ b/scepter/studio/self_train/self_train_ui/model_ui.py @@ -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, diff --git a/scepter/studio/self_train/self_train_ui/trainer_ui.py b/scepter/studio/self_train/self_train_ui/trainer_ui.py index 948d5da..c84d3d0 100644 --- a/scepter/studio/self_train/self_train_ui/trainer_ui.py +++ b/scepter/studio/self_train/self_train_ui/trainer_ui.py @@ -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) diff --git a/scepter/tools/webui.py b/scepter/tools/webui.py index be33021..7e2ba60 100644 --- a/scepter/tools/webui.py +++ b/scepter/tools/webui.py @@ -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) diff --git a/scepter/version.py b/scepter/version.py index 79b1b0a..8c43cf5 100644 --- a/scepter/version.py +++ b/scepter/version.py @@ -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]) diff --git a/scepter/workflow/config/ace_0.6b_1024_pro.yaml b/scepter/workflow/config/ace_0.6b_1024_pro.yaml new file mode 100644 index 0000000..2d99b72 --- /dev/null +++ b/scepter/workflow/config/ace_0.6b_1024_pro.yaml @@ -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 \ No newline at end of file diff --git a/scepter/workflow/config/ace_0.6b_1024_refiner_pro.yaml b/scepter/workflow/config/ace_0.6b_1024_refiner_pro.yaml new file mode 100644 index 0000000..d40b62d --- /dev/null +++ b/scepter/workflow/config/ace_0.6b_1024_refiner_pro.yaml @@ -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 \ No newline at end of file diff --git a/scepter/workflow/config/ace_0.6b_512_pro.yaml b/scepter/workflow/config/ace_0.6b_512_pro.yaml index e7c29f7..c5bc98e 100644 --- a/scepter/workflow/config/ace_0.6b_512_pro.yaml +++ b/scepter/workflow/config/ace_0.6b_512_pro.yaml @@ -39,7 +39,7 @@ DEFAULT_PARAS: # COND_STAGE_MODEL: FUNCTION: - - NAME: encode_list + - NAME: encode_list_of_list DTYPE: bfloat16 INPUT: ["PROMPT"] # diff --git a/scepter/workflow/config/flux1.0_dev_pro.yaml b/scepter/workflow/config/flux1.0_dev_pro.yaml index 3e1f5d9..8e2d0dc 100644 --- a/scepter/workflow/config/flux1.0_dev_pro.yaml +++ b/scepter/workflow/config/flux1.0_dev_pro.yaml @@ -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: diff --git a/scepter/workflow/config/flux1.0_schnell_pro.yaml b/scepter/workflow/config/flux1.0_schnell_pro.yaml index b86a958..175b8bf 100644 --- a/scepter/workflow/config/flux1.0_schnell_pro.yaml +++ b/scepter/workflow/config/flux1.0_schnell_pro.yaml @@ -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: diff --git a/scepter/workflow/config/scepter_workflow.yaml b/scepter/workflow/config/scepter_workflow.yaml index 1afc620..fbb2fa5 100644 --- a/scepter/workflow/config/scepter_workflow.yaml +++ b/scepter/workflow/config/scepter_workflow.yaml @@ -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" diff --git a/scepter/workflow/model_node.py b/scepter/workflow/model_node.py index 0e1b3f9..99c6555 100644 --- a/scepter/workflow/model_node.py +++ b/scepter/workflow/model_node.py @@ -94,10 +94,14 @@ class ModelNode: elif source == 'Local': cfg_new = copy.deepcopy(cfg) cfg_new.MODEL = cfg_new.MODEL_LOCAL + if hasattr(cfg_new, 'EFINER_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, 'EFINER_MODEL_HF'): + cfg_new.REFINER_MODEL = cfg_new.REFINER_MODEL_HF return cfg_new else: raise NotImplementedError(f'Unknown model source: {source}') diff --git a/tests/modules/test_diffusion_inference.py b/tests/modules/test_diffusion_inference.py index 7dd035c..1d2346d 100644 --- a/tests/modules/test_diffusion_inference.py +++ b/tests/modules/test_diffusion_inference.py @@ -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()