fix bf16
This commit is contained in:
@@ -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",
|
||||
)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
import os
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from .utils import is_dist_avail_and_initialized, get_rank
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user