RoPE fixes for 1.5, bfloat16 support in prepare_dataset, gradient_accumulation grad norm undefined fix (#107)

* rope scaling fixes for cog 1.5

* grad norm log fix

* bf16 prepare dataset fix

* update

* style
This commit is contained in:
Aryan
2024-12-02 15:12:12 +05:30
committed by GitHub
parent 76a9f2b911
commit 6c41984d8a
5 changed files with 58 additions and 30 deletions
+17 -9
View File
@@ -398,6 +398,8 @@ def main(args):
VAE_SCALING_FACTOR = vae.config.scaling_factor
VAE_SCALE_FACTOR_SPATIAL = 2 ** (len(vae.config.block_out_channels) - 1)
RoPE_BASE_HEIGHT = transformer.config.sample_height * VAE_SCALE_FACTOR_SPATIAL
RoPE_BASE_WIDTH = transformer.config.sample_width * VAE_SCALE_FACTOR_SPATIAL
# For mixed precision training we cast all non-trainable weights (vae, text_encoder and transformer) to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
@@ -715,6 +717,7 @@ def main(args):
for step, batch in enumerate(train_dataloader):
models_to_accumulate = [transformer]
logs = {}
with accelerator.accumulate(models_to_accumulate):
images = batch["images"].to(accelerator.device, non_blocking=True)
@@ -790,6 +793,8 @@ def main(args):
patch_size_t=model_config.patch_size_t if hasattr(model_config, "patch_size_t") else None,
attention_head_dim=model_config.attention_head_dim,
device=accelerator.device,
base_height=RoPE_BASE_HEIGHT,
base_width=RoPE_BASE_WIDTH,
)
if model_config.use_rotary_positional_embeddings
else None
@@ -828,6 +833,12 @@ def main(args):
gradient_norm_before_clip = get_gradient_norm(transformer.parameters())
accelerator.clip_grad_norm_(transformer.parameters(), args.max_grad_norm)
gradient_norm_after_clip = get_gradient_norm(transformer.parameters())
logs.update(
{
"gradient_norm_before_clip": gradient_norm_before_clip,
"gradient_norm_after_clip": gradient_norm_after_clip,
}
)
if accelerator.state.deepspeed_plugin is None:
optimizer.step()
@@ -876,15 +887,12 @@ def main(args):
run_validation(args, accelerator, transformer, scheduler, model_config, weight_dtype)
last_lr = lr_scheduler.get_last_lr()[0] if lr_scheduler is not None else args.learning_rate
logs = {"loss": loss.detach().item(), "lr": last_lr}
# gradnorm + deepspeed: https://github.com/microsoft/DeepSpeed/issues/4555
if accelerator.sync_gradients and accelerator.distributed_type != DistributedType.DEEPSPEED:
logs.update(
{
"gradient_norm_before_clip": gradient_norm_before_clip,
"gradient_norm_after_clip": gradient_norm_after_clip,
}
)
logs.update(
{
"loss": loss.detach().item(),
"lr": last_lr,
}
)
progress_bar.set_postfix(**logs)
accelerator.log(logs, step=global_step)
+17 -9
View File
@@ -328,6 +328,8 @@ def main(args):
VAE_SCALING_FACTOR = vae.config.scaling_factor
VAE_SCALE_FACTOR_SPATIAL = 2 ** (len(vae.config.block_out_channels) - 1)
RoPE_BASE_HEIGHT = transformer.config.sample_height * VAE_SCALE_FACTOR_SPATIAL
RoPE_BASE_WIDTH = transformer.config.sample_width * VAE_SCALE_FACTOR_SPATIAL
# For mixed precision training we cast all non-trainable weights (vae, text_encoder and transformer) to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
@@ -644,6 +646,7 @@ def main(args):
for step, batch in enumerate(train_dataloader):
models_to_accumulate = [transformer]
logs = {}
with accelerator.accumulate(models_to_accumulate):
videos = batch["videos"].to(accelerator.device, non_blocking=True)
@@ -699,6 +702,8 @@ def main(args):
patch_size_t=model_config.patch_size_t if hasattr(model_config, "patch_size_t") else None,
attention_head_dim=model_config.attention_head_dim,
device=accelerator.device,
base_height=RoPE_BASE_HEIGHT,
base_width=RoPE_BASE_WIDTH,
)
if model_config.use_rotary_positional_embeddings
else None
@@ -736,6 +741,12 @@ def main(args):
gradient_norm_before_clip = get_gradient_norm(transformer.parameters())
accelerator.clip_grad_norm_(transformer.parameters(), args.max_grad_norm)
gradient_norm_after_clip = get_gradient_norm(transformer.parameters())
logs.update(
{
"gradient_norm_before_clip": gradient_norm_before_clip,
"gradient_norm_after_clip": gradient_norm_after_clip,
}
)
if accelerator.state.deepspeed_plugin is None:
optimizer.step()
@@ -776,15 +787,12 @@ def main(args):
logger.info(f"Saved state to {save_path}")
last_lr = lr_scheduler.get_last_lr()[0] if lr_scheduler is not None else args.learning_rate
logs = {"loss": loss.detach().item(), "lr": last_lr}
# gradnorm + deepspeed: https://github.com/microsoft/DeepSpeed/issues/4555
if accelerator.sync_gradients and accelerator.distributed_type != DistributedType.DEEPSPEED:
logs.update(
{
"gradient_norm_before_clip": gradient_norm_before_clip,
"gradient_norm_after_clip": gradient_norm_after_clip,
}
)
logs.update(
{
"loss": loss.detach().item(),
"lr": last_lr,
}
)
progress_bar.set_postfix(**logs)
accelerator.log(logs, step=global_step)
+17 -9
View File
@@ -317,6 +317,8 @@ def main(args):
VAE_SCALING_FACTOR = vae.config.scaling_factor
VAE_SCALE_FACTOR_SPATIAL = 2 ** (len(vae.config.block_out_channels) - 1)
RoPE_BASE_HEIGHT = transformer.config.sample_height * VAE_SCALE_FACTOR_SPATIAL
RoPE_BASE_WIDTH = transformer.config.sample_width * VAE_SCALE_FACTOR_SPATIAL
# For mixed precision training we cast all non-trainable weights (vae, text_encoder and transformer) to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
@@ -610,6 +612,7 @@ def main(args):
for step, batch in enumerate(train_dataloader):
models_to_accumulate = [transformer]
logs = {}
with accelerator.accumulate(models_to_accumulate):
videos = batch["videos"].to(accelerator.device, non_blocking=True)
@@ -665,6 +668,8 @@ def main(args):
patch_size_t=model_config.patch_size_t if hasattr(model_config, "patch_size_t") else None,
attention_head_dim=model_config.attention_head_dim,
device=accelerator.device,
base_height=RoPE_BASE_HEIGHT,
base_width=RoPE_BASE_WIDTH,
)
if model_config.use_rotary_positional_embeddings
else None
@@ -702,6 +707,12 @@ def main(args):
gradient_norm_before_clip = get_gradient_norm(transformer.parameters())
accelerator.clip_grad_norm_(transformer.parameters(), args.max_grad_norm)
gradient_norm_after_clip = get_gradient_norm(transformer.parameters())
logs.update(
{
"gradient_norm_before_clip": gradient_norm_before_clip,
"gradient_norm_after_clip": gradient_norm_after_clip,
}
)
if accelerator.state.deepspeed_plugin is None:
optimizer.step()
@@ -742,15 +753,12 @@ def main(args):
logger.info(f"Saved state to {save_path}")
last_lr = lr_scheduler.get_last_lr()[0] if lr_scheduler is not None else args.learning_rate
logs = {"loss": loss.detach().item(), "lr": last_lr}
# gradnorm + deepspeed: https://github.com/microsoft/DeepSpeed/issues/4555
if accelerator.sync_gradients and accelerator.distributed_type != DistributedType.DEEPSPEED:
logs.update(
{
"gradient_norm_before_clip": gradient_norm_before_clip,
"gradient_norm_after_clip": gradient_norm_after_clip,
}
)
logs.update(
{
"loss": loss.detach().item(),
"lr": last_lr,
}
)
progress_bar.set_postfix(**logs)
accelerator.log(logs, step=global_step)
+3 -1
View File
@@ -287,11 +287,13 @@ to_pil_image = transforms.ToPILImage(mode="RGB")
def save_image(image: torch.Tensor, path: pathlib.Path) -> None:
image = to_pil_image(image)
image = image.to(dtype=torch.float32).clamp(-1, 1)
image = to_pil_image(image.float())
image.save(path)
def save_video(video: torch.Tensor, path: pathlib.Path, fps: int = 8) -> None:
video = video.to(dtype=torch.float32).clamp(-1, 1)
video = [to_pil_image(frame) for frame in video]
export_to_video(video, path, fps=fps)
+4 -2
View File
@@ -208,9 +208,12 @@ def prepare_rotary_positional_embeddings(
grid_width = width // (vae_scale_factor_spatial * patch_size)
base_size_width = base_width // (vae_scale_factor_spatial * patch_size)
base_size_height = base_height // (vae_scale_factor_spatial * patch_size)
if patch_size_t is None:
# CogVideoX 1.0
grid_crops_coords = get_resize_crop_region_for_grid((grid_height, grid_width), base_size_width, base_size_height)
grid_crops_coords = get_resize_crop_region_for_grid(
(grid_height, grid_width), base_size_width, base_size_height
)
freqs_cos, freqs_sin = get_3d_rotary_pos_embed(
embed_dim=attention_head_dim,
crops_coords=grid_crops_coords,
@@ -230,7 +233,6 @@ def prepare_rotary_positional_embeddings(
max_size=(base_size_height, base_size_width),
)
freqs_cos = freqs_cos.to(device=device)
freqs_sin = freqs_sin.to(device=device)
return freqs_cos, freqs_sin