From 8c18421294f433d75b40210e46088c4877329f2e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jukka=20Sepp=C3=A4nen?= <40791699+kijai@users.noreply.github.com> Date: Thu, 10 Oct 2024 21:43:12 +0300 Subject: [PATCH] fix bf16 --- nodes.py | 8 +-- pyramid_dit/modeling_mmdit_block.py | 16 ++--- pyramid_dit/modeling_pyramid_mmdit.py | 71 ++++--------------- .../pyramid_dit_for_video_gen_pipeline.py | 40 ++++++----- trainer_misc/sp_utils.py | 1 - 5 files changed, 50 insertions(+), 86 deletions(-) diff --git a/nodes.py b/nodes.py index f09978b..ac6d9c2 100644 --- a/nodes.py +++ b/nodes.py @@ -196,15 +196,16 @@ class PyramidFlowSampler: torch.cuda.manual_seed(seed) autocastcondition = not model["model"].dtype == torch.float32 - autocast_context = torch.autocast(mm.get_autocast_device(device)) if autocastcondition else nullcontext() + #autocastcondition = True + autocast_context = torch.autocast(mm.get_autocast_device(device), dtype=model["model"].dtype) if autocastcondition else nullcontext() if input_latent is None: with autocast_context: latents = model["model"].generate( prompt_embeds_dict = prompt_embeds, device=device, - num_inference_steps=[first_frame_steps, first_frame_steps, first_frame_steps], #why's this a list - video_num_inference_steps=[video_steps, video_steps, video_steps], #why's this a list + num_inference_steps=[first_frame_steps, first_frame_steps, first_frame_steps], + video_num_inference_steps=[video_steps, video_steps, video_steps], height=height, width=width, temp=temp, @@ -222,7 +223,6 @@ class PyramidFlowSampler: height=height, width=width, temp=temp, - guidance_scale=guidance_scale, # The guidance for the first frame video_guidance_scale=video_guidance_scale, # The guidance for the other video latent output_type="latent", ) diff --git a/pyramid_dit/modeling_mmdit_block.py b/pyramid_dit/modeling_mmdit_block.py index ffd97ff..0160f47 100644 --- a/pyramid_dit/modeling_mmdit_block.py +++ b/pyramid_dit/modeling_mmdit_block.py @@ -476,15 +476,15 @@ class JointAttention(nn.Module): # print(f"Using flash-attention: {self.use_flash_attn}") if self.use_flash_attn: - if is_sequence_parallel_initialized(): - self.var_flash_attn = SequenceParallelVarlenFlashSelfAttentionWithT5Mask() - else: - self.var_flash_attn = VarlenFlashSelfAttentionWithT5Mask() + #if is_sequence_parallel_initialized(): + # self.var_flash_attn = SequenceParallelVarlenFlashSelfAttentionWithT5Mask() + #else: + self.var_flash_attn = VarlenFlashSelfAttentionWithT5Mask() else: - if is_sequence_parallel_initialized(): - self.var_len_attn = SequenceParallelVarlenSelfAttentionWithT5Mask() - else: - self.var_len_attn = VarlenSelfAttentionWithT5Mask() + #if is_sequence_parallel_initialized(): + #self.var_len_attn = SequenceParallelVarlenSelfAttentionWithT5Mask() + #else: + self.var_len_attn = VarlenSelfAttentionWithT5Mask() def forward( diff --git a/pyramid_dit/modeling_pyramid_mmdit.py b/pyramid_dit/modeling_pyramid_mmdit.py index dc65ad2..3113d3a 100644 --- a/pyramid_dit/modeling_pyramid_mmdit.py +++ b/pyramid_dit/modeling_pyramid_mmdit.py @@ -301,13 +301,6 @@ class PyramidDiffusionMMDiT(ModelMixin, ConfigMixin): text_ids = torch.zeros(pad_batch_size, encoder_attention_mask.shape[1], 1).to(device=device) input_ids_list = [torch.cat([text_ids, image_ids], dim=1) for image_ids in image_ids_list] image_rotary_emb = [self.temp_rope_embed(input_ids) for input_ids in input_ids_list] # [bs, seq_len, 1, head_dim // 2, 2, 2] - - if is_sequence_parallel_initialized(): - sp_group = get_sequence_parallel_group() - sp_group_size = get_sequence_parallel_world_size() - image_rotary_emb = [all_to_all(x_.repeat(1, 1, sp_group_size, 1, 1, 1), sp_group, sp_group_size, scatter_dim=2, gather_dim=0) for x_ in image_rotary_emb] - input_ids_list = [all_to_all(input_ids.repeat(1, 1, sp_group_size), sp_group, sp_group_size, scatter_dim=2, gather_dim=0) for input_ids in input_ids_list] - else: image_rotary_emb = None @@ -425,59 +418,23 @@ class PyramidDiffusionMMDiT(ModelMixin, ConfigMixin): hidden_states, hidden_length, temps, heights, widths, trainable_token_list, encoder_attention_mask, \ attention_mask, image_rotary_emb = self.merge_input(sample, encoder_hidden_length, encoder_attention_mask) - # split the long latents if necessary - if is_sequence_parallel_initialized(): - sp_group = get_sequence_parallel_group() - sp_group_size = get_sequence_parallel_world_size() - - # sync the input hidden states - batch_hidden_states = [] - for i_p, hidden_states_ in enumerate(hidden_states): - assert hidden_states_.shape[1] % sp_group_size == 0, "The sequence length should be divided by sequence parallel size" - hidden_states_ = all_to_all(hidden_states_, sp_group, sp_group_size, scatter_dim=1, gather_dim=0) - hidden_length[i_p] = hidden_length[i_p] // sp_group_size - batch_hidden_states.append(hidden_states_) - - # sync the encoder hidden states - hidden_states = torch.cat(batch_hidden_states, dim=1) - encoder_hidden_states = all_to_all(encoder_hidden_states, sp_group, sp_group_size, scatter_dim=1, gather_dim=0) - temb = all_to_all(temb.unsqueeze(1).repeat(1, sp_group_size, 1), sp_group, sp_group_size, scatter_dim=1, gather_dim=0) - temb = temb.squeeze(1) - else: - hidden_states = torch.cat(hidden_states, dim=1) + hidden_states = torch.cat(hidden_states, dim=1) # print(hidden_length) for i_b, block in enumerate(self.transformer_blocks): - if self.training and self.gradient_checkpointing and (i_b >= 2): - def create_custom_forward(module): - def custom_forward(*inputs): - return module(*inputs) - - return custom_forward - - ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} - encoder_hidden_states, hidden_states = torch.utils.checkpoint.checkpoint( - create_custom_forward(block), - hidden_states, - encoder_hidden_states, - encoder_attention_mask, - temb, - attention_mask, - hidden_length, - image_rotary_emb, - **ckpt_kwargs, - ) - - else: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - encoder_attention_mask=encoder_attention_mask, - temb=temb, - attention_mask=attention_mask, - hidden_length=hidden_length, - image_rotary_emb=image_rotary_emb, - ) + encoder_hidden_states, hidden_states = block( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + encoder_attention_mask=encoder_attention_mask, + temb=temb, + attention_mask=attention_mask, + hidden_length=hidden_length, + image_rotary_emb=image_rotary_emb, + ) + + # nan_mask = torch.isnan(hidden_states) + # if torch.any(nan_mask): + # raise ValueError("nan in hidden_states") hidden_states = self.norm_out(hidden_states, temb, hidden_length=hidden_length) hidden_states = self.proj_out(hidden_states) diff --git a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py index e590add..0dfad62 100644 --- a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py +++ b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py @@ -178,23 +178,23 @@ class PyramidDiTForVideoGeneration: latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype) return latents - # def sample_block_noise(self, bs, ch, temp, height, width): - # gamma = self.scheduler.config.gamma - # dist = torch.distributions.multivariate_normal.MultivariateNormal(torch.zeros(4), torch.eye(4) * (1 + gamma) - torch.ones(4, 4) * gamma) - # block_number = bs * ch * temp * (height // 2) * (width // 2) - # noise = torch.stack([dist.sample() for _ in range(block_number)]) # [block number, 4] - # noise = rearrange(noise, '(b c t h w) (p q) -> b c t (h p) (w q)',b=bs,c=ch,t=temp,h=height//2,w=width//2,p=2,q=2) - # return noise - def sample_block_noise(self, bs, ch, temp, height, width): gamma = self.scheduler.config.gamma - epsilon = 1e-5 # Small value to ensure positive definiteness - covariance_matrix = torch.eye(4) * (1 + gamma) - torch.ones(4, 4) * gamma + torch.eye(4) * epsilon - dist = torch.distributions.multivariate_normal.MultivariateNormal(torch.zeros(4), covariance_matrix) + dist = torch.distributions.multivariate_normal.MultivariateNormal(torch.zeros(4), torch.eye(4) * (1 + gamma) - torch.ones(4, 4) * gamma) block_number = bs * ch * temp * (height // 2) * (width // 2) - noise = torch.stack([dist.sample() for _ in range(block_number)]) # [block number, 4] - noise = rearrange(noise, '(b c t h w) (p q) -> b c t (h p) (w q)', b=bs, c=ch, t=temp, h=height//2, w=width//2, p=2, q=2) + noise = torch.stack([dist.sample() for _ in range(block_number)]) # [block number, 4] + noise = rearrange(noise, '(b c t h w) (p q) -> b c t (h p) (w q)',b=bs,c=ch,t=temp,h=height//2,w=width//2,p=2,q=2) return noise + + # def sample_block_noise(self, bs, ch, temp, height, width): + # gamma = self.scheduler.config.gamma + # epsilon = 1e-5 # Small value to ensure positive definiteness + # covariance_matrix = torch.eye(4) * (1 + gamma) - torch.ones(4, 4) * gamma + torch.eye(4) * epsilon + # dist = torch.distributions.multivariate_normal.MultivariateNormal(torch.zeros(4), covariance_matrix) + # block_number = bs * ch * temp * (height // 2) * (width // 2) + # noise = torch.stack([dist.sample() for _ in range(block_number)]) # [block number, 4] + # noise = rearrange(noise, '(b c t h w) (p q) -> b c t (h p) (w q)', b=bs, c=ch, t=temp, h=height//2, w=width//2, p=2, q=2) + # return noise @torch.no_grad() def generate_one_unit( @@ -255,6 +255,10 @@ class PyramidDiTForVideoGeneration: ) noise_pred = noise_pred[0] + + # nan_mask = torch.isnan(noise_pred) + # if torch.any(nan_mask): + # raise ValueError("nan in hidden_states") # perform guidance if self.do_classifier_free_guidance: @@ -271,6 +275,9 @@ class PyramidDiTForVideoGeneration: sample=latents, generator=generator, ).prev_sample + nan_mask = torch.isnan(latents) + if torch.any(nan_mask): + raise ValueError("nan in latents") intermed_latents.append(latents) @@ -432,6 +439,7 @@ class PyramidDiTForVideoGeneration: generator, is_first_frame=False, ) + comfy_pbar.update(1) generated_latents_list.append(intermed_latents[-1]) last_generated_latents = intermed_latents @@ -516,9 +524,9 @@ class PyramidDiTForVideoGeneration: pooled_prompt_embeds = torch.cat([negative_pooled_prompt_embeds, positive_pooled_prompt_embeds], dim=0) prompt_attention_mask = torch.cat([negative_prompt_attention_mask, positive_prompt_attention_mask], dim=0) - prompt_embeds = prompt_embeds.to(dtype) - pooled_prompt_embeds = pooled_prompt_embeds.to(dtype) - prompt_attention_mask = prompt_attention_mask.to(dtype) + # prompt_embeds = prompt_embeds.to(dtype) + # pooled_prompt_embeds = pooled_prompt_embeds.to(dtype) + # prompt_attention_mask = prompt_attention_mask.to(dtype) # Create the initial random noise num_channels_latents = self.dit.config.in_channels diff --git a/trainer_misc/sp_utils.py b/trainer_misc/sp_utils.py index d261bf6..e0c078c 100644 --- a/trainer_misc/sp_utils.py +++ b/trainer_misc/sp_utils.py @@ -1,6 +1,5 @@ import os import torch -import torch.distributed as dist from .utils import is_dist_avail_and_initialized, get_rank