Fix bug in ZImage training (#407)
This commit is contained in:
@@ -106,7 +106,7 @@ export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# 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 ZImageTransformer2DModel --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/z_image/train.py \
|
||||
accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap ZImageTransformerBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/z_image/train.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
|
||||
@@ -112,7 +112,7 @@ export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
|
||||
# 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 ZImageTransformer2DModel --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/z_image/train_lora.py \
|
||||
accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap ZImageTransformerBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/z_image/train_lora.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_data_dir=$DATASET_NAME \
|
||||
--train_data_meta=$DATASET_META_NAME \
|
||||
|
||||
@@ -80,7 +80,7 @@ from videox_fun.models import (AutoencoderKL, AutoTokenizer,
|
||||
Qwen2Tokenizer, Qwen3ForCausalLM,
|
||||
QwenImageTransformer2DModel,
|
||||
ZImageTransformer2DModel)
|
||||
from videox_fun.pipeline import Flux2Pipeline
|
||||
from videox_fun.pipeline import ZImagePipeline
|
||||
from videox_fun.utils.discrete_sampler import DiscreteSampling
|
||||
from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid
|
||||
|
||||
@@ -201,7 +201,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, a
|
||||
subfolder="scheduler"
|
||||
)
|
||||
transformer3d = transformer3d.to("cpu")
|
||||
pipeline = Flux2Pipeline(
|
||||
pipeline = ZImagePipeline(
|
||||
vae=accelerator.unwrap_model(vae).to(weight_dtype),
|
||||
text_encoder=accelerator.unwrap_model(text_encoder),
|
||||
tokenizer=tokenizer,
|
||||
@@ -1423,7 +1423,6 @@ def main():
|
||||
indices = idx_sampling(bsz, generator=torch_rng, device=latents.device)
|
||||
indices = indices.long().cpu()
|
||||
|
||||
sigmas = np.linspace(1.0, 1 / args.train_sampling_steps, args.train_sampling_steps)
|
||||
image_seq_len = latents.shape[1]
|
||||
mu = calculate_shift(
|
||||
image_seq_len,
|
||||
@@ -1432,7 +1431,8 @@ def main():
|
||||
noise_scheduler.config.get("base_shift", 0.5),
|
||||
noise_scheduler.config.get("max_shift", 1.15),
|
||||
)
|
||||
noise_scheduler.set_timesteps(sigmas=sigmas, device=latents.device, mu=mu)
|
||||
noise_scheduler.sigma_min = 0.0
|
||||
noise_scheduler.set_timesteps(args.train_sampling_steps, device=latents.device, mu=mu)
|
||||
timesteps = noise_scheduler.timesteps[indices].to(device=latents.device)
|
||||
|
||||
def get_sigmas(timesteps, n_dim=4, dtype=torch.float32):
|
||||
@@ -1455,6 +1455,7 @@ def main():
|
||||
target = noise - latents
|
||||
|
||||
timesteps = (1000 - timesteps) / 1000
|
||||
|
||||
# Predict the noise residual
|
||||
with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device):
|
||||
noise_pred = transformer3d(
|
||||
@@ -1476,7 +1477,7 @@ def main():
|
||||
return final_loss
|
||||
|
||||
weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas)
|
||||
loss = custom_mse_loss(noise_pred.float(), target.float(), weighting.float())
|
||||
loss = custom_mse_loss(-noise_pred.float(), target.float(), weighting.float())
|
||||
loss = loss.mean()
|
||||
|
||||
# Gather the losses across all processes for logging (if we use distributed training).
|
||||
|
||||
@@ -80,8 +80,10 @@ from videox_fun.models import (AutoencoderKL, AutoTokenizer,
|
||||
Qwen2Tokenizer, Qwen3ForCausalLM,
|
||||
QwenImageTransformer2DModel,
|
||||
ZImageTransformer2DModel)
|
||||
from videox_fun.pipeline import Flux2Pipeline
|
||||
from videox_fun.pipeline import ZImagePipeline
|
||||
from videox_fun.utils.discrete_sampler import DiscreteSampling
|
||||
from videox_fun.utils.lora_utils import (create_network, merge_lora,
|
||||
unmerge_lora)
|
||||
from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid
|
||||
|
||||
if is_wandb_available():
|
||||
@@ -201,7 +203,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, a
|
||||
subfolder="scheduler"
|
||||
)
|
||||
transformer3d = transformer3d.to("cpu")
|
||||
pipeline = Flux2Pipeline(
|
||||
pipeline = ZImagePipeline(
|
||||
vae=accelerator.unwrap_model(vae).to(weight_dtype),
|
||||
text_encoder=accelerator.unwrap_model(text_encoder),
|
||||
tokenizer=tokenizer,
|
||||
|
||||
Reference in New Issue
Block a user