delete useless code

This commit is contained in:
bubbliiiing
2026-01-05 20:28:36 +08:00
parent 643d1c324c
commit 47412d2440
6 changed files with 28 additions and 44 deletions
+1 -4
View File
@@ -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,
-1
View File
@@ -1563,7 +1563,6 @@ def main():
vae,
text_encoder,
tokenizer,
tokenizer_2,
transformer3d,
network,
args,
+5 -6
View File
@@ -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,
+3 -6
View File
@@ -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,
+11 -15
View File
@@ -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,
+8 -12
View File
@@ -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,