Update Wan 2.2 5b (#278)

This commit is contained in:
Bubbliiiing
2025-08-11 14:59:34 +08:00
committed by GitHub
parent eb915f0da3
commit 24a5eed03b
16 changed files with 2510 additions and 90 deletions
+52 -2
View File
@@ -19,9 +19,9 @@ Some parameters in the sh file can be confusing, and they are explained in this
- At 768x768 resolution, the number of video frames is 21 (~= 512 * 512 * 49 / 768 / 768).
- At 1024x1024 resolution, the number of video frames is 9 (~= 512 * 512 * 49 / 1024 / 1024).
- These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes.
- `train_mode` is used to specify the training mode, which can be either normal or i2v. Since Wan uses the inpaint model to achieve image-to-video generation, the default is set to inpaint mode. If you only wish to achieve text-to-video generation, you can remove this line, and it will default to the text-to-video mode.
- `train_mode` is used to specify the training mode, which can be either normal, i2v or ti2v. The t2v is used for 14B T2V model. The i2v is used for 14B I2V model. The ti2v is used in 5B TI2V model.
- `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.
- `boundary_type`: The Wan2.2 series includes two distinct models that handle different noise levels, specified via the `boundary_type` parameter. `low`: Corresponds to the **low noise model** (low_noise_model). `high`: Corresponds to the **high noise model**. (high_noise_model)
- `boundary_type`: The Wan2.2 series includes two distinct models that handle different noise levels, specified via the `boundary_type` parameter. `low`: Corresponds to the **low noise model** (low_noise_model). `high`: Corresponds to the **high noise model**. (high_noise_model). `full`: Corresponds to the ti2v 5B model (single mode).
Wan2.2 T2V without deepspeed:
@@ -174,6 +174,7 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag
--uniform_sampling \
--low_vram \
--use_deepspeed \
--boundary_type="low" \
--train_mode="normal" \
--trainable_modules "."
```
@@ -222,6 +223,55 @@ accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TR
--enable_bucket \
--uniform_sampling \
--low_vram \
--boundary_type="low" \
--train_mode="normal" \
--trainable_modules "."
```
If you want to train 5B Wan2.2 TI2V model, please set config to `config/wan2.2/wan_civitai_5b.yaml`, set train_mode to `ti2v` and set boundary_type to `full`. Training shell command is as follows:
```sh
export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-TI2V-5B"
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/wan2.2/train.py \
--config_path="config/wan2.2/wan_civitai_5b.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=1024 \
--video_sample_size=256 \
--token_sample_size=512 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--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 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--low_vram \
--boundary_type="full" \
--train_mode="ti2v" \
--trainable_modules "."
```
+47 -3
View File
@@ -17,9 +17,9 @@ Some parameters in the sh file can be confusing, and they are explained in this
- At 768x768 resolution, the number of video frames is 21 (~= 512 * 512 * 49 / 768 / 768).
- At 1024x1024 resolution, the number of video frames is 9 (~= 512 * 512 * 49 / 1024 / 1024).
- These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes.
- `train_mode` is used to specify the training mode, which can be either normal or i2v. Since Wan uses the inpaint model to achieve image-to-video generation, the default is set to inpaint mode. If you only wish to achieve text-to-video generation, you can remove this line, and it will default to the text-to-video mode.
- `train_mode` is used to specify the training mode, which can be either normal, i2v or ti2v. The t2v is used for 14B T2V model. The i2v is used for 14B I2V model. The ti2v is used in 5B TI2V model.
- `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.
- `boundary_type`: The Wan2.2 series includes two distinct models that handle different noise levels, specified via the `boundary_type` parameter. `low`: Corresponds to the **low noise model** (low_noise_model). `high`: Corresponds to the **high noise model**. (high_noise_model)
- `boundary_type`: The Wan2.2 series includes two distinct models that handle different noise levels, specified via the `boundary_type` parameter. `low`: Corresponds to the **low noise model** (low_noise_model). `high`: Corresponds to the **high noise model**. (high_noise_model). `full`: Corresponds to the ti2v 5B model (single mode).
Wan2.2 T2V without deepspeed:
@@ -211,6 +211,50 @@ accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TR
--uniform_sampling \
--boundary_type="low" \
--train_mode="normal" \
--use_deepspeed \
--low_vram
```
If you want to train 5B Wan2.2 TI2V model, please set config to `config/wan2.2/wan_civitai_5b.yaml`, set train_mode to `ti2v` and set boundary_type to `full`. Training shell command is as follows:
```sh
export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-T2V-A14B"
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/wan2.2/train_lora.py \
--config_path="config/wan2.2/wan_civitai_5b.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--image_sample_size=1024 \
--video_sample_size=256 \
--token_sample_size=512 \
--video_sample_stride=2 \
--video_sample_n_frames=81 \
--train_batch_size=1 \
--video_repeat=1 \
--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 \
--random_hw_adapt \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--boundary_type="full" \
--train_mode="ti2v" \
--low_vram
```
+43 -13
View File
@@ -71,7 +71,7 @@ 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 (AutoencoderKLWan, WanT5EncoderModel,
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, WanT5EncoderModel,
Wan2_2Transformer3DModel)
from videox_fun.pipeline import WanPipeline, WanI2VPipeline
from videox_fun.utils.discrete_sampler import DiscreteSampling
@@ -892,15 +892,21 @@ def main():
)
text_encoder = text_encoder.eval()
# Get Vae
vae = AutoencoderKLWan.from_pretrained(
Choosen_AutoencoderKL = {
"AutoencoderKLWan": AutoencoderKLWan,
"AutoencoderKLWan3_8": AutoencoderKLWan3_8
}[config['vae_kwargs'].get('vae_type', 'AutoencoderKLWan')]
vae = Choosen_AutoencoderKL.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
)
vae.eval()
# Get Transformer
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') \
if args.boundary_type == "low" else config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
if args.boundary_type == "low" or args.boundary_type == "full":
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
else:
sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
transformer3d = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
@@ -1135,6 +1141,7 @@ def main():
# Get the training dataset
sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio
spatial_compression_ratio = vae.config.spatial_compression_ratio
if args.fix_sample_size is not None and args.enable_bucket:
args.video_sample_size = max(max(args.fix_sample_size), args.video_sample_size)
@@ -1228,7 +1235,7 @@ def main():
aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()}
if args.fix_sample_size is not None:
fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size]
fix_sample_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in args.fix_sample_size]
elif args.random_ratio_crop:
if rng is None:
random_sample_size = aspect_ratio_random_crop_sample_size[
@@ -1238,10 +1245,10 @@ def main():
random_sample_size = aspect_ratio_random_crop_sample_size[
rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
random_sample_size = [int(x / 16) * 16 for x in random_sample_size]
random_sample_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in random_sample_size]
else:
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 16) * 16 for x in closest_size]
closest_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in closest_size]
for example in examples:
if args.fix_sample_size is not None:
@@ -1485,7 +1492,8 @@ def main():
split_timesteps = args.train_sampling_steps * boundary
differences = torch.abs(noise_scheduler.timesteps - split_timesteps)
closest_index = torch.argmin(differences).item()
print(f"The boundary is {boundary} and the boundary_type is {args.boundary_type}. The closest_index we calculate is {closest_index}")
if args.boundary_type == "high" or args.boundary_type == "low":
print(f"The boundary is {boundary} and the boundary_type is {args.boundary_type}. The closest_index we calculate is {closest_index}")
if args.boundary_type == "high":
start_num_idx = 0
train_sampling_steps = closest_index
@@ -1653,16 +1661,23 @@ def main():
)
mask = mask.view(mask.shape[0], mask.shape[2] // 4, 4, mask.shape[3], mask.shape[4])
mask = mask.transpose(1, 2)
mask = resize_mask(1 - mask, latents)
if args.train_mode != "ti2v":
mask = resize_mask(1 - mask, latents)
else:
mask = F.interpolate(mask[:, :1], size=latents.size()[-3:], mode='trilinear', align_corners=True).to(accelerator.device, weight_dtype)
# Encode inpaint latents.
mask_latents = _batch_encode_vae(mask_pixel_values)
if vae_stream_2 is not None:
torch.cuda.current_stream().wait_stream(vae_stream_2)
inpaint_latents = torch.concat([mask, mask_latents], dim=1)
inpaint_latents = t2v_flag[:, None, None, None, None] * inpaint_latents
if args.train_mode != "ti2v":
inpaint_latents = torch.concat([mask, mask_latents], dim=1)
inpaint_latents = t2v_flag[:, None, None, None, None] * inpaint_latents
else:
inpaint_latents = mask_latents
# wait for latents = vae.encode(pixel_values) to complete
if vae_stream_1 is not None:
torch.cuda.current_stream().wait_stream(vae_stream_1)
@@ -1742,6 +1757,21 @@ def main():
target_shape[1]
)
if args.train_mode == "ti2v":
if rng is None:
t2v_in_ti2v = np.random.choice([0, 1], p = [0.50, 0.50])
else:
t2v_in_ti2v = rng.choice([0, 1], p = [0.50, 0.50])
mask_bs = mask.size()[0]
if t2v_in_ti2v:
noisy_latents = (1 - mask) * inpaint_latents + mask * noisy_latents
temp_ts = (mask[:, 0, :, ::2, ::2] * timesteps[:, None, None, None]).flatten(1)
timesteps = torch.cat([temp_ts, temp_ts.new_ones(mask_bs, seq_len - temp_ts.size(1)) * timesteps[:, None,]], dim = 1)
else:
timesteps = mask.new_ones(mask_bs, seq_len) * timesteps[:, None,]
# Predict the noise residual
with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device):
noise_pred = transformer3d(
@@ -1749,7 +1779,7 @@ def main():
context=prompt_embeds,
t=timesteps,
seq_len=seq_len,
y=inpaint_latents if args.train_mode != "normal" else None,
y=inpaint_latents if args.train_mode != "normal" and args.train_mode != "ti2v" else None,
)
def custom_mse_loss(noise_pred, target, weighting=None, threshold=50):
+45 -14
View File
@@ -68,7 +68,7 @@ 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 (AutoencoderKLWan, WanT5EncoderModel,
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, WanT5EncoderModel,
Wan2_2Transformer3DModel)
from videox_fun.pipeline import Wan2_2Pipeline, Wan2_2I2VPipeline
from videox_fun.utils.discrete_sampler import DiscreteSampling
@@ -891,15 +891,21 @@ def main():
)
text_encoder = text_encoder.eval()
# Get Vae
vae = AutoencoderKLWan.from_pretrained(
Choosen_AutoencoderKL = {
"AutoencoderKLWan": AutoencoderKLWan,
"AutoencoderKLWan3_8": AutoencoderKLWan3_8
}[config['vae_kwargs'].get('vae_type', 'AutoencoderKLWan')]
vae = Choosen_AutoencoderKL.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')),
additional_kwargs=OmegaConf.to_container(config['vae_kwargs']),
)
vae.eval()
# Get Transformer
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') \
if args.boundary_type == "low" else config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
if args.boundary_type == "low" or args.boundary_type == "full":
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
else:
sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
transformer3d = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
@@ -1070,6 +1076,7 @@ def main():
# Get the training dataset
sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio
spatial_compression_ratio = vae.config.spatial_compression_ratio
if args.fix_sample_size is not None and args.enable_bucket:
args.video_sample_size = max(max(args.fix_sample_size), args.video_sample_size)
@@ -1164,7 +1171,7 @@ def main():
aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()}
if args.fix_sample_size is not None:
fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size]
fix_sample_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in args.fix_sample_size]
elif args.random_ratio_crop:
if rng is None:
random_sample_size = aspect_ratio_random_crop_sample_size[
@@ -1174,10 +1181,10 @@ def main():
random_sample_size = aspect_ratio_random_crop_sample_size[
rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB)
]
random_sample_size = [int(x / 16) * 16 for x in random_sample_size]
random_sample_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in random_sample_size]
else:
closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size)
closest_size = [int(x / 16) * 16 for x in closest_size]
closest_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in closest_size]
for example in examples:
if args.fix_sample_size is not None:
@@ -1492,7 +1499,8 @@ def main():
split_timesteps = args.train_sampling_steps * boundary
differences = torch.abs(noise_scheduler.timesteps - split_timesteps)
closest_index = torch.argmin(differences).item()
print(f"The boundary is {boundary} and the boundary_type is {args.boundary_type}. The closest_index we calculate is {closest_index}")
if args.boundary_type == "high" or args.boundary_type == "low":
print(f"The boundary is {boundary} and the boundary_type is {args.boundary_type}. The closest_index we calculate is {closest_index}")
if args.boundary_type == "high":
start_num_idx = 0
train_sampling_steps = closest_index
@@ -1659,16 +1667,23 @@ def main():
)
mask = mask.view(mask.shape[0], mask.shape[2] // 4, 4, mask.shape[3], mask.shape[4])
mask = mask.transpose(1, 2)
mask = resize_mask(1 - mask, latents)
if args.train_mode != "ti2v":
mask = resize_mask(1 - mask, latents)
else:
mask = F.interpolate(mask[:, :1], size=latents.size()[-3:], mode='trilinear', align_corners=True).to(accelerator.device, weight_dtype)
# Encode inpaint latents.
mask_latents = _batch_encode_vae(mask_pixel_values)
if vae_stream_2 is not None:
torch.cuda.current_stream().wait_stream(vae_stream_2)
inpaint_latents = torch.concat([mask, mask_latents], dim=1)
inpaint_latents = t2v_flag[:, None, None, None, None] * inpaint_latents
if args.train_mode != "ti2v":
inpaint_latents = torch.concat([mask, mask_latents], dim=1)
inpaint_latents = t2v_flag[:, None, None, None, None] * inpaint_latents
else:
inpaint_latents = mask_latents
# wait for latents = vae.encode(pixel_values) to complete
if vae_stream_1 is not None:
torch.cuda.current_stream().wait_stream(vae_stream_1)
@@ -1747,6 +1762,22 @@ def main():
(accelerator.unwrap_model(transformer3d).config.patch_size[1] * accelerator.unwrap_model(transformer3d).config.patch_size[2]) *
target_shape[1]
)
if args.train_mode == "ti2v":
if rng is None:
t2v_in_ti2v = np.random.choice([0, 1], p = [0.50, 0.50])
else:
t2v_in_ti2v = rng.choice([0, 1], p = [0.50, 0.50])
mask_bs = mask.size()[0]
if t2v_in_ti2v:
noisy_latents = (1 - mask) * inpaint_latents + mask * noisy_latents
temp_ts = (mask[:, 0, :, ::2, ::2] * timesteps[:, None, None, None]).flatten(1)
timesteps = torch.cat([temp_ts, temp_ts.new_ones(mask_bs, seq_len - temp_ts.size(1)) * timesteps[:, None,]], dim = 1)
else:
timesteps = mask.new_ones(mask_bs, seq_len) * timesteps[:, None,]
# Predict the noise residual
with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device):
noise_pred = transformer3d(
@@ -1754,9 +1785,9 @@ def main():
context=prompt_embeds,
t=timesteps,
seq_len=seq_len,
y=inpaint_latents if args.train_mode != "normal" else None,
y=inpaint_latents if args.train_mode != "normal" and args.train_mode != "ti2v" else None,
)
def custom_mse_loss(noise_pred, target, weighting=None, threshold=50):
noise_pred = noise_pred.float()
target = target.float()