Update Flux2 Training Code (#436)

This commit is contained in:
Bubbliiiing
2026-01-16 17:21:25 +08:00
committed by GitHub
parent b534ebfef1
commit ac114cc142
6 changed files with 428 additions and 385 deletions
+90 -85
View File
@@ -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()
+79 -73
View File
@@ -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()
+84 -80
View File
@@ -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()
+73 -68
View File
@@ -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()
+102 -78
View File
@@ -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()