Update flux2 (#383)

This commit is contained in:
Bubbliiiing
2025-11-28 10:31:06 +08:00
committed by GitHub
parent 5794017c00
commit 200e1f3224
27 changed files with 7168 additions and 265 deletions
+15 -62
View File
@@ -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
View File
@@ -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()
+133
View File
@@ -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 "."
```
+136
View File
@@ -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
+32
View File
@@ -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
+33
View File
@@ -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
+1 -1
View File
@@ -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: