Update flux2 (#383)
This commit is contained in:
+15
-62
@@ -24,6 +24,7 @@ import pickle
|
||||
import random
|
||||
import shutil
|
||||
import sys
|
||||
from typing import List, NamedTuple, Optional, Union
|
||||
|
||||
import accelerate
|
||||
import diffusers
|
||||
@@ -33,9 +34,6 @@ import torch.nn.functional as F
|
||||
import torch.utils.checkpoint
|
||||
import torchvision.transforms.functional as TF
|
||||
import transformers
|
||||
from typing import NamedTuple, List, Optional, Union
|
||||
|
||||
|
||||
from accelerate import Accelerator
|
||||
from accelerate.logging import get_logger
|
||||
from accelerate.state import AcceleratorState
|
||||
@@ -60,7 +58,6 @@ from transformers.utils import ContextManagers
|
||||
|
||||
import datasets
|
||||
|
||||
|
||||
current_file_path = os.path.abspath(__file__)
|
||||
project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))]
|
||||
for project_root in project_roots:
|
||||
@@ -74,14 +71,14 @@ from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512,
|
||||
from videox_fun.data.dataset_image_video import (ImageVideoDataset,
|
||||
ImageVideoSampler,
|
||||
get_random_mask)
|
||||
from videox_fun.models import (AutoencoderKL, AutoencoderKLWan,
|
||||
Qwen2_5_VLForConditionalGeneration,
|
||||
Qwen2Tokenizer, QwenImageTransformer2DModel)
|
||||
from videox_fun.dist import set_multi_gpus_devices, shard_model
|
||||
from videox_fun.models import (CLIPImageProcessor, CLIPTextModel,
|
||||
from videox_fun.models import (AutoencoderKL, AutoencoderKLWan,
|
||||
CLIPImageProcessor, CLIPTextModel,
|
||||
CLIPTokenizer, CLIPVisionModelWithProjection,
|
||||
FluxTransformer2DModel, T5EncoderModel,
|
||||
T5TokenizerFast)
|
||||
FluxTransformer2DModel,
|
||||
Qwen2_5_VLForConditionalGeneration,
|
||||
Qwen2Tokenizer, QwenImageTransformer2DModel,
|
||||
T5EncoderModel, T5TokenizerFast)
|
||||
from videox_fun.pipeline import FluxPipeline
|
||||
from videox_fun.utils.discrete_sampler import DiscreteSampling
|
||||
from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid
|
||||
@@ -109,12 +106,6 @@ def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=Non
|
||||
u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator)
|
||||
t = 1 / (1 + torch.exp(-u)) * (high - low) + low
|
||||
return torch.clip(t.to(torch.int32), low, high - 1)
|
||||
|
||||
def _pack_latents(latents, batch_size, num_channels_latents, height, width):
|
||||
latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2)
|
||||
latents = latents.permute(0, 2, 4, 1, 3, 5)
|
||||
latents = latents.reshape(batch_size, (height // 2) * (width // 2), num_channels_latents * 4)
|
||||
return latents
|
||||
|
||||
def calculate_shift(
|
||||
image_seq_len,
|
||||
@@ -128,6 +119,12 @@ def calculate_shift(
|
||||
mu = image_seq_len * m + b
|
||||
return mu
|
||||
|
||||
def _pack_latents(latents, batch_size, num_channels_latents, height, width):
|
||||
latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2)
|
||||
latents = latents.permute(0, 2, 4, 1, 3, 5)
|
||||
latents = latents.reshape(batch_size, (height // 2) * (width // 2), num_channels_latents * 4)
|
||||
return latents
|
||||
|
||||
def _extract_masked_hidden(hidden_states: torch.Tensor, mask: torch.Tensor):
|
||||
bool_mask = mask.bool()
|
||||
valid_lengths = bool_mask.sum(dim=1)
|
||||
@@ -149,20 +146,6 @@ def _prepare_latent_image_ids(batch_size, height, width, device, dtype):
|
||||
|
||||
return latent_image_ids.to(device=device, dtype=dtype)
|
||||
|
||||
def _pack_latents(latents, batch_size, num_channels_latents, height, width):
|
||||
latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2)
|
||||
latents = latents.permute(0, 2, 4, 1, 3, 5)
|
||||
latents = latents.reshape(batch_size, (height // 2) * (width // 2), num_channels_latents * 4)
|
||||
return latents
|
||||
|
||||
def _extract_masked_hidden(hidden_states: torch.Tensor, mask: torch.Tensor):
|
||||
bool_mask = mask.bool()
|
||||
valid_lengths = bool_mask.sum(dim=1)
|
||||
selected = hidden_states[bool_mask]
|
||||
split_result = torch.split(selected, valid_lengths.tolist(), dim=0)
|
||||
|
||||
return split_result
|
||||
|
||||
def _get_t5_prompt_embeds(
|
||||
prompt = None,
|
||||
max_sequence_length = 512,
|
||||
@@ -645,12 +628,6 @@ def parse_args():
|
||||
default=[],
|
||||
help='Enter a list of trainable modules with lower learning rate'
|
||||
)
|
||||
parser.add_argument(
|
||||
'--tokenizer_max_length',
|
||||
type=int,
|
||||
default=1024,
|
||||
help='Max length of tokenizer'
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use_deepspeed", action="store_true", help="Whether or not to use deepspeed."
|
||||
)
|
||||
@@ -660,31 +637,6 @@ def parse_args():
|
||||
parser.add_argument(
|
||||
"--low_vram", action="store_true", help="Whether enable low_vram mode."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt_template_encode",
|
||||
type=str,
|
||||
default="<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n",
|
||||
help=(
|
||||
'The prompt template for text encoder.'
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt_template_encode_start_idx",
|
||||
type=int,
|
||||
default=34,
|
||||
help=(
|
||||
'The start idx for prompt template.'
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_mode",
|
||||
type=str,
|
||||
default="normal",
|
||||
help=(
|
||||
'The format of training data. Support `"normal"`'
|
||||
' (default), `"i2v"`.'
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--abnormal_norm_clip_start",
|
||||
type=int,
|
||||
@@ -1343,6 +1295,7 @@ def main():
|
||||
|
||||
if fsdp_stage != 0:
|
||||
from functools import partial
|
||||
|
||||
from videox_fun.dist import set_multi_gpus_devices, shard_model
|
||||
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.text_model.encoder.layers)
|
||||
text_encoder = shard_fn(text_encoder)
|
||||
@@ -1435,7 +1388,7 @@ def main():
|
||||
disable=not accelerator.is_local_main_process,
|
||||
)
|
||||
|
||||
if args.multi_stream and args.train_mode != "normal":
|
||||
if args.multi_stream:
|
||||
# create extra cuda streams to speedup inpaint vae computation
|
||||
vae_stream_1 = torch.cuda.Stream()
|
||||
vae_stream_2 = torch.cuda.Stream()
|
||||
|
||||
+12
-44
@@ -16,6 +16,7 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
|
||||
import argparse
|
||||
import copy
|
||||
import gc
|
||||
import logging
|
||||
import math
|
||||
@@ -24,7 +25,7 @@ import pickle
|
||||
import random
|
||||
import shutil
|
||||
import sys
|
||||
import copy
|
||||
from typing import List, NamedTuple, Optional, Union
|
||||
|
||||
import accelerate
|
||||
import diffusers
|
||||
@@ -34,9 +35,6 @@ import torch.nn.functional as F
|
||||
import torch.utils.checkpoint
|
||||
import torchvision.transforms.functional as TF
|
||||
import transformers
|
||||
from typing import NamedTuple, List, Optional, Union
|
||||
|
||||
|
||||
from accelerate import Accelerator
|
||||
from accelerate.logging import get_logger
|
||||
from accelerate.state import AcceleratorState
|
||||
@@ -73,14 +71,14 @@ from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512,
|
||||
from videox_fun.data.dataset_image_video import (ImageVideoDataset,
|
||||
ImageVideoSampler,
|
||||
get_random_mask)
|
||||
from videox_fun.models import (AutoencoderKL, AutoencoderKLWan,
|
||||
Qwen2_5_VLForConditionalGeneration,
|
||||
Qwen2Tokenizer, QwenImageTransformer2DModel)
|
||||
from videox_fun.dist import set_multi_gpus_devices, shard_model
|
||||
from videox_fun.models import (CLIPImageProcessor, CLIPTextModel,
|
||||
from videox_fun.models import (AutoencoderKL, AutoencoderKLWan,
|
||||
CLIPImageProcessor, CLIPTextModel,
|
||||
CLIPTokenizer, CLIPVisionModelWithProjection,
|
||||
FluxTransformer2DModel, T5EncoderModel,
|
||||
T5TokenizerFast)
|
||||
FluxTransformer2DModel,
|
||||
Qwen2_5_VLForConditionalGeneration,
|
||||
Qwen2Tokenizer, QwenImageTransformer2DModel,
|
||||
T5EncoderModel, T5TokenizerFast)
|
||||
from videox_fun.pipeline import FluxPipeline
|
||||
from videox_fun.utils.discrete_sampler import DiscreteSampling
|
||||
from videox_fun.utils.lora_utils import (create_network, merge_lora,
|
||||
@@ -642,12 +640,6 @@ def parse_args():
|
||||
)
|
||||
parser.add_argument("--save_state", action="store_true", help="Whether or not to save state.")
|
||||
|
||||
parser.add_argument(
|
||||
'--tokenizer_max_length',
|
||||
type=int,
|
||||
default=1024,
|
||||
help='Max length of tokenizer'
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use_deepspeed", action="store_true", help="Whether or not to use deepspeed."
|
||||
)
|
||||
@@ -657,31 +649,6 @@ def parse_args():
|
||||
parser.add_argument(
|
||||
"--low_vram", action="store_true", help="Whether enable low_vram mode."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt_template_encode",
|
||||
type=str,
|
||||
default="<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n",
|
||||
help=(
|
||||
'The prompt template for text encoder.'
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt_template_encode_start_idx",
|
||||
type=int,
|
||||
default=34,
|
||||
help=(
|
||||
'The start idx for prompt template.'
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_mode",
|
||||
type=str,
|
||||
default="normal",
|
||||
help=(
|
||||
'The format of training data. Support `"normal"`'
|
||||
' (default), `"inpaint"`.'
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--weighting_scheme",
|
||||
type=str,
|
||||
@@ -904,7 +871,8 @@ def main():
|
||||
|
||||
# Lora will work with this...
|
||||
if args.use_peft_lora:
|
||||
from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict
|
||||
from peft import (LoraConfig, get_peft_model_state_dict,
|
||||
inject_adapter_in_model)
|
||||
lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(","))
|
||||
transformer3d = inject_adapter_in_model(lora_config, transformer3d)
|
||||
|
||||
@@ -1310,7 +1278,7 @@ def main():
|
||||
from functools import partial
|
||||
|
||||
from videox_fun.dist import set_multi_gpus_devices, shard_model
|
||||
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=transformer3d.transformer_blocks)
|
||||
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=list(transformer3d.transformer_blocks) + list(transformer3d.single_transformer_blocks))
|
||||
transformer3d = shard_fn(transformer3d)
|
||||
|
||||
if fsdp_stage != 0 or zero_stage != 0:
|
||||
@@ -1474,7 +1442,7 @@ def main():
|
||||
disable=not accelerator.is_local_main_process,
|
||||
)
|
||||
|
||||
if args.multi_stream and args.train_mode != "normal":
|
||||
if args.multi_stream:
|
||||
# create extra cuda streams to speedup inpaint vae computation
|
||||
vae_stream_1 = torch.cuda.Stream()
|
||||
vae_stream_2 = torch.cuda.Stream()
|
||||
|
||||
@@ -0,0 +1,133 @@
|
||||
## Training Code
|
||||
|
||||
We can choose whether to use deepspeed or fsdp in flux2, which can save a lot of video memory.
|
||||
|
||||
Some parameters in the sh file can be confusing, and they are explained in this document:
|
||||
|
||||
- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the images at the center, but instead, it trains the entire images after grouping them into buckets based on resolution.
|
||||
- `random_hw_adapt` is used to enable automatic height and width scaling for images. When `random_hw_adapt` is enabled, the training images will have their height and width set to `image_sample_size` as the maximum and `512` as the minimum.
|
||||
- For example, when `random_hw_adapt` is enabled, `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024`
|
||||
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
|
||||
|
||||
When train model with multi machines, please set the params as follows:
|
||||
```sh
|
||||
export MASTER_ADDR="your master address"
|
||||
export MASTER_PORT=10086
|
||||
export WORLD_SIZE=1 # The number of machines
|
||||
export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8
|
||||
export RANK=0 # The rank of this machine
|
||||
|
||||
accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py
|
||||
```
|
||||
|
||||
Without deepspeed:
|
||||
|
||||
Training flux2 without DeepSpeed may result in insufficient GPU memory.
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-dev"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" scripts/flux2/train.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--train_batch_size=1 \
|
||||
--image_sample_size=1024 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--trainable_modules "."
|
||||
```
|
||||
|
||||
With Deepspeed Zero-2:
|
||||
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-dev"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/flux2/train.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--train_batch_size=1 \
|
||||
--image_sample_size=1024 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--trainable_modules "."
|
||||
```
|
||||
|
||||
With FSDP:
|
||||
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-dev"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap Flux2SingleTransformerBlock,Flux2TransformerBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/flux2/train.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--train_batch_size=1 \
|
||||
--image_sample_size=1024 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--trainable_modules "."
|
||||
```
|
||||
@@ -0,0 +1,136 @@
|
||||
## Lora Training Code
|
||||
|
||||
We can choose whether to use deepspeed or fsdp in flux, which can save a lot of video memory.
|
||||
|
||||
Some parameters in the sh file can be confusing, and they are explained in this document:
|
||||
|
||||
- `enable_bucket` is used to enable bucket training. When enabled, the model does not crop the images at the center, but instead, it trains the entire images after grouping them into buckets based on resolution.
|
||||
- `random_hw_adapt` is used to enable automatic height and width scaling for images. When `random_hw_adapt` is enabled, the training images will have their height and width set to `image_sample_size` as the maximum and `512` as the minimum.
|
||||
- For example, when `random_hw_adapt` is enabled, `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024`
|
||||
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
|
||||
- `target_name` represents the components/modules to which LoRA will be applied, separated by commas.
|
||||
- `use_peft_lora` indicates whether to use the PEFT module for adding LoRA. Using this module will be more memory-efficient.
|
||||
- `rank` means the dimension of the LoRA update matrices.
|
||||
- `network_alpha` means the scale of the LoRA update matrices.
|
||||
|
||||
When train model with multi machines, please set the params as follows:
|
||||
```sh
|
||||
export MASTER_ADDR="your master address"
|
||||
export MASTER_PORT=10086
|
||||
export WORLD_SIZE=1 # The number of machines
|
||||
export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8
|
||||
export RANK=0 # The rank of this machine
|
||||
|
||||
accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py
|
||||
```
|
||||
|
||||
Without deepspeed:
|
||||
|
||||
Training flux without DeepSpeed may result in insufficient GPU memory.
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-dev"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" scripts/flux/train_lora.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--train_batch_size=1 \
|
||||
--image_sample_size=1024 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=1e-04 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--enable_bucket \
|
||||
--rank=64 \
|
||||
--network_alpha=32 \
|
||||
--target_name="to_q,to_k,to_v,ff.0,ff.2,ff_context.0,ff_context.2" \
|
||||
--use_peft_lora \
|
||||
--uniform_sampling
|
||||
```
|
||||
|
||||
With Deepspeed Zero-2:
|
||||
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-dev"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json --deepspeed_multinode_launcher standard scripts/flux/train_lora.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--train_batch_size=1 \
|
||||
--image_sample_size=1024 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=1e-04 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--enable_bucket \
|
||||
--rank=64 \
|
||||
--network_alpha=32 \
|
||||
--target_name="to_q,to_k,to_v,ff.0,ff.2,ff_context.0,ff_context.2" \
|
||||
--use_peft_lora \
|
||||
--uniform_sampling
|
||||
```
|
||||
|
||||
With FSDP:
|
||||
|
||||
```sh
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-InP"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap Flux2SingleTransformerBlock,Flux2TransformerBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/flux/train_lora.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--train_batch_size=1 \
|
||||
--image_sample_size=1024 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=1e-04 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--enable_bucket \
|
||||
--uniform_sampling
|
||||
```
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,32 @@
|
||||
export MODEL_NAME="models/Diffusion_Transformer/FLUX.2-dev"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" scripts/flux2/train.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--train_batch_size=1 \
|
||||
--image_sample_size=1328 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=2e-05 \
|
||||
--lr_scheduler="constant_with_warmup" \
|
||||
--lr_warmup_steps=100 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--enable_bucket \
|
||||
--uniform_sampling \
|
||||
--trainable_modules "."
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,33 @@
|
||||
export MODEL_NAME="models/Diffusion_Transformer/FLUX.2-dev"
|
||||
export DATASET_NAME="datasets/internal_datasets/"
|
||||
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
|
||||
# export NCCL_IB_DISABLE=1
|
||||
# export NCCL_P2P_DISABLE=1
|
||||
NCCL_DEBUG=INFO
|
||||
|
||||
accelerate launch --mixed_precision="bf16" scripts/flux2/train_lora.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
--train_batch_size=1 \
|
||||
--image_sample_size=1328 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--dataloader_num_workers=8 \
|
||||
--num_train_epochs=100 \
|
||||
--checkpointing_steps=50 \
|
||||
--learning_rate=1e-04 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir_lora" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--vae_mini_batch=1 \
|
||||
--max_grad_norm=0.05 \
|
||||
--enable_bucket \
|
||||
--rank=64 \
|
||||
--network_alpha=32 \
|
||||
--target_name="to_q,to_k,to_v,ff.0,ff.2,ff_context.0,ff_context.2" \
|
||||
--use_peft_lora \
|
||||
--uniform_sampling
|
||||
@@ -1423,7 +1423,7 @@ def main():
|
||||
from functools import partial
|
||||
|
||||
from videox_fun.dist import set_multi_gpus_devices, shard_model
|
||||
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=transformer3d.transformer_blocks)
|
||||
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=list(transformer3d.transformer_blocks) + list(transformer3d.single_transformer_blocks))
|
||||
transformer3d = shard_fn(transformer3d)
|
||||
|
||||
if fsdp_stage != 0 or zero_stage != 0:
|
||||
|
||||
Reference in New Issue
Block a user