Update Flux2 Training Code (#436)
This commit is contained in:
+90
-85
@@ -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
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user