This commit is contained in:
Jukka Seppänen
2024-10-10 21:43:12 +03:00
parent 0bc9a12e7a
commit 8c18421294
5 changed files with 50 additions and 86 deletions
+4 -4
View File
@@ -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",
)
+8 -8
View File
@@ -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(
+14 -57
View File
@@ -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
View File
@@ -1,6 +1,5 @@
import os
import torch
import torch.distributed as dist
from .utils import is_dist_avail_and_initialized, get_rank