Fix bug in v3 training (#98)
This commit is contained in:
+22
-22
@@ -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
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user