Fix bug in v3 training (#98)

This commit is contained in:
Bubbliiiing
2024-08-22 16:03:14 +08:00
committed by GitHub
parent 5ea1bf2450
commit d3b8bbbd14
4 changed files with 49 additions and 42 deletions
+22 -22
View File
@@ -1557,8 +1557,8 @@ def main():
if args.train_mode != "normal":
clip_pixel_values = batch["clip_pixel_values"]
mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype)
mask = batch["mask"].to(weight_dtype)
mask_pixel_values = batch["mask_pixel_values"].to(accelerator.device, dtype=weight_dtype)
mask = batch["mask"].to(accelerator.device, dtype=weight_dtype)
if args.training_with_video_token_length:
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
clip_pixel_values = torch.tile(clip_pixel_values, (4, 1, 1, 1))
@@ -1568,7 +1568,7 @@ def main():
clip_pixel_values = torch.tile(clip_pixel_values, (2, 1, 1, 1))
mask_pixel_values = torch.tile(mask_pixel_values, (2, 1, 1, 1, 1))
mask = torch.tile(mask, (2, 1, 1, 1, 1))
def create_special_list(length):
if length == 1:
return [1.0]
@@ -1621,7 +1621,7 @@ def main():
text_encoder.to(accelerator.device)
if text_encoder_2 is not None:
text_encoder_2.to(accelerator.device)
with torch.no_grad():
if vae.quant_conv.weight.ndim==5:
# This way is quicker when batch grows up
@@ -1648,7 +1648,7 @@ def main():
new_pixel_values = []
for i in range(0, pixel_values.shape[0], bs):
pixel_values_bs = pixel_values[i : i + bs]
pixel_values_bs = vae.encode(pixel_values_bs.to(dtype=weight_dtype)).latent_dist
pixel_values_bs = vae.encode(pixel_values_bs.to(accelerator.device, dtype=weight_dtype)).latent_dist
pixel_values_bs = pixel_values_bs.sample()
new_pixel_values.append(pixel_values_bs)
latents = torch.cat(new_pixel_values, dim = 0)
@@ -1695,7 +1695,7 @@ def main():
new_mask_pixel_values = []
for i in range(0, mask_pixel_values.shape[0], bs):
mask_pixel_values_bs = mask_pixel_values[i : i + bs]
mask_pixel_values_bs = vae.encode(mask_pixel_values_bs.to(dtype=weight_dtype)).latent_dist
mask_pixel_values_bs = vae.encode(mask_pixel_values_bs.to(accelerator.device, dtype=weight_dtype)).latent_dist
mask_pixel_values_bs = mask_pixel_values_bs.sample()
new_mask_pixel_values.append(mask_pixel_values_bs)
mask_latents = torch.cat(new_mask_pixel_values, dim = 0)
@@ -1746,13 +1746,13 @@ def main():
if args.enable_text_encoder_in_dataloader:
if config.get('enable_multi_text_encoder', False):
prompt_embeds = batch['prompt_embeds'].to(device=latents.device)
prompt_attention_mask = batch['prompt_attention_mask'].to(device=latents.device)
prompt_embeds_2 = batch['prompt_embeds_2'].to(device=latents.device)
prompt_attention_mask_2 = batch['prompt_attention_mask_2'].to(device=latents.device)
prompt_embeds = batch['prompt_embeds'].to(accelerator.device, dtype=weight_dtype)
prompt_attention_mask = batch['prompt_attention_mask'].to(accelerator.device, dtype=weight_dtype)
prompt_embeds_2 = batch['prompt_embeds_2'].to(accelerator.device, dtype=weight_dtype)
prompt_attention_mask_2 = batch['prompt_attention_mask_2'].to(accelerator.device, dtype=weight_dtype)
else:
encoder_attention_mask = batch['encoder_attention_mask'].to(device=latents.device)
encoder_hidden_states = batch['encoder_hidden_states'].to(device=latents.device)
encoder_attention_mask = batch['encoder_attention_mask'].to(accelerator.device, dtype=weight_dtype)
encoder_hidden_states = batch['encoder_hidden_states'].to(accelerator.device, dtype=weight_dtype)
else:
if config.get('enable_multi_text_encoder', False):
with torch.no_grad():
@@ -1772,11 +1772,11 @@ def main():
return_tensors="pt"
)
encoder_hidden_states = text_encoder(
prompt_ids.input_ids.to(latents.device),
attention_mask=prompt_ids.attention_mask.to(latents.device),
prompt_ids.input_ids.to(accelerator.device),
attention_mask=prompt_ids.attention_mask.to(accelerator.device),
return_dict=False
)[0]
encoder_attention_mask = prompt_ids.attention_mask.to(latents.device)
)[0].to(accelerator.device, dtype=weight_dtype)
encoder_attention_mask = prompt_ids.attention_mask.to(accelerator.device, dtype=weight_dtype)
if args.low_vram and not args.enable_text_encoder_in_dataloader:
text_encoder.to('cpu')
@@ -1796,7 +1796,7 @@ def main():
noise = torch.randn(latents.size(), device=latents.device, generator=torch_rng, dtype=weight_dtype)
# Sample a random timestep for each image
timesteps = generate_timestep_with_lognorm(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng)
timesteps = timesteps.long()
timesteps = timesteps.long().to(accelerator.device)
if config.get('enable_multi_text_encoder', False):
height, width = batch["pixel_values"].size()[-2], batch["pixel_values"].size()[-1]
@@ -1805,7 +1805,7 @@ def main():
grid_width = width // 8 // accelerator.unwrap_model(transformer3d).config.patch_size
base_size = 512 // 8 // accelerator.unwrap_model(transformer3d).config.patch_size
grid_crops_coords = get_resize_crop_region_for_grid((grid_height, grid_width), base_size)
image_rotary_emb = _get_2d_rotary_pos_embed_cached(
image_rotary_emb = get_2d_rotary_pos_embed(
accelerator.unwrap_model(transformer3d).inner_dim // accelerator.unwrap_model(transformer3d).num_heads, grid_crops_coords, (grid_height, grid_width)
)
@@ -1862,8 +1862,8 @@ def main():
bs, height, width = bsz, batch["pixel_values"].size()[-2], batch["pixel_values"].size()[-1]
resolution = torch.tensor([height, width]).repeat(bs, 1)
aspect_ratio = torch.tensor([float(height / width)]).repeat(bs, 1)
resolution = resolution.to(dtype=encoder_hidden_states.dtype, device=latents.device)
aspect_ratio = aspect_ratio.to(dtype=encoder_hidden_states.dtype, device=latents.device)
resolution = resolution.to(accelerator.device, dtype=weight_dtype)
aspect_ratio = aspect_ratio.to(accelerator.device, dtype=weight_dtype)
added_cond_kwargs = {"resolution": resolution, "aspect_ratio": aspect_ratio}
loss_term = train_diffusion.training_losses(
@@ -1872,8 +1872,8 @@ def main():
timesteps,
noise=noise,
model_kwargs=dict(
encoder_hidden_states=encoder_hidden_states.to(latents.device, latents.dtype),
encoder_attention_mask=encoder_attention_mask.to(latents.device, latents.dtype),
encoder_hidden_states=encoder_hidden_states,
encoder_attention_mask=encoder_attention_mask,
added_cond_kwargs=added_cond_kwargs,
inpaint_latents=inpaint_latents if args.train_mode != "normal" else None,
clip_encoder_hidden_states=clip_encoder_hidden_states if args.train_mode != "normal" else None,
+22 -17
View File
@@ -1495,8 +1495,8 @@ def main():
if args.train_mode != "normal":
clip_pixel_values = batch["clip_pixel_values"]
mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype)
mask = batch["mask"].to(weight_dtype)
mask_pixel_values = batch["mask_pixel_values"].to(accelerator.device, dtype=weight_dtype)
mask = batch["mask"].to(accelerator.device, dtype=weight_dtype)
if args.training_with_video_token_length:
if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]:
clip_pixel_values = torch.tile(clip_pixel_values, (4, 1, 1, 1))
@@ -1555,6 +1555,10 @@ def main():
if args.low_vram:
torch.cuda.empty_cache()
vae.to(accelerator.device)
if not args.enable_text_encoder_in_dataloader:
text_encoder.to(accelerator.device)
if text_encoder_2 is not None:
text_encoder_2.to(accelerator.device)
with torch.no_grad():
if vae.quant_conv.weight.ndim==5:
@@ -1582,7 +1586,7 @@ def main():
new_pixel_values = []
for i in range(0, pixel_values.shape[0], bs):
pixel_values_bs = pixel_values[i : i + bs]
pixel_values_bs = vae.encode(pixel_values_bs.to(dtype=weight_dtype)).latent_dist
pixel_values_bs = vae.encode(pixel_values_bs.to(accelerator.device, dtype=weight_dtype)).latent_dist
pixel_values_bs = pixel_values_bs.sample()
new_pixel_values.append(pixel_values_bs)
latents = torch.cat(new_pixel_values, dim = 0)
@@ -1629,7 +1633,7 @@ def main():
new_mask_pixel_values = []
for i in range(0, mask_pixel_values.shape[0], bs):
mask_pixel_values_bs = mask_pixel_values[i : i + bs]
mask_pixel_values_bs = vae.encode(mask_pixel_values_bs.to(dtype=weight_dtype)).latent_dist
mask_pixel_values_bs = vae.encode(mask_pixel_values_bs.to(accelerator.device, dtype=weight_dtype)).latent_dist
mask_pixel_values_bs = mask_pixel_values_bs.sample()
new_mask_pixel_values.append(mask_pixel_values_bs)
mask_latents = torch.cat(new_mask_pixel_values, dim = 0)
@@ -1680,13 +1684,13 @@ def main():
if args.enable_text_encoder_in_dataloader:
if config.get('enable_multi_text_encoder', False):
prompt_embeds = batch['prompt_embeds'].to(device=latents.device)
prompt_attention_mask = batch['prompt_attention_mask'].to(device=latents.device)
prompt_embeds_2 = batch['prompt_embeds_2'].to(device=latents.device)
prompt_attention_mask_2 = batch['prompt_attention_mask_2'].to(device=latents.device)
prompt_embeds = batch['prompt_embeds'].to(accelerator.device, dtype=weight_dtype)
prompt_attention_mask = batch['prompt_attention_mask'].to(accelerator.device, dtype=weight_dtype)
prompt_embeds_2 = batch['prompt_embeds_2'].to(accelerator.device, dtype=weight_dtype)
prompt_attention_mask_2 = batch['prompt_attention_mask_2'].to(accelerator.device, dtype=weight_dtype)
else:
encoder_attention_mask = batch['encoder_attention_mask'].to(device=latents.device)
encoder_hidden_states = batch['encoder_hidden_states'].to(device=latents.device)
encoder_attention_mask = batch['encoder_attention_mask'].to(accelerator.device, dtype=weight_dtype)
encoder_hidden_states = batch['encoder_hidden_states'].to(accelerator.device, dtype=weight_dtype)
else:
if config.get('enable_multi_text_encoder', False):
with torch.no_grad():
@@ -1706,11 +1710,11 @@ def main():
return_tensors="pt"
)
encoder_hidden_states = text_encoder(
prompt_ids.input_ids.to(latents.device),
attention_mask=prompt_ids.attention_mask.to(latents.device),
prompt_ids.input_ids.to(accelerator.device),
attention_mask=prompt_ids.attention_mask.to(accelerator.device),
return_dict=False
)[0]
encoder_attention_mask = prompt_ids.attention_mask.to(latents.device)
)[0].to(accelerator.device, dtype=weight_dtype)
encoder_attention_mask = prompt_ids.attention_mask.to(accelerator.device, dtype=weight_dtype)
if args.low_vram and not args.enable_text_encoder_in_dataloader:
text_encoder.to('cpu')
@@ -1730,7 +1734,7 @@ def main():
noise = torch.randn(latents.size(), device=latents.device, generator=torch_rng, dtype=weight_dtype)
# Sample a random timestep for each image
timesteps = generate_timestep_with_lognorm(0, args.train_sampling_steps, (bsz,), device=latents.device, generator=torch_rng)
timesteps = timesteps.long()
timesteps = timesteps.long().to(accelerator.device)
if config.get('enable_multi_text_encoder', False):
height, width = batch["pixel_values"].size()[-2], batch["pixel_values"].size()[-1]
@@ -1766,6 +1770,7 @@ def main():
else:
raise ValueError(f"Unknown prediction type {noise_scheduler.config.prediction_type}")
# predict the noise residual
noise_pred = transformer3d(
noisy_latents,
timesteps.to(noisy_latents.dtype),
@@ -1795,8 +1800,8 @@ def main():
bs, height, width = bsz, batch["pixel_values"].size()[-2], batch["pixel_values"].size()[-1]
resolution = torch.tensor([height, width]).repeat(bs, 1)
aspect_ratio = torch.tensor([float(height / width)]).repeat(bs, 1)
resolution = resolution.to(dtype=encoder_hidden_states.dtype, device=latents.device)
aspect_ratio = aspect_ratio.to(dtype=encoder_hidden_states.dtype, device=latents.device)
resolution = resolution.to(accelerator.device, dtype=weight_dtype)
aspect_ratio = aspect_ratio.to(accelerator.device, dtype=weight_dtype)
added_cond_kwargs = {"resolution": resolution, "aspect_ratio": aspect_ratio}
loss_term = train_diffusion.training_losses(