From ac114cc14285c8e0073a3e08e27525263d1264a7 Mon Sep 17 00:00:00 2001 From: Bubbliiiing <47347516+bubbliiiing@users.noreply.github.com> Date: Fri, 16 Jan 2026 17:21:25 +0800 Subject: [PATCH] Update Flux2 Training Code (#436) --- scripts/flux/train.py | 175 +++++++++++++------------ scripts/flux/train_lora.py | 152 +++++++++++----------- scripts/flux2/train.py | 164 +++++++++++------------ scripts/flux2/train_lora.py | 141 ++++++++++---------- scripts/flux2_fun/train_control.py | 180 +++++++++++++++----------- videox_fun/pipeline/pipeline_flux2.py | 1 - 6 files changed, 428 insertions(+), 385 deletions(-) diff --git a/scripts/flux/train.py b/scripts/flux/train.py index 1bb94d3..f6a5783 100644 --- a/scripts/flux/train.py +++ b/scripts/flux/train.py @@ -244,60 +244,69 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, text_encoder_2, tokenizer, tokenizer_2, transformer3d, network, args, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, text_encoder_2, tokenizer, tokenizer_2, transformer3d, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = FluxPipeline( + vae=vae, + text_encoder=text_encoder, + text_encoder_2=text_encoder_2, + tokenizer=tokenizer, + tokenizer_2=tokenizer_2, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - transformer3d_val = FluxTransformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, - low_cpu_mem_usage=True, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - pipeline = FluxPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - text_encoder_2=accelerator.unwrap_model(text_encoder_2), - tokenizer=tokenizer, - tokenizer_2=tokenizer_2, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): sample = pipeline( - args.validation_prompts[i], - negative_prompt = "bad detailed", + prompt = args.validation_prompts[i], height = args.image_sample_size, width = args.image_sample_size, - generator = generator + generator = generator, + num_inference_steps = 20, ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -1632,28 +1641,26 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - text_encoder_2, - tokenizer, - tokenizer_2, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + text_encoder_2, + tokenizer, + tokenizer_2, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1661,28 +1668,26 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - text_encoder_2, - tokenizer, - tokenizer_2, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + text_encoder_2, + tokenizer, + tokenizer_2, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/flux/train_lora.py b/scripts/flux/train_lora.py index 5eb45e3..c6dae26 100644 --- a/scripts/flux/train_lora.py +++ b/scripts/flux/train_lora.py @@ -249,61 +249,69 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, text_encoder_2, tokenizer, tokenizer_2, transformer3d, network, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = FluxPipeline( + vae=vae, + text_encoder=text_encoder, + text_encoder_2=text_encoder_2, + tokenizer=tokenizer, + tokenizer_2=tokenizer_2, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - transformer3d_val = FluxTransformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, - low_cpu_mem_usage=True, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - pipeline = FluxPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - text_encoder_2=accelerator.unwrap_model(text_encoder_2), - tokenizer=tokenizer, - tokenizer_2=tokenizer_2, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): sample = pipeline( - args.validation_prompts[i], - negative_prompt = "bad detailed", + prompt = args.validation_prompts[i], height = args.image_sample_size, width = args.image_sample_size, - generator = generator + generator = generator, + num_inference_steps = 20, ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -1700,21 +1708,20 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - text_encoder_2, - tokenizer, - tokenizer_2, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + text_encoder_2, + tokenizer, + tokenizer_2, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1722,21 +1729,20 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - text_encoder_2, - tokenizer, - tokenizer_2, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + text_encoder_2, + tokenizer, + tokenizer_2, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/flux2/train.py b/scripts/flux2/train.py index 23455e4..7c9a423 100644 --- a/scripts/flux2/train.py +++ b/scripts/flux2/train.py @@ -313,58 +313,67 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = Flux2Pipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - transformer3d_val = Flux2Transformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, - low_cpu_mem_usage=True, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - pipeline = Flux2Pipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): sample = pipeline( - args.validation_prompts[i], - negative_prompt = "bad detailed", + prompt = args.validation_prompts[i], height = args.image_sample_size, width = args.image_sample_size, - generator = generator + generator = generator, + num_inference_steps = 20, ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -1712,26 +1721,24 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1739,27 +1746,24 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - tokenizer_2, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/flux2/train_lora.py b/scripts/flux2/train_lora.py index 1b42963..2469bfe 100644 --- a/scripts/flux2/train_lora.py +++ b/scripts/flux2/train_lora.py @@ -318,59 +318,67 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = Flux2Pipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - transformer3d_val = Flux2Transformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, - low_cpu_mem_usage=True, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - pipeline = Flux2Pipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): sample = pipeline( - args.validation_prompts[i], - negative_prompt = "bad detailed", + prompt = args.validation_prompts[i], height = args.image_sample_size, width = args.image_sample_size, - generator = generator + generator = generator, + num_inference_steps = 20, ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -1690,19 +1698,18 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1710,20 +1717,18 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - tokenizer_2, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/flux2_fun/train_control.py b/scripts/flux2_fun/train_control.py index c91df69..7ea5e49 100644 --- a/scripts/flux2_fun/train_control.py +++ b/scripts/flux2_fun/train_control.py @@ -83,8 +83,10 @@ from videox_fun.models import (AutoencoderKLFlux2, AutoProcessor, PixtralProcessor) from videox_fun.pipeline import Flux2ControlPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid -from videox_fun.utils.utils_yolo import ObjectInstanceDetector +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) + if is_wandb_available(): import wandb @@ -317,55 +319,72 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + pipeline = Flux2ControlPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - transformer3d_val = Flux2ControlTransformer2DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, low_cpu_mem_usage=True, - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - transformer3d = transformer3d.to("cpu") - pipeline = Flux2ControlPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + for i in range(len(args.validation_prompts)): + control_image = Image.open(args.validation_paths[i]) + width, height = control_image.width, control_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + control_image = get_image_latent(control_image, sample_size=(height, width))[:, :, 0] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): sample = pipeline( - args.validation_prompts[i], - negative_prompt = "bad detailed", - height = args.image_sample_size, - width = args.image_sample_size, - generator = generator + prompt = args.validation_prompts[i], + height = height, + width = width, + generator = generator, + num_inference_steps = 20, + control_context_scale = 0.90, + control_image = control_image, ).images os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + image = sample[0].save( + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.jpg" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - transformer3d = transformer3d.to(accelerator.device) + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - transformer3d = transformer3d.to(accelerator.device) + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -424,6 +443,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -1823,25 +1849,24 @@ def main(): transformer3d.requires_grad_(True) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1849,25 +1874,24 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/videox_fun/pipeline/pipeline_flux2.py b/videox_fun/pipeline/pipeline_flux2.py index 26a8741..536924d 100644 --- a/videox_fun/pipeline/pipeline_flux2.py +++ b/videox_fun/pipeline/pipeline_flux2.py @@ -878,7 +878,6 @@ class Flux2Pipeline(DiffusionPipeline): if output_type == "latent": image = latents else: - torch.save({"pred": latents}, "pred_d.pt") latents = self._unpack_latents_with_ids(latents, latent_ids) latents_bn_mean = self.vae.bn.running_mean.view(1, -1, 1, 1).to(latents.device, latents.dtype)