From 47412d24407e3573452412d4b4de491f761d9803 Mon Sep 17 00:00:00 2001 From: bubbliiiing <3323290568@qq.com> Date: Mon, 5 Jan 2026 20:28:36 +0800 Subject: [PATCH] delete useless code --- scripts/z_image/train.py | 5 +--- scripts/z_image/train_lora.py | 1 - scripts/z_image/train_lora_omni.py | 11 ++++----- scripts/z_image/train_omni.py | 9 +++---- scripts/z_image_fun/train_control.py | 26 +++++++++----------- scripts/z_image_fun/train_control_distill.py | 20 ++++++--------- 6 files changed, 28 insertions(+), 44 deletions(-) diff --git a/scripts/z_image/train.py b/scripts/z_image/train.py index a0799c9..7aeeb28 100644 --- a/scripts/z_image/train.py +++ b/scripts/z_image/train.py @@ -187,7 +187,7 @@ 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... ") @@ -1562,7 +1562,6 @@ def main(): text_encoder, tokenizer, transformer3d, - network, args, accelerator, weight_dtype, @@ -1588,9 +1587,7 @@ def main(): vae, text_encoder, tokenizer, - tokenizer_2, transformer3d, - network, args, accelerator, weight_dtype, diff --git a/scripts/z_image/train_lora.py b/scripts/z_image/train_lora.py index 9782be6..2242878 100644 --- a/scripts/z_image/train_lora.py +++ b/scripts/z_image/train_lora.py @@ -1563,7 +1563,6 @@ def main(): vae, text_encoder, tokenizer, - tokenizer_2, transformer3d, network, args, diff --git a/scripts/z_image/train_lora_omni.py b/scripts/z_image/train_lora_omni.py index a3c5692..bf5103f 100644 --- a/scripts/z_image/train_lora_omni.py +++ b/scripts/z_image/train_lora_omni.py @@ -77,12 +77,12 @@ from videox_fun.data.dataset_image_video import (ImageVideoControlDataset, from videox_fun.dist import set_multi_gpus_devices, shard_model from videox_fun.models import (AutoencoderKL, AutoProcessor, AutoTokenizer, CLIPImageProcessor, - CLIPVisionModelWithProjection, - Qwen2Tokenizer, Qwen3ForCausalLM, - QwenImageTransformer2DModel, Siglip2VisionModel, + CLIPVisionModelWithProjection, Qwen2Tokenizer, + Qwen3ForCausalLM, QwenImageTransformer2DModel, + Siglip2VisionModel, ZImageOmniTransformer2DModel) from videox_fun.models.flux2_image_processor import Flux2ImageProcessor -from videox_fun.pipeline import Flux2Pipeline +from videox_fun.pipeline import ZImageOmniPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (create_network, merge_lora, unmerge_lora) @@ -216,7 +216,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, a subfolder="scheduler" ) transformer3d = transformer3d.to("cpu") - pipeline = Flux2Pipeline( + pipeline = ZImageOmniPipeline( vae=accelerator.unwrap_model(vae).to(weight_dtype), text_encoder=accelerator.unwrap_model(text_encoder), tokenizer=tokenizer, @@ -1672,7 +1672,6 @@ def main(): vae, text_encoder, tokenizer, - tokenizer_2, transformer3d, network, args, diff --git a/scripts/z_image/train_omni.py b/scripts/z_image/train_omni.py index 7a76691..9cea7b3 100644 --- a/scripts/z_image/train_omni.py +++ b/scripts/z_image/train_omni.py @@ -80,7 +80,7 @@ from videox_fun.models import (AutoencoderKL, AutoProcessor, AutoTokenizer, QwenImageTransformer2DModel, Siglip2VisionModel, ZImageOmniTransformer2DModel) from videox_fun.models.flux2_image_processor import Flux2ImageProcessor -from videox_fun.pipeline import Flux2Pipeline +from videox_fun.pipeline import ZImageOmniPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid @@ -199,7 +199,7 @@ 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... ") @@ -213,7 +213,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, a subfolder="scheduler" ) transformer3d = transformer3d.to("cpu") - pipeline = Flux2Pipeline( + pipeline = ZImageOmniPipeline( vae=accelerator.unwrap_model(vae).to(weight_dtype), text_encoder=accelerator.unwrap_model(text_encoder), tokenizer=tokenizer, @@ -1698,7 +1698,6 @@ def main(): text_encoder, tokenizer, transformer3d, - network, args, accelerator, weight_dtype, @@ -1724,9 +1723,7 @@ def main(): vae, text_encoder, tokenizer, - tokenizer_2, transformer3d, - network, args, accelerator, weight_dtype, diff --git a/scripts/z_image_fun/train_control.py b/scripts/z_image_fun/train_control.py index 484f625..cea7c6f 100644 --- a/scripts/z_image_fun/train_control.py +++ b/scripts/z_image_fun/train_control.py @@ -81,9 +81,8 @@ from videox_fun.models import (AutoencoderKL, AutoProcessor, AutoTokenizer, CLIPImageProcessor, CLIPVisionModelWithProjection, Qwen2Tokenizer, Qwen3ForCausalLM, QwenImageTransformer2DModel, - ZImageControlTransformer2DModel, - ZImageTransformer2DModel) -from videox_fun.pipeline import Flux2Pipeline + ZImageControlTransformer2DModel) +from videox_fun.pipeline import ZImageControlPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid @@ -190,11 +189,11 @@ 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... ") - transformer3d_val = ZImageTransformer2DModel.from_pretrained( + transformer3d_val = ZImageControlTransformer2DModel.from_pretrained( args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, low_cpu_mem_usage=True, ).to(weight_dtype) @@ -204,7 +203,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, a subfolder="scheduler" ) transformer3d = transformer3d.to("cpu") - pipeline = Flux2Pipeline( + pipeline = ZImageControlPipeline( vae=accelerator.unwrap_model(vae).to(weight_dtype), text_encoder=accelerator.unwrap_model(text_encoder), tokenizer=tokenizer, @@ -892,13 +891,13 @@ def main(): if zero_stage == 3: raise NotImplementedError("FSDP does not support EMA.") - ema_transformer3d = ZImageTransformer2DModel.from_pretrained( + ema_transformer3d = ZImageControlTransformer2DModel.from_pretrained( args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, ).to(weight_dtype) - ema_transformer3d = EMAModel(ema_transformer3d.parameters(), model_cls=ZImageTransformer2DModel, model_config=ema_transformer3d.config) + ema_transformer3d = EMAModel(ema_transformer3d.parameters(), model_cls=ZImageControlTransformer2DModel, model_config=ema_transformer3d.config) # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): @@ -960,11 +959,11 @@ def main(): def load_model_hook(models, input_dir): if args.use_ema: ema_path = os.path.join(input_dir, "transformer_ema") - _, ema_kwargs = ZImageTransformer2DModel.load_config(ema_path, return_unused_kwargs=True) - load_model = ZImageTransformer2DModel.from_pretrained( + _, ema_kwargs = ZImageControlTransformer2DModel.load_config(ema_path, return_unused_kwargs=True) + load_model = ZImageControlTransformer2DModel.from_pretrained( input_dir, subfolder="transformer_ema", ) - load_model = EMAModel(load_model.parameters(), model_cls=ZImageTransformer2DModel, model_config=load_model.config) + load_model = EMAModel(load_model.parameters(), model_cls=ZImageControlTransformer2DModel, model_config=load_model.config) load_model.load_state_dict(ema_kwargs) ema_transformer3d.load_state_dict(load_model.state_dict()) @@ -976,7 +975,7 @@ def main(): model = models.pop() # load diffusers style into model - load_model = ZImageTransformer2DModel.from_pretrained( + load_model = ZImageControlTransformer2DModel.from_pretrained( input_dir, subfolder="transformer" ) model.register_to_config(**load_model.config) @@ -1676,7 +1675,6 @@ def main(): text_encoder, tokenizer, transformer3d, - network, args, accelerator, weight_dtype, @@ -1702,9 +1700,7 @@ def main(): vae, text_encoder, tokenizer, - tokenizer_2, transformer3d, - network, args, accelerator, weight_dtype, diff --git a/scripts/z_image_fun/train_control_distill.py b/scripts/z_image_fun/train_control_distill.py index c0f491b..8aac3d4 100644 --- a/scripts/z_image_fun/train_control_distill.py +++ b/scripts/z_image_fun/train_control_distill.py @@ -87,9 +87,8 @@ from videox_fun.models import (AutoencoderKL, AutoProcessor, AutoTokenizer, CLIPImageProcessor, CLIPVisionModelWithProjection, Qwen2Tokenizer, Qwen3ForCausalLM, QwenImageTransformer2DModel, - ZImageControlTransformer2DModel, - ZImageTransformer2DModel) -from videox_fun.pipeline import Flux2Pipeline + ZImageControlTransformer2DModel) +from videox_fun.pipeline import ZImageControlPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid @@ -196,11 +195,11 @@ 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... ") - transformer3d_val = ZImageTransformer2DModel.from_pretrained( + transformer3d_val = ZImageControlTransformer2DModel.from_pretrained( args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, low_cpu_mem_usage=True, ).to(weight_dtype) @@ -210,7 +209,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, a subfolder="scheduler" ) transformer3d = transformer3d.to("cpu") - pipeline = Flux2Pipeline( + pipeline = ZImageControlPipeline( vae=accelerator.unwrap_model(vae).to(weight_dtype), text_encoder=accelerator.unwrap_model(text_encoder), tokenizer=tokenizer, @@ -964,13 +963,13 @@ def main(): if zero_stage == 3: raise NotImplementedError("FSDP does not support EMA.") - ema_transformer3d = ZImageTransformer2DModel.from_pretrained( + ema_transformer3d = ZImageControlTransformer2DModel.from_pretrained( args.pretrained_model_name_or_path, subfolder="transformer", torch_dtype=weight_dtype, ).to(weight_dtype) - ema_transformer3d = EMAModel(ema_transformer3d.parameters(), model_cls=ZImageTransformer2DModel, model_config=ema_transformer3d.config) + ema_transformer3d = EMAModel(ema_transformer3d.parameters(), model_cls=ZImageControlTransformer2DModel, model_config=ema_transformer3d.config) # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): @@ -1032,7 +1031,7 @@ def main(): model = models.pop() # load diffusers style into model - load_model = ZImageTransformer2DModel.from_pretrained( + load_model = ZImageControlTransformer2DModel.from_pretrained( input_dir, subfolder="transformer" ) model.register_to_config(**load_model.config) @@ -2010,7 +2009,6 @@ def main(): text_encoder, tokenizer, transformer3d, - network, args, accelerator, weight_dtype, @@ -2036,9 +2034,7 @@ def main(): vae, text_encoder, tokenizer, - tokenizer_2, transformer3d, - network, args, accelerator, weight_dtype,