Add latent preview, cleanup

This commit is contained in:
kijai
2024-11-01 14:38:28 +02:00
parent bc8d4360fe
commit 3aeea32237
14 changed files with 244 additions and 1455 deletions
+87
View File
@@ -0,0 +1,87 @@
import io
import torch
from PIL import Image
import struct
import numpy as np
from comfy.cli_args import args, LatentPreviewMethod
from comfy.taesd.taesd import TAESD
import comfy.model_management
import folder_paths
import comfy.utils
import logging
MAX_PREVIEW_RESOLUTION = args.preview_size
def preview_to_image(latent_image):
latents_ubyte = (((latent_image + 1.0) / 2.0).clamp(0, 1) # change scale from -1..1 to 0..1
.mul(0xFF) # to 0..255
).to(device="cpu", dtype=torch.uint8, non_blocking=comfy.model_management.device_supports_non_blocking(latent_image.device))
return Image.fromarray(latents_ubyte.numpy())
class LatentPreviewer:
def decode_latent_to_preview(self, x0):
pass
def decode_latent_to_preview_image(self, preview_format, x0):
preview_image = self.decode_latent_to_preview(x0)
return ("GIF", preview_image, MAX_PREVIEW_RESOLUTION)
class Latent2RGBPreviewer(LatentPreviewer):
def __init__(self, latent_rgb_factors, latent_rgb_factors_bias=None):
latent_rgb_factors = [[0.05389399697934166, 0.025018778505575393, -0.009193515248318657], [0.02318250640590553, -0.026987363837713156, 0.040172639061236956], [0.046035451343323666, -0.02039565868920197, 0.01275569344290342], [-0.015559161155025095, 0.051403973219861246, 0.03179031307996347], [-0.02766167769640129, 0.03749545161530447, 0.003335141009473408], [0.05824598730479011, 0.021744367381243884, -0.01578925627951616], [0.05260929401500947, 0.0560165014956886, -0.027477296572565126], [0.018513891242931686, 0.041961785217662514, 0.004490763489747966], [0.024063060899760215, 0.065082853069653, 0.044343437673514896], [0.05250992323006226, 0.04361117432588933, 0.01030076055524387], [0.0038921710021782366, -0.025299228133723792, 0.019370764014574535], [-0.00011950534333568519, 0.06549370069727675, -0.03436712163379723], [-0.026020578032683626, -0.013341758571090847, -0.009119046570271953], [0.024412451175602937, 0.030135064560817174, -0.008355486384198006], [0.04002209845752687, -0.017341304390739463, 0.02818338690302971], [-0.032575108695213684, -0.009588338926775117, -0.03077312160940468]]
self.latent_rgb_factors = torch.tensor(latent_rgb_factors, device="cpu").transpose(0, 1)
self.latent_rgb_factors_bias = None
# if latent_rgb_factors_bias is not None:
# self.latent_rgb_factors_bias = torch.tensor(latent_rgb_factors_bias, device="cpu")
def decode_latent_to_preview(self, x0):
self.latent_rgb_factors = self.latent_rgb_factors.to(dtype=x0.dtype, device=x0.device)
if self.latent_rgb_factors_bias is not None:
self.latent_rgb_factors_bias = self.latent_rgb_factors_bias.to(dtype=x0.dtype, device=x0.device)
latent_image = torch.nn.functional.linear(x0[0].permute(1, 2, 0), self.latent_rgb_factors,
bias=self.latent_rgb_factors_bias)
return preview_to_image(latent_image)
def get_previewer(device, latent_format):
previewer = None
method = args.preview_method
if method != LatentPreviewMethod.NoPreviews:
# TODO previewer methods
taesd_decoder_path = None
if latent_format.taesd_decoder_name is not None:
taesd_decoder_path = next(
(fn for fn in folder_paths.get_filename_list("vae_approx")
if fn.startswith(latent_format.taesd_decoder_name)),
""
)
taesd_decoder_path = folder_paths.get_full_path("vae_approx", taesd_decoder_path)
if method == LatentPreviewMethod.Auto:
method = LatentPreviewMethod.Latent2RGB
if previewer is None:
if latent_format.latent_rgb_factors is not None:
previewer = Latent2RGBPreviewer(latent_format.latent_rgb_factors, latent_format.latent_rgb_factors_bias)
return previewer
def prepare_callback(model, steps, x0_output_dict=None):
preview_format = "JPEG"
if preview_format not in ["JPEG", "PNG"]:
preview_format = "JPEG"
previewer = get_previewer(model.load_device, model.model.latent_format)
pbar = comfy.utils.ProgressBar(steps)
def callback(step, x0, x, total_steps):
if x0_output_dict is not None:
x0_output_dict["x0"] = x0
preview_bytes = None
if previewer:
preview_bytes = previewer.decode_latent_to_preview_image(preview_format, x0)
pbar.update_absolute(step + 1, total_steps, preview_bytes)
return callback
+88 -10
View File
@@ -261,12 +261,14 @@ class PyramidFlowSampler:
if isinstance(model, dict):
pyramid_model = model["model"]
#dtype = model["dtype"]
else:
pyramid_model = model
dtype = pyramid_model.dit.dtype
from .latent_preview import prepare_callback
callback = prepare_callback(model, temp)
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
@@ -289,20 +291,20 @@ class PyramidFlowSampler:
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",
callback=callback,
)
else:
with autocast_context:
latents = pyramid_model.generate_i2v(
prompt_embeds_dict = prompt_embeds,
input_image_latent=input_latent,
input_image_latent=input_latent["samples"],
device=device,
num_inference_steps=video_steps,
height=height,
width=width,
temp=temp,
video_guidance_scale=video_guidance_scale, # The guidance for the other video latent
output_type="latent",
callback=callback,
)
if not keep_model_loaded and not pyramid_model.sequential_offload_enabled:
@@ -384,6 +386,7 @@ class PyramidFlowVAEEncode:
CATEGORY = "PyramidFlowWrapper"
def sample(self, vae, image, enable_tiling):
B, H, W, C = image.shape
mm.soft_empty_cache()
dtype = vae.dtype
@@ -400,15 +403,16 @@ class PyramidFlowVAEEncode:
normalize = transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))
input_image_tensor = rearrange(image, 'b h w c -> b c h w')
input_image_tensor = normalize(input_image_tensor)
input_image_tensor = input_image_tensor.unsqueeze(2) # Add temporal dimension t=1
input_image_tensor = normalize(input_image_tensor).unsqueeze(0)
input_image_tensor = rearrange(input_image_tensor, 'b t c h w -> b c t h w', t=B)
#input_image_tensor = input_image_tensor.unsqueeze(2) # Add temporal dimension t=1
input_image_tensor = input_image_tensor.to(dtype=dtype, device=device)
vae.to(device)
input_image_latent = (vae.encode(input_image_tensor).latent_dist.sample() - vae_shift_factor) * vae_scale_factor # [b c 1 h w]
vae.to(offload_device)
return (input_image_latent,)
return ({"samples": input_image_latent},)
class PyramidFlowVAEDecode:
@classmethod
@@ -467,7 +471,79 @@ class PyramidFlowVAEDecode:
return (image,)
class PyramidFlowLatentPreview:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"samples": ("LATENT",),
# "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
# "min_val": ("FLOAT", {"default": -0.15, "min": -1.0, "max": 0.0, "step": 0.001}),
# "max_val": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.001}),
},
}
RETURN_TYPES = ("IMAGE", "STRING", )
RETURN_NAMES = ("images", "latent_rgb_factors", )
FUNCTION = "sample"
CATEGORY = "PyramidFlowWrapper"
def sample(self, samples):#, seed, min_val, max_val):
mm.soft_empty_cache()
latents = samples["samples"].clone()
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
# For the image latent
vae_shift_factor = 0.1490
vae_scale_factor = 1 / 1.8415
# For the video latent
vae_video_shift_factor = -0.2343
vae_video_scale_factor = 1 / 3.0986
if latents.shape[2] == 1:
latents = (latents / vae_scale_factor) + vae_shift_factor
else:
latents[:, :, :1] = (latents[:, :, :1] / vae_scale_factor) + vae_shift_factor
latents[:, :, 1:] = (latents[:, :, 1:] / vae_video_scale_factor) + vae_video_shift_factor
latent_rgb_factors = [[0.05389399697934166, 0.025018778505575393, -0.009193515248318657], [0.02318250640590553, -0.026987363837713156, 0.040172639061236956], [0.046035451343323666, -0.02039565868920197, 0.01275569344290342], [-0.015559161155025095, 0.051403973219861246, 0.03179031307996347], [-0.02766167769640129, 0.03749545161530447, 0.003335141009473408], [0.05824598730479011, 0.021744367381243884, -0.01578925627951616], [0.05260929401500947, 0.0560165014956886, -0.027477296572565126], [0.018513891242931686, 0.041961785217662514, 0.004490763489747966], [0.024063060899760215, 0.065082853069653, 0.044343437673514896], [0.05250992323006226, 0.04361117432588933, 0.01030076055524387], [0.0038921710021782366, -0.025299228133723792, 0.019370764014574535], [-0.00011950534333568519, 0.06549370069727675, -0.03436712163379723], [-0.026020578032683626, -0.013341758571090847, -0.009119046570271953], [0.024412451175602937, 0.030135064560817174, -0.008355486384198006], [0.04002209845752687, -0.017341304390739463, 0.02818338690302971], [-0.032575108695213684, -0.009588338926775117, -0.03077312160940468]]
#import random
#random.seed(seed)
#latent_rgb_factors = [[random.uniform(min_val, max_val) for _ in range(3)] for _ in range(16)]
out_factors = latent_rgb_factors
print(latent_rgb_factors)
latent_rgb_factors_bias = [0,0,0]
latent_rgb_factors = torch.tensor(latent_rgb_factors, device=latents.device, dtype=latents.dtype).transpose(0, 1)
latent_rgb_factors_bias = torch.tensor(latent_rgb_factors_bias, device=latents.device, dtype=latents.dtype)
print("latent_rgb_factors", latent_rgb_factors.shape)
latent_images = []
for t in range(latents.shape[2]):
latent = latents[:, :, t, :, :]
latent = latent[0].permute(1, 2, 0)
latent_image = torch.nn.functional.linear(
latent,
latent_rgb_factors,
bias=latent_rgb_factors_bias
)
latent_images.append(latent_image)
latent_images = torch.stack(latent_images, dim=0)
print("latent_images", latent_images.shape)
latent_images_min = latent_images.min()
latent_images_max = latent_images.max()
latent_images = (latent_images - latent_images_min) / (latent_images_max - latent_images_min)
return (latent_images.float().cpu(), out_factors)
NODE_CLASS_MAPPINGS = {
"PyramidFlowSampler": PyramidFlowSampler,
@@ -476,7 +552,8 @@ NODE_CLASS_MAPPINGS = {
"PyramidFlowVAEEncode": PyramidFlowVAEEncode,
"PyramidFlowTorchCompileSettings": PyramidFlowTorchCompileSettings,
"PyramidFlowTransformerLoader": PyramidFlowModelLoader,
"PyramidFlowVAELoader": PyramidFlowVAELoader
"PyramidFlowVAELoader": PyramidFlowVAELoader,
"PyramidFlowLatentPreview": PyramidFlowLatentPreview
}
NODE_DISPLAY_NAME_MAPPINGS = {
@@ -487,5 +564,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"PyramidFlowVAEEncode": "PyramidFlow VAE Encode",
"PyramidFlowTorchCompileSettings": "PyramidFlow Torch Compile Settings",
"PyramidFlowTransformerLoader": "PyramidFlow Model Loader",
"PyramidFlowVAELoader": "PyramidFlow VAE Loader"
"PyramidFlowVAELoader": "PyramidFlow VAE Loader",
"PyramidFlowLatentPreview": "PyramidFlow Latent Preview"
}
-1
View File
@@ -1,3 +1,2 @@
from .modeling_pyramid_flux import PyramidFluxTransformer
from .modeling_text_encoder import FluxTextEncoderWithMask
from .modeling_flux_block import FluxSingleTransformerBlock, FluxTransformerBlock
+15 -307
View File
@@ -3,33 +3,26 @@ from typing import Any, Dict, List, Optional, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
import inspect
from einops import rearrange
from diffusers.utils import deprecate
from diffusers.models.activations import GEGLU, GELU, ApproximateGELU, SwiGLU
from .modeling_normalization import (
AdaLayerNormContinuous, AdaLayerNormZero,
AdaLayerNormZeroSingle, FP32LayerNorm, RMSNorm
)
from ...trainer_misc import (
is_sequence_parallel_initialized,
get_sequence_parallel_group,
get_sequence_parallel_world_size,
all_to_all,
)
# try:
# from flash_attn.ops.triton.layer_norm import RMSNorm as FlashRMSNorm #slightly faster
# @torch.compiler.disable() #cause NaNs when compiled for some reason
# class RMSNorm(FlashRMSNorm):
# pass
# except:
from .modeling_normalization import RMSNorm
from .modeling_normalization import (AdaLayerNormZero, AdaLayerNormZeroSingle, FP32LayerNorm)
try:
from flash_attn import flash_attn_qkvpacked_func, flash_attn_func
from flash_attn.bert_padding import pad_input, unpad_input, index_first_axis
from flash_attn.flash_attn_interface import flash_attn_varlen_func
except:
flash_attn_func = None
flash_attn_qkvpacked_func = None
flash_attn_varlen_func = None
def apply_rope(xq, xk, freqs_cis):
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
@@ -100,92 +93,6 @@ class FeedForward(nn.Module):
return hidden_states
class SequenceParallelVarlenFlashSelfAttentionWithT5Mask:
def __init__(self):
pass
def __call__(
self, query, key, value, encoder_query, encoder_key, encoder_value,
heads, scale, hidden_length=None, image_rotary_emb=None, encoder_attention_mask=None,
):
assert encoder_attention_mask is not None, "The encoder-hidden mask needed to be set"
batch_size = query.shape[0]
qkv_list = []
num_stages = len(hidden_length)
encoder_qkv = torch.stack([encoder_query, encoder_key, encoder_value], dim=2) # [bs, sub_seq, 3, head, head_dim]
qkv = torch.stack([query, key, value], dim=2) # [bs, sub_seq, 3, head, head_dim]
# To sync the encoder query, key and values
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
encoder_qkv = all_to_all(encoder_qkv, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
output_hidden = torch.zeros_like(qkv[:,:,0])
output_encoder_hidden = torch.zeros_like(encoder_qkv[:,:,0])
encoder_length = encoder_qkv.shape[1]
i_sum = 0
for i_p, length in enumerate(hidden_length):
# get the query, key, value from padding sequence
encoder_qkv_tokens = encoder_qkv[i_p::num_stages]
qkv_tokens = qkv[:, i_sum:i_sum+length]
qkv_tokens = all_to_all(qkv_tokens, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
concat_qkv_tokens = torch.cat([encoder_qkv_tokens, qkv_tokens], dim=1) # [bs, pad_seq, 3, nhead, dim]
if image_rotary_emb is not None:
concat_qkv_tokens[:,:,0], concat_qkv_tokens[:,:,1] = apply_rope(concat_qkv_tokens[:,:,0], concat_qkv_tokens[:,:,1], image_rotary_emb[i_p])
indices = encoder_attention_mask[i_p]['indices']
qkv_list.append(index_first_axis(rearrange(concat_qkv_tokens, "b s ... -> (b s) ..."), indices))
i_sum += length
token_lengths = [x_.shape[0] for x_ in qkv_list]
qkv = torch.cat(qkv_list, dim=0)
query, key, value = qkv.unbind(1)
cu_seqlens = torch.cat([x_['seqlens_in_batch'] for x_ in encoder_attention_mask], dim=0)
max_seqlen_q = cu_seqlens.max().item()
max_seqlen_k = max_seqlen_q
cu_seqlens_q = F.pad(torch.cumsum(cu_seqlens, dim=0, dtype=torch.int32), (1, 0))
cu_seqlens_k = cu_seqlens_q.clone()
output = flash_attn_varlen_func(
query,
key,
value,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
max_seqlen_q=max_seqlen_q,
max_seqlen_k=max_seqlen_k,
dropout_p=0.0,
causal=False,
softmax_scale=scale,
)
# To merge the tokens
i_sum = 0;token_sum = 0
for i_p, length in enumerate(hidden_length):
tot_token_num = token_lengths[i_p]
stage_output = output[token_sum : token_sum + tot_token_num]
stage_output = pad_input(stage_output, encoder_attention_mask[i_p]['indices'], batch_size, encoder_length + length * sp_group_size)
stage_encoder_hidden_output = stage_output[:, :encoder_length]
stage_hidden_output = stage_output[:, encoder_length:]
stage_hidden_output = all_to_all(stage_hidden_output, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_hidden[:, i_sum:i_sum+length] = stage_hidden_output
output_encoder_hidden[i_p::num_stages] = stage_encoder_hidden_output
token_sum += tot_token_num
i_sum += length
output_encoder_hidden = all_to_all(output_encoder_hidden, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_hidden = output_hidden.flatten(2, 3)
output_encoder_hidden = output_encoder_hidden.flatten(2, 3)
return output_hidden, output_encoder_hidden
class VarlenFlashSelfAttentionWithT5Mask:
def __init__(self):
@@ -262,69 +169,6 @@ class VarlenFlashSelfAttentionWithT5Mask:
return output_hidden, output_encoder_hidden
class SequenceParallelVarlenSelfAttentionWithT5Mask:
def __init__(self):
pass
def __call__(
self, query, key, value, encoder_query, encoder_key, encoder_value,
heads, scale, hidden_length=None, image_rotary_emb=None, attention_mask=None,
):
assert attention_mask is not None, "The attention mask needed to be set"
num_stages = len(hidden_length)
encoder_qkv = torch.stack([encoder_query, encoder_key, encoder_value], dim=2) # [bs, sub_seq, 3, head, head_dim]
qkv = torch.stack([query, key, value], dim=2) # [bs, sub_seq, 3, head, head_dim]
# To sync the encoder query, key and values
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
encoder_qkv = all_to_all(encoder_qkv, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
encoder_length = encoder_qkv.shape[1]
i_sum = 0
output_encoder_hidden_list = []
output_hidden_list = []
for i_p, length in enumerate(hidden_length):
encoder_qkv_tokens = encoder_qkv[i_p::num_stages]
qkv_tokens = qkv[:, i_sum:i_sum+length]
qkv_tokens = all_to_all(qkv_tokens, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
concat_qkv_tokens = torch.cat([encoder_qkv_tokens, qkv_tokens], dim=1) # [bs, tot_seq, 3, nhead, dim]
if image_rotary_emb is not None:
concat_qkv_tokens[:,:,0], concat_qkv_tokens[:,:,1] = apply_rope(concat_qkv_tokens[:,:,0], concat_qkv_tokens[:,:,1], image_rotary_emb[i_p])
query, key, value = concat_qkv_tokens.unbind(2) # [bs, tot_seq, nhead, dim]
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
stage_hidden_states = F.scaled_dot_product_attention(
query, key, value, dropout_p=0.0, is_causal=False, attn_mask=attention_mask[i_p],
)
stage_hidden_states = stage_hidden_states.transpose(1, 2) # [bs, tot_seq, nhead, dim]
output_encoder_hidden_list.append(stage_hidden_states[:, :encoder_length])
output_hidden = stage_hidden_states[:, encoder_length:]
output_hidden = all_to_all(output_hidden, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_hidden_list.append(output_hidden)
i_sum += length
output_encoder_hidden = torch.stack(output_encoder_hidden_list, dim=1) # [b n s nhead d]
output_encoder_hidden = rearrange(output_encoder_hidden, 'b n s h d -> (b n) s h d')
output_encoder_hidden = all_to_all(output_encoder_hidden, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_encoder_hidden = output_encoder_hidden.flatten(2, 3)
output_hidden = torch.cat(output_hidden_list, dim=1).flatten(2, 3)
return output_hidden, output_encoder_hidden
class VarlenSelfAttentionWithT5Mask:
def __init__(self):
@@ -375,80 +219,6 @@ class VarlenSelfAttentionWithT5Mask:
return output_hidden, output_encoder_hidden
class SequenceParallelVarlenFlashAttnSingle:
def __init__(self):
pass
def __call__(
self, query, key, value, heads, scale,
hidden_length=None, image_rotary_emb=None, encoder_attention_mask=None,
):
assert encoder_attention_mask is not None, "The encoder-hidden mask needed to be set"
batch_size = query.shape[0]
qkv_list = []
num_stages = len(hidden_length)
qkv = torch.stack([query, key, value], dim=2) # [bs, sub_seq, 3, head, head_dim]
output_hidden = torch.zeros_like(qkv[:,:,0])
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
i_sum = 0
for i_p, length in enumerate(hidden_length):
# get the query, key, value from padding sequence
qkv_tokens = qkv[:, i_sum:i_sum+length]
qkv_tokens = all_to_all(qkv_tokens, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
if image_rotary_emb is not None:
qkv_tokens[:,:,0], qkv_tokens[:,:,1] = apply_rope(qkv_tokens[:,:,0], qkv_tokens[:,:,1], image_rotary_emb[i_p])
indices = encoder_attention_mask[i_p]['indices']
qkv_list.append(index_first_axis(rearrange(qkv_tokens, "b s ... -> (b s) ..."), indices))
i_sum += length
token_lengths = [x_.shape[0] for x_ in qkv_list]
qkv = torch.cat(qkv_list, dim=0)
query, key, value = qkv.unbind(1)
cu_seqlens = torch.cat([x_['seqlens_in_batch'] for x_ in encoder_attention_mask], dim=0)
max_seqlen_q = cu_seqlens.max().item()
max_seqlen_k = max_seqlen_q
cu_seqlens_q = F.pad(torch.cumsum(cu_seqlens, dim=0, dtype=torch.int32), (1, 0))
cu_seqlens_k = cu_seqlens_q.clone()
output = flash_attn_varlen_func(
query,
key,
value,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
max_seqlen_q=max_seqlen_q,
max_seqlen_k=max_seqlen_k,
dropout_p=0.0,
causal=False,
softmax_scale=scale,
)
# To merge the tokens
i_sum = 0;token_sum = 0
for i_p, length in enumerate(hidden_length):
tot_token_num = token_lengths[i_p]
stage_output = output[token_sum : token_sum + tot_token_num]
stage_output = pad_input(stage_output, encoder_attention_mask[i_p]['indices'], batch_size, length * sp_group_size)
stage_hidden_output = all_to_all(stage_output, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_hidden[:, i_sum:i_sum+length] = stage_hidden_output
token_sum += tot_token_num
i_sum += length
output_hidden = output_hidden.flatten(2, 3)
return output_hidden
class VarlenFlashSelfAttnSingle:
def __init__(self):
@@ -515,56 +285,6 @@ class VarlenFlashSelfAttnSingle:
return output_hidden
class SequenceParallelVarlenAttnSingle:
def __init__(self):
pass
def __call__(
self, query, key, value, heads, scale,
hidden_length=None, image_rotary_emb=None, attention_mask=None,
):
assert attention_mask is not None, "The attention mask needed to be set"
num_stages = len(hidden_length)
qkv = torch.stack([query, key, value], dim=2) # [bs, sub_seq, 3, head, head_dim]
# To sync the encoder query, key and values
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
i_sum = 0
output_hidden_list = []
for i_p, length in enumerate(hidden_length):
qkv_tokens = qkv[:, i_sum:i_sum+length]
qkv_tokens = all_to_all(qkv_tokens, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
if image_rotary_emb is not None:
qkv_tokens[:,:,0], qkv_tokens[:,:,1] = apply_rope(qkv_tokens[:,:,0], qkv_tokens[:,:,1], image_rotary_emb[i_p])
query, key, value = qkv_tokens.unbind(2) # [bs, tot_seq, nhead, dim]
query = query.transpose(1, 2).contiguous()
key = key.transpose(1, 2).contiguous()
value = value.transpose(1, 2).contiguous()
stage_hidden_states = F.scaled_dot_product_attention(
query, key, value, dropout_p=0.0, is_causal=False, attn_mask=attention_mask[i_p],
)
stage_hidden_states = stage_hidden_states.transpose(1, 2) # [bs, tot_seq, nhead, dim]
output_hidden = stage_hidden_states
output_hidden = all_to_all(output_hidden, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_hidden_list.append(output_hidden)
i_sum += length
output_hidden = torch.cat(output_hidden_list, dim=1).flatten(2, 3)
return output_hidden
class VarlenSelfAttnSingle:
def __init__(self):
@@ -575,8 +295,7 @@ class VarlenSelfAttnSingle:
hidden_length=None, image_rotary_emb=None, attention_mask=None,
):
assert attention_mask is not None, "The attention mask needed to be set"
num_stages = len(hidden_length)
qkv = torch.stack([query, key, value], dim=2) # [bs, sub_seq, 3, head, head_dim]
i_sum = 0
@@ -732,15 +451,9 @@ class FluxSingleAttnProcessor2_0:
self.use_flash_attn = use_flash_attn
if self.use_flash_attn:
if is_sequence_parallel_initialized():
self.varlen_flash_attn = SequenceParallelVarlenFlashAttnSingle()
else:
self.varlen_flash_attn = VarlenFlashSelfAttnSingle()
self.varlen_flash_attn = VarlenFlashSelfAttnSingle()
else:
if is_sequence_parallel_initialized():
self.varlen_attn = SequenceParallelVarlenAttnSingle()
else:
self.varlen_attn = VarlenSelfAttnSingle()
self.varlen_attn = VarlenSelfAttnSingle()
def __call__(
self,
@@ -792,15 +505,9 @@ class FluxAttnProcessor2_0:
self.use_flash_attn = use_flash_attn
if self.use_flash_attn:
if is_sequence_parallel_initialized():
self.varlen_flash_attn = SequenceParallelVarlenFlashSelfAttentionWithT5Mask()
else:
self.varlen_flash_attn = VarlenFlashSelfAttentionWithT5Mask()
self.varlen_flash_attn = VarlenFlashSelfAttentionWithT5Mask()
else:
if is_sequence_parallel_initialized():
self.varlen_attn = SequenceParallelVarlenSelfAttentionWithT5Mask()
else:
self.varlen_attn = VarlenSelfAttentionWithT5Mask()
self.varlen_attn = VarlenSelfAttentionWithT5Mask()
def __call__(
self,
@@ -813,6 +520,7 @@ class FluxAttnProcessor2_0:
image_rotary_emb: Optional[torch.Tensor] = None,
) -> torch.FloatTensor:
# `sample` projections.
query = attn.to_q(hidden_states)
key = attn.to_k(hidden_states)
value = attn.to_v(hidden_states)
+20 -152
View File
@@ -1,30 +1,17 @@
from typing import Any, Dict, List, Optional, Union
import torch
import os
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from tqdm import tqdm
from diffusers.utils.torch_utils import randn_tensor
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.models.modeling_utils import ModelMixin
from diffusers.utils import is_torch_version
from .modeling_normalization import AdaLayerNormContinuous
from .modeling_embedding import CombinedTimestepGuidanceTextProjEmbeddings, CombinedTimestepTextProjEmbeddings
from .modeling_embedding import CombinedTimestepTextProjEmbeddings
from .modeling_flux_block import FluxTransformerBlock, FluxSingleTransformerBlock
from ...trainer_misc import (
is_sequence_parallel_initialized,
get_sequence_parallel_group,
get_sequence_parallel_world_size,
get_sequence_parallel_rank,
all_to_all,
)
def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor:
assert dim % 2 == 0, "The dimension must be even."
@@ -269,13 +256,6 @@ class PyramidFluxTransformer(ModelMixin, ConfigMixin):
input_ids_list = [torch.cat([text_ids, image_ids], dim=1) for image_ids in image_ids_list]
image_rotary_emb = [self.pos_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()
concat_output = True if self.training else False
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, concat_output=concat_output) 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, concat_output=concat_output) for input_ids in input_ids_list]
hidden_states, hidden_length = [], []
for sample_ in sample:
@@ -298,12 +278,6 @@ class PyramidFluxTransformer(ModelMixin, ConfigMixin):
for i_p, length in enumerate(hidden_length):
pad_attention_mask = torch.ones((pad_batch_size, length), dtype=encoder_attention_mask.dtype).to(device)
pad_attention_mask = torch.cat([encoder_attention_mask[i_p::num_stages], pad_attention_mask], dim=1)
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
pad_attention_mask = all_to_all(pad_attention_mask.unsqueeze(2).repeat(1, 1, sp_group_size), sp_group, sp_group_size, scatter_dim=2, gather_dim=0)
pad_attention_mask = pad_attention_mask.squeeze(2)
seqlens_in_batch = pad_attention_mask.sum(dim=-1, dtype=torch.int32)
indices = torch.nonzero(pad_attention_mask.flatten(), as_tuple=False).flatten()
@@ -330,13 +304,6 @@ class PyramidFluxTransformer(ModelMixin, ConfigMixin):
image_ids_list = []
for i_p, length in enumerate(hidden_length):
image_ids_list.append(image_ids[i_p::num_stages][:, :length])
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
concat_output = True if self.training else False
text_ids = all_to_all(text_ids.unsqueeze(2).repeat(1, 1, sp_group_size), sp_group, sp_group_size, scatter_dim=2, gather_dim=0, concat_output=concat_output).squeeze(2)
image_ids_list = [all_to_all(image_ids_.unsqueeze(2).repeat(1, 1, sp_group_size), sp_group, sp_group_size, scatter_dim=2, gather_dim=0, concat_output=concat_output).squeeze(2) for image_ids_ in image_ids_list]
attention_mask = []
for i_p in range(len(hidden_length)):
@@ -357,20 +324,11 @@ class PyramidFluxTransformer(ModelMixin, ConfigMixin):
output_hidden_list = []
batch_hidden_states = torch.split(batch_hidden_states, hidden_length, dim=1)
if is_sequence_parallel_initialized():
sp_group_size = get_sequence_parallel_world_size()
batch_size = batch_size // sp_group_size
for i_p, length in enumerate(hidden_length):
width, height, temp = widths[i_p], heights[i_p], temps[i_p]
trainable_token_num = trainable_token_list[i_p]
hidden_states = batch_hidden_states[i_p]
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
hidden_states = all_to_all(hidden_states, sp_group, sp_group_size, scatter_dim=0, gather_dim=1)
# only the trainable token are taking part in loss computation
hidden_states = hidden_states[:, -trainable_token_num:]
@@ -400,134 +358,44 @@ class PyramidFluxTransformer(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()
concat_output = True if self.training else False
# 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, concat_output=concat_output)
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, concat_output=concat_output)
temb = all_to_all(temb.unsqueeze(1).repeat(1, sp_group_size, 1), sp_group, sp_group_size, scatter_dim=1, gather_dim=0, concat_output=concat_output)
temb = temb.squeeze(1)
else:
hidden_states = torch.cat(hidden_states, dim=1)
hidden_states = torch.cat(hidden_states, dim=1)
for index_block, block in enumerate(self.transformer_blocks):
if self.training and self.gradient_checkpointing and (index_block <= int(len(self.transformer_blocks) * self.gradient_checkpointing_ratio)):
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,
)
# remerge for single attention block
num_stages = len(hidden_length)
batch_hidden_states = list(torch.split(hidden_states, hidden_length, dim=1))
concat_hidden_length = []
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
encoder_hidden_states = all_to_all(encoder_hidden_states, sp_group, sp_group_size, scatter_dim=0, gather_dim=1)
for i_p in range(len(hidden_length)):
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
batch_hidden_states[i_p] = all_to_all(batch_hidden_states[i_p], sp_group, sp_group_size, scatter_dim=0, gather_dim=1)
batch_hidden_states[i_p] = torch.cat([encoder_hidden_states[i_p::num_stages], batch_hidden_states[i_p]], dim=1)
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
batch_hidden_states[i_p] = all_to_all(batch_hidden_states[i_p], sp_group, sp_group_size, scatter_dim=1, gather_dim=0)
concat_hidden_length.append(batch_hidden_states[i_p].shape[1])
hidden_states = torch.cat(batch_hidden_states, dim=1)
for index_block, block in enumerate(self.single_transformer_blocks):
if self.training and self.gradient_checkpointing and (index_block <= int(len(self.single_transformer_blocks) * self.gradient_checkpointing_ratio)):
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 {}
hidden_states = torch.utils.checkpoint.checkpoint(
create_custom_forward(block),
hidden_states,
temb,
encoder_attention_mask,
attention_mask,
concat_hidden_length,
image_rotary_emb,
**ckpt_kwargs,
)
else:
hidden_states = block(
hidden_states=hidden_states,
temb=temb,
encoder_attention_mask=encoder_attention_mask, # used for
attention_mask=attention_mask,
hidden_length=concat_hidden_length,
image_rotary_emb=image_rotary_emb,
)
hidden_states = block(
hidden_states=hidden_states,
temb=temb,
encoder_attention_mask=encoder_attention_mask,
attention_mask=attention_mask,
hidden_length=concat_hidden_length,
image_rotary_emb=image_rotary_emb,
)
batch_hidden_states = list(torch.split(hidden_states, concat_hidden_length, dim=1))
for i_p in range(len(concat_hidden_length)):
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
batch_hidden_states[i_p] = all_to_all(batch_hidden_states[i_p], sp_group, sp_group_size, scatter_dim=0, gather_dim=1)
for i_p in range(len(concat_hidden_length)):
batch_hidden_states[i_p] = batch_hidden_states[i_p][:, encoder_hidden_length :, ...]
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
batch_hidden_states[i_p] = all_to_all(batch_hidden_states[i_p], sp_group, sp_group_size, scatter_dim=1, gather_dim=0)
hidden_states = torch.cat(batch_hidden_states, dim=1)
hidden_states = self.norm_out(hidden_states, temb, hidden_length=hidden_length)
@@ -1,141 +0,0 @@
import torch
import torch.nn as nn
import os
from transformers import (
CLIPTextModel,
CLIPTokenizer,
T5EncoderModel,
T5TokenizerFast,
)
from typing import Any, Callable, Dict, List, Optional, Union
class FluxTextEncoderWithMask(nn.Module):
def __init__(self, model_path, torch_dtype):
super().__init__()
# CLIP-G
self.tokenizer = CLIPTokenizer.from_pretrained(os.path.join(model_path, 'tokenizer'), torch_dtype=torch_dtype)
self.tokenizer_max_length = (
self.tokenizer.model_max_length if hasattr(self, "tokenizer") and self.tokenizer is not None else 77
)
self.text_encoder = CLIPTextModel.from_pretrained(os.path.join(model_path, 'text_encoder'), torch_dtype=torch_dtype)
# T5
self.tokenizer_2 = T5TokenizerFast.from_pretrained(os.path.join(model_path, 'tokenizer_2'))
self.text_encoder_2 = T5EncoderModel.from_pretrained(os.path.join(model_path, 'text_encoder_2'), torch_dtype=torch_dtype)
self._freeze()
def _freeze(self):
for param in self.parameters():
param.requires_grad = False
def _get_t5_prompt_embeds(
self,
prompt: Union[str, List[str]] = None,
num_images_per_prompt: int = 1,
max_sequence_length: int = 128,
device: Optional[torch.device] = None,
):
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
text_inputs = self.tokenizer_2(
prompt,
padding="max_length",
max_length=max_sequence_length,
truncation=True,
return_length=False,
return_overflowing_tokens=False,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
prompt_attention_mask = text_inputs.attention_mask
prompt_attention_mask = prompt_attention_mask.to(device)
prompt_embeds = self.text_encoder_2(text_input_ids.to(device), attention_mask=prompt_attention_mask, output_hidden_states=False)[0]
dtype = self.text_encoder_2.dtype
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
_, seq_len, _ = prompt_embeds.shape
# duplicate text embeddings and attention mask for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1)
prompt_embeds = prompt_embeds.view(batch_size * num_images_per_prompt, seq_len, -1)
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
prompt_attention_mask = prompt_attention_mask.repeat(num_images_per_prompt, 1)
return prompt_embeds, prompt_attention_mask
def _get_clip_prompt_embeds(
self,
prompt: Union[str, List[str]],
num_images_per_prompt: int = 1,
device: Optional[torch.device] = None,
):
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
text_inputs = self.tokenizer(
prompt,
padding="max_length",
max_length=self.tokenizer_max_length,
truncation=True,
return_overflowing_tokens=False,
return_length=False,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
prompt_embeds = self.text_encoder(text_input_ids.to(device), output_hidden_states=False)
# Use pooled output of CLIPTextModel
prompt_embeds = prompt_embeds.pooler_output
prompt_embeds = prompt_embeds.to(dtype=self.text_encoder.dtype, device=device)
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt)
prompt_embeds = prompt_embeds.view(batch_size * num_images_per_prompt, -1)
return prompt_embeds
def encode_prompt(self,
prompt,
num_images_per_prompt=1,
device=None,
):
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
pooled_prompt_embeds = self._get_clip_prompt_embeds(
prompt=prompt,
device=device,
num_images_per_prompt=num_images_per_prompt,
)
prompt_embeds, prompt_attention_mask = self._get_t5_prompt_embeds(
prompt=prompt,
num_images_per_prompt=num_images_per_prompt,
device=device,
)
print("prompt_embeds_shape: ",prompt_embeds.shape)
print("pooled_prompt_embeds_shape: ",pooled_prompt_embeds.shape)
print("prompt_attention_mask_shape: ",prompt_attention_mask.shape)
# prompt_embeds_shape: torch.Size([1, 128, 4096])
# pooled_prompt_embeds_shape: torch.Size([1, 768])
# prompt_attention_mask_shape: torch.Size([1, 128])
return prompt_embeds, prompt_attention_mask, pooled_prompt_embeds
def forward(self, input_prompts, device):
with torch.no_grad():
prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = self.encode_prompt(input_prompts, 1, device=device)
return prompt_embeds, prompt_attention_mask, pooled_prompt_embeds
+1 -2
View File
@@ -1,2 +1 @@
from .modeling_pyramid_mmdit import PyramidDiffusionMMDiT
from .modeling_text_encoder import SD3TextEncoderWithMask
from .modeling_pyramid_mmdit import PyramidDiffusionMMDiT
@@ -15,13 +15,6 @@ except:
flash_attn_varlen_func = None
print("Please install flash attention")
from ...trainer_misc import (
is_sequence_parallel_initialized,
get_sequence_parallel_group,
get_sequence_parallel_world_size,
all_to_all,
)
from .modeling_normalization import AdaLayerNormZero, AdaLayerNormContinuous, RMSNorm
@@ -167,99 +160,6 @@ class VarlenFlashSelfAttentionWithT5Mask:
return output_hidden, output_encoder_hidden
class SequenceParallelVarlenFlashSelfAttentionWithT5Mask:
def __init__(self):
pass
def apply_rope(self, xq, xk, freqs_cis):
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)
xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]
xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]
return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)
def __call__(
self, query, key, value, encoder_query, encoder_key, encoder_value,
heads, scale, hidden_length=None, image_rotary_emb=None, encoder_attention_mask=None,
):
assert encoder_attention_mask is not None, "The encoder-hidden mask needed to be set"
batch_size = query.shape[0]
qkv_list = []
num_stages = len(hidden_length)
encoder_qkv = torch.stack([encoder_query, encoder_key, encoder_value], dim=2) # [bs, sub_seq, 3, head, head_dim]
qkv = torch.stack([query, key, value], dim=2) # [bs, sub_seq, 3, head, head_dim]
# To sync the encoder query, key and values
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
encoder_qkv = all_to_all(encoder_qkv, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
output_hidden = torch.zeros_like(qkv[:,:,0])
output_encoder_hidden = torch.zeros_like(encoder_qkv[:,:,0])
encoder_length = encoder_qkv.shape[1]
i_sum = 0
for i_p, length in enumerate(hidden_length):
# get the query, key, value from padding sequence
encoder_qkv_tokens = encoder_qkv[i_p::num_stages]
qkv_tokens = qkv[:, i_sum:i_sum+length]
qkv_tokens = all_to_all(qkv_tokens, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
concat_qkv_tokens = torch.cat([encoder_qkv_tokens, qkv_tokens], dim=1) # [bs, pad_seq, 3, nhead, dim]
if image_rotary_emb is not None:
concat_qkv_tokens[:,:,0], concat_qkv_tokens[:,:,1] = self.apply_rope(concat_qkv_tokens[:,:,0], concat_qkv_tokens[:,:,1], image_rotary_emb[i_p])
indices = encoder_attention_mask[i_p]['indices']
qkv_list.append(index_first_axis(rearrange(concat_qkv_tokens, "b s ... -> (b s) ..."), indices))
i_sum += length
token_lengths = [x_.shape[0] for x_ in qkv_list]
qkv = torch.cat(qkv_list, dim=0)
query, key, value = qkv.unbind(1)
cu_seqlens = torch.cat([x_['seqlens_in_batch'] for x_ in encoder_attention_mask], dim=0)
max_seqlen_q = cu_seqlens.max().item()
max_seqlen_k = max_seqlen_q
cu_seqlens_q = F.pad(torch.cumsum(cu_seqlens, dim=0, dtype=torch.int32), (1, 0))
cu_seqlens_k = cu_seqlens_q.clone()
output = flash_attn_varlen_func(
query,
key,
value,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
max_seqlen_q=max_seqlen_q,
max_seqlen_k=max_seqlen_k,
dropout_p=0.0,
causal=False,
softmax_scale=scale,
)
# To merge the tokens
i_sum = 0;token_sum = 0
for i_p, length in enumerate(hidden_length):
tot_token_num = token_lengths[i_p]
stage_output = output[token_sum : token_sum + tot_token_num]
stage_output = pad_input(stage_output, encoder_attention_mask[i_p]['indices'], batch_size, encoder_length + length * sp_group_size)
stage_encoder_hidden_output = stage_output[:, :encoder_length]
stage_hidden_output = stage_output[:, encoder_length:]
stage_hidden_output = all_to_all(stage_hidden_output, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_hidden[:, i_sum:i_sum+length] = stage_hidden_output
output_encoder_hidden[i_p::num_stages] = stage_encoder_hidden_output
token_sum += tot_token_num
i_sum += length
output_encoder_hidden = all_to_all(output_encoder_hidden, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_hidden = output_hidden.flatten(2, 3)
output_encoder_hidden = output_encoder_hidden.flatten(2, 3)
return output_hidden, output_encoder_hidden
class VarlenSelfAttentionWithT5Mask:
"""
@@ -321,79 +221,6 @@ class VarlenSelfAttentionWithT5Mask:
return output_hidden, output_encoder_hidden
class SequenceParallelVarlenSelfAttentionWithT5Mask:
"""
For chunk stage attention without using flash attention
"""
def __init__(self):
pass
def apply_rope(self, xq, xk, freqs_cis):
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)
xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]
xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]
return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)
def __call__(
self, query, key, value, encoder_query, encoder_key, encoder_value,
heads, scale, hidden_length=None, image_rotary_emb=None, attention_mask=None,
):
assert attention_mask is not None, "The attention mask needed to be set"
num_stages = len(hidden_length)
encoder_qkv = torch.stack([encoder_query, encoder_key, encoder_value], dim=2) # [bs, sub_seq, 3, head, head_dim]
qkv = torch.stack([query, key, value], dim=2) # [bs, sub_seq, 3, head, head_dim]
# To sync the encoder query, key and values
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
encoder_qkv = all_to_all(encoder_qkv, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
encoder_length = encoder_qkv.shape[1]
i_sum = 0
output_encoder_hidden_list = []
output_hidden_list = []
for i_p, length in enumerate(hidden_length):
encoder_qkv_tokens = encoder_qkv[i_p::num_stages]
qkv_tokens = qkv[:, i_sum:i_sum+length]
qkv_tokens = all_to_all(qkv_tokens, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
concat_qkv_tokens = torch.cat([encoder_qkv_tokens, qkv_tokens], dim=1) # [bs, tot_seq, 3, nhead, dim]
if image_rotary_emb is not None:
concat_qkv_tokens[:,:,0], concat_qkv_tokens[:,:,1] = self.apply_rope(concat_qkv_tokens[:,:,0], concat_qkv_tokens[:,:,1], image_rotary_emb[i_p])
query, key, value = concat_qkv_tokens.unbind(2) # [bs, tot_seq, nhead, dim]
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
stage_hidden_states = F.scaled_dot_product_attention(
query, key, value, dropout_p=0.0, is_causal=False, attn_mask=attention_mask[i_p],
)
stage_hidden_states = stage_hidden_states.transpose(1, 2) # [bs, tot_seq, nhead, dim]
output_encoder_hidden_list.append(stage_hidden_states[:, :encoder_length])
output_hidden = stage_hidden_states[:, encoder_length:]
output_hidden = all_to_all(output_hidden, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_hidden_list.append(output_hidden)
i_sum += length
output_encoder_hidden = torch.stack(output_encoder_hidden_list, dim=1) # [b n s nhead d]
output_encoder_hidden = rearrange(output_encoder_hidden, 'b n s h d -> (b n) s h d')
output_encoder_hidden = all_to_all(output_encoder_hidden, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_encoder_hidden = output_encoder_hidden.flatten(2, 3)
output_hidden = torch.cat(output_hidden_list, dim=1).flatten(2, 3)
return output_hidden, output_encoder_hidden
class JointAttention(nn.Module):
def __init__(
@@ -476,14 +303,8 @@ 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()
else:
#if is_sequence_parallel_initialized():
#self.var_len_attn = SequenceParallelVarlenSelfAttentionWithT5Mask()
#else:
self.var_len_attn = VarlenSelfAttentionWithT5Mask()
@@ -13,16 +13,6 @@ from .modeling_embedding import PatchEmbed3D, CombinedTimestepConditionEmbedding
from .modeling_normalization import AdaLayerNormContinuous
from .modeling_mmdit_block import JointTransformerBlock
from ...trainer_misc import (
is_sequence_parallel_initialized,
get_sequence_parallel_group,
get_sequence_parallel_world_size,
get_sequence_parallel_rank,
all_to_all,
)
#from IPython import embed
def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor:
assert dim % 2 == 0, "The dimension must be even."
@@ -315,12 +305,6 @@ class PyramidDiffusionMMDiT(ModelMixin, ConfigMixin):
for i_p, length in enumerate(hidden_length):
pad_attention_mask = torch.ones((pad_batch_size, length), dtype=encoder_attention_mask.dtype).to(device)
pad_attention_mask = torch.cat([encoder_attention_mask[i_p::num_stages], pad_attention_mask], dim=1)
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
pad_attention_mask = all_to_all(pad_attention_mask.unsqueeze(2).repeat(1, 1, sp_group_size), sp_group, sp_group_size, scatter_dim=2, gather_dim=0)
pad_attention_mask = pad_attention_mask.squeeze(2)
seqlens_in_batch = pad_attention_mask.sum(dim=-1, dtype=torch.int32)
indices = torch.nonzero(pad_attention_mask.flatten(), as_tuple=False).flatten()
@@ -347,12 +331,6 @@ class PyramidDiffusionMMDiT(ModelMixin, ConfigMixin):
for i_p, length in enumerate(hidden_length):
image_ids_list.append(image_ids[i_p::num_stages][:, :length])
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
text_ids = all_to_all(text_ids.unsqueeze(2).repeat(1, 1, sp_group_size), sp_group, sp_group_size, scatter_dim=2, gather_dim=0).squeeze(2)
image_ids_list = [all_to_all(image_ids_.unsqueeze(2).repeat(1, 1, sp_group_size), sp_group, sp_group_size, scatter_dim=2, gather_dim=0).squeeze(2) for image_ids_ in image_ids_list]
attention_mask = []
for i_p in range(len(hidden_length)):
image_ids = image_ids_list[i_p]
@@ -372,20 +350,11 @@ class PyramidDiffusionMMDiT(ModelMixin, ConfigMixin):
output_hidden_list = []
batch_hidden_states = torch.split(batch_hidden_states, hidden_length, dim=1)
if is_sequence_parallel_initialized():
sp_group_size = get_sequence_parallel_world_size()
batch_size = batch_size // sp_group_size
for i_p, length in enumerate(hidden_length):
width, height, temp = widths[i_p], heights[i_p], temps[i_p]
trainable_token_num = trainable_token_list[i_p]
hidden_states = batch_hidden_states[i_p]
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
hidden_states = all_to_all(hidden_states, sp_group, sp_group_size, scatter_dim=0, gather_dim=1)
# only the trainable token are taking part in loss computation
hidden_states = hidden_states[:, -trainable_token_num:]
@@ -1,3 +1,5 @@
from types import SimpleNamespace
import torch
import torch.nn.functional as F
@@ -7,7 +9,7 @@ from diffusers.utils.torch_utils import randn_tensor
import math
from tqdm import tqdm
from typing import List, Optional, Union
from typing import List, Optional, Union, Callable
from ..diffusion_schedulers import PyramidFlowMatchEulerDiscreteScheduler
from accelerate import cpu_offload
from comfy.utils import ProgressBar
@@ -56,6 +58,10 @@ class PyramidDiTForVideoGeneration:
self.device = main_device
self.sequential_offload_enabled = False
from comfy import latent_formats
self.model = SimpleNamespace(latent_format=latent_formats.Flux())
self.load_device = main_device
if model_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
self.dtype = torch.bfloat16
else:
@@ -89,6 +95,7 @@ class PyramidDiTForVideoGeneration:
)
#round the gamma as 1/3 seems to have issues on some systems
self.gamma = round(self.scheduler.config.gamma, 5)
self.dist = torch.distributions.MultivariateNormal(torch.zeros(4), torch.eye(4) * (1 + self.gamma) - torch.ones(4, 4) * self.gamma)
print(f"The start sigmas and end sigmas of each stage is Start: {self.scheduler.start_sigmas}, End: {self.scheduler.end_sigmas}, Ori_start: {self.scheduler.ori_start_sigmas}")
@@ -142,14 +149,16 @@ class PyramidDiTForVideoGeneration:
)
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
return latents
def sample_block_noise(self, bs, ch, temp, height, width):
dist = torch.distributions.multivariate_normal.MultivariateNormal(torch.zeros(4), torch.eye(4) * (1 + self.gamma) - torch.ones(4, 4) * self.gamma)
block_number = bs * ch * temp * (height // 2) * (width // 2)
noise = torch.stack([dist.sample() for _ in range(block_number)]) # [block number, 4]
noise = torch.stack([self.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(
self,
@@ -166,6 +175,7 @@ class PyramidDiTForVideoGeneration:
dtype,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
is_first_frame: bool = False,
callback = None,
):
stages = self.stages
intermed_latents = []
@@ -198,7 +208,6 @@ class PyramidDiTForVideoGeneration:
timestep = t.expand(latent_model_input.shape[0]).to(latent_model_input.dtype)
latent_model_input = past_conditions[i_s] + [latent_model_input]
noise_pred = self.dit(
sample=[latent_model_input],
timestep_ratio=timestep,
@@ -208,10 +217,6 @@ 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:
@@ -228,11 +233,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)
return intermed_latents
@@ -253,32 +256,18 @@ class PyramidDiTForVideoGeneration:
alpha: float = 0.5,
num_images_per_prompt: Optional[int] = 1,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
output_type: Optional[str] = "pil",
callback: Optional[Callable] = None,
):
#device = self.device
dtype = self.dtype
assert temp % self.frame_per_unit == 0, "The frames should be divided by frame_per unit"
batch_size = prompt_embeds_dict['prompt_embeds'].shape[0]
# if isinstance(prompt, str):
# batch_size = 1
# prompt = prompt + ", hyper quality, Ultra HD, 8K" # adding this prompt to improve aesthetics
# else:
# assert isinstance(prompt, list)
# batch_size = len(prompt)
# prompt = [_ + ", hyper quality, Ultra HD, 8K" for _ in prompt]
if isinstance(num_inference_steps, int):
num_inference_steps = [num_inference_steps] * len(self.stages)
elif isinstance(num_inference_steps, list) and len(num_inference_steps) < len(self.stages):
num_inference_steps = (num_inference_steps * len(self.stages))[:len(self.stages)]
# negative_prompt = negative_prompt or ""
# # Get the text embeddings
# prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = self.text_encoder(prompt, device)
# negative_prompt_embeds, negative_prompt_attention_mask, negative_pooled_prompt_embeds = self.text_encoder(negative_prompt, device)
if use_linear_guidance:
max_guidance_scale = guidance_scale
guidance_scale_list = [max(max_guidance_scale - alpha * t_, min_guidance_scale) for t_ in range(temp+1)]
@@ -300,11 +289,6 @@ 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)
# Create the initial random noise
num_channels_latents = (self.dit.config.in_channels // 4) if self.model_name == "pyramid_flux" else self.dit.config.in_channels
latents = self.prepare_latents(
@@ -330,17 +314,9 @@ class PyramidDiTForVideoGeneration:
num_units = temp // self.frame_per_unit
stages = self.stages
# # encode the image latents
# image_transform = transforms.Compose([
# transforms.ToTensor(),
# transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5)),
# ])
#input_image_tensor = image_transform(input_image).unsqueeze(0).unsqueeze(2) # [b c 1 h w]
input_image_latent = input_image_latent.to(dtype).to(device)
generated_latents_list = [input_image_latent] # The generated results
#last_generated_latents = input_image_latent
if not self.sequential_offload_enabled:
self.dit.to(device)
@@ -394,20 +370,19 @@ class PyramidDiTForVideoGeneration:
dtype,
generator,
is_first_frame=False,
callback=callback
)
comfy_pbar.update(1)
if callback is not None:
callback(unit_index, intermed_latents[-1].detach()[0].permute(1,0,2,3), None, temp)
else:
comfy_pbar.update(1)
generated_latents_list.append(intermed_latents[-1])
#last_generated_latents = intermed_latents
generated_latents = torch.cat(generated_latents_list, dim=2)
if output_type == "latent":
image = generated_latents
else:
image = self.decode_latent(generated_latents)
return image
return generated_latents
@torch.no_grad()
def generate(
@@ -425,19 +400,11 @@ class PyramidDiTForVideoGeneration:
alpha: float = 0.5,
num_images_per_prompt: Optional[int] = 1,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
output_type: Optional[str] = "pil",
device: Optional[torch.device] = None,
callback = None
):
assert (temp - 1) % self.frame_per_unit == 0, "The frames should be divided by frame_per unit"
# if isinstance(prompt, str):
# batch_size = 1
# prompt = prompt + ", hyper quality, Ultra HD, 8K" # adding this prompt to improve aesthetics
# else:
# assert isinstance(prompt, list)
# batch_size = len(prompt)
# prompt = [_ + ", hyper quality, Ultra HD, 8K" for _ in prompt]
if isinstance(num_inference_steps, int):
num_inference_steps = [num_inference_steps] * len(self.stages)
elif isinstance(num_inference_steps, list) and len(num_inference_steps) < len(self.stages):
@@ -448,19 +415,11 @@ class PyramidDiTForVideoGeneration:
elif isinstance(video_num_inference_steps, list) and len(video_num_inference_steps) < len(self.stages):
video_num_inference_steps = (video_num_inference_steps * len(self.stages))[:len(self.stages)]
#negative_prompt = negative_prompt or ""
# # Get the text embeddings
# self.text_encoder.to(device)
# prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = self.text_encoder(prompt, device)
# negative_prompt_embeds, negative_prompt_attention_mask, negative_pooled_prompt_embeds = self.text_encoder(negative_prompt, device)
# self.text_encoder.to('cpu')
batch_size = prompt_embeds_dict['prompt_embeds'].shape[0]
if use_linear_guidance:
max_guidance_scale = guidance_scale
# guidance_scale_list = torch.linspace(max_guidance_scale, min_guidance_scale, temp).tolist()
guidance_scale_list = [max(max_guidance_scale - alpha * t_, min_guidance_scale) for t_ in range(temp)]
print(guidance_scale_list)
@@ -512,7 +471,7 @@ class PyramidDiTForVideoGeneration:
stages = self.stages
generated_latents_list = [] # The generated results
#last_generated_latents = None
if not self.sequential_offload_enabled:
self.dit.to(device)
comfy_pbar = ProgressBar(num_units)
@@ -538,6 +497,7 @@ class PyramidDiTForVideoGeneration:
self.dtype,
generator,
is_first_frame=True,
callback=callback
)
else:
# prepare the condition latents
@@ -583,23 +543,19 @@ class PyramidDiTForVideoGeneration:
self.dtype,
generator,
is_first_frame=False,
callback=callback
)
comfy_pbar.update(1)
generated_latents_list.append(intermed_latents[-1])
#last_generated_latents = intermed_latents
if callback is not None:
callback(unit_index, intermed_latents[-1].detach()[0].permute(1,0,2,3), None, temp)
else:
comfy_pbar.update(1)
generated_latents = torch.cat(generated_latents_list, dim=2)
if output_type == "latent":
image = generated_latents
else:
image = self.decode_latent(generated_latents, device)
return image
# @property
# def device(self):
# return next(self.dit.parameters()).device
return generated_latents
@property
def guidance_scale(self):
-25
View File
@@ -1,25 +0,0 @@
from .utils import (
create_optimizer,
get_rank,
get_world_size,
is_main_process,
is_dist_avail_and_initialized,
init_distributed_mode,
setup_for_distributed,
cosine_scheduler,
constant_scheduler,
)
from .sp_utils import (
is_sequence_parallel_initialized,
init_sequence_parallel_group,
get_sequence_parallel_group,
get_sequence_parallel_world_size,
get_sequence_parallel_rank,
get_sequence_parallel_group_rank,
get_sequence_parallel_proc_num,
init_sync_input_group,
get_sync_input_group,
)
from .communicate import all_to_all
-56
View File
@@ -1,56 +0,0 @@
import torch
import torch.distributed as dist
def _all_to_all(
input_: torch.Tensor,
world_size: int,
group: dist.ProcessGroup,
scatter_dim: int,
gather_dim: int,
):
if world_size == 1:
return input_
input_list = [t.contiguous() for t in torch.tensor_split(input_, world_size, scatter_dim)]
output_list = [torch.empty_like(input_list[0]) for _ in range(world_size)]
dist.all_to_all(output_list, input_list, group=group)
return torch.cat(output_list, dim=gather_dim).contiguous()
class _AllToAll(torch.autograd.Function):
@staticmethod
def forward(ctx, input_, process_group, world_size, scatter_dim, gather_dim):
ctx.process_group = process_group
ctx.scatter_dim = scatter_dim
ctx.gather_dim = gather_dim
ctx.world_size = world_size
output = _all_to_all(input_, ctx.world_size, process_group, scatter_dim, gather_dim)
return output
@staticmethod
def backward(ctx, grad_output):
grad_output = _all_to_all(
grad_output,
ctx.world_size,
ctx.process_group,
ctx.gather_dim,
ctx.scatter_dim,
)
return (
grad_output,
None,
None,
None,
None,
)
def all_to_all(
input_: torch.Tensor,
process_group: dist.ProcessGroup,
world_size: int = 1,
scatter_dim: int = 2,
gather_dim: int = 1,
):
return _AllToAll.apply(input_, process_group, world_size, scatter_dim, gather_dim)
-97
View File
@@ -1,97 +0,0 @@
import os
import torch
from .utils import is_dist_avail_and_initialized, get_rank
SEQ_PARALLEL_GROUP = None
SEQ_PARALLEL_SIZE = None
SEQ_PARALLEL_PROC_NUM = None # using how many process for sequence parallel
SYNC_INPUT_GROUP = None
SYNC_INPUT_SIZE = None
def is_sequence_parallel_initialized():
if SEQ_PARALLEL_GROUP is None:
return False
else:
return True
def init_sequence_parallel_group(args):
global SEQ_PARALLEL_GROUP
global SEQ_PARALLEL_SIZE
global SEQ_PARALLEL_PROC_NUM
assert SEQ_PARALLEL_GROUP is None, "sequence parallel group is already initialized"
assert is_dist_avail_and_initialized(), "The pytorch distributed should be initialized"
SEQ_PARALLEL_SIZE = args.sp_group_size
print(f"Setting the Sequence Parallel Size {SEQ_PARALLEL_SIZE}")
rank = torch.distributed.get_rank()
world_size = torch.distributed.get_world_size()
if args.sp_proc_num == -1:
SEQ_PARALLEL_PROC_NUM = world_size
else:
SEQ_PARALLEL_PROC_NUM = args.sp_proc_num
assert SEQ_PARALLEL_PROC_NUM % SEQ_PARALLEL_SIZE == 0, "The process needs to be evenly divided"
for i in range(0, SEQ_PARALLEL_PROC_NUM, SEQ_PARALLEL_SIZE):
ranks = list(range(i, i + SEQ_PARALLEL_SIZE))
group = torch.distributed.new_group(ranks)
if rank in ranks:
SEQ_PARALLEL_GROUP = group
break
def init_sync_input_group(args):
global SYNC_INPUT_GROUP
global SYNC_INPUT_SIZE
assert SYNC_INPUT_GROUP is None, "parallel group is already initialized"
assert is_dist_avail_and_initialized(), "The pytorch distributed should be initialized"
SYNC_INPUT_SIZE = args.max_frames
rank = torch.distributed.get_rank()
world_size = torch.distributed.get_world_size()
for i in range(0, world_size, SYNC_INPUT_SIZE):
ranks = list(range(i, i + SYNC_INPUT_SIZE))
group = torch.distributed.new_group(ranks)
if rank in ranks:
SYNC_INPUT_GROUP = group
break
def get_sequence_parallel_group():
assert SEQ_PARALLEL_GROUP is not None, "sequence parallel group is not initialized"
return SEQ_PARALLEL_GROUP
def get_sync_input_group():
return SYNC_INPUT_GROUP
def get_sequence_parallel_world_size():
assert SEQ_PARALLEL_SIZE is not None, "sequence parallel size is not initialized"
return SEQ_PARALLEL_SIZE
def get_sequence_parallel_rank():
assert SEQ_PARALLEL_SIZE is not None, "sequence parallel size is not initialized"
rank = get_rank()
cp_rank = rank % SEQ_PARALLEL_SIZE
return cp_rank
def get_sequence_parallel_group_rank():
assert SEQ_PARALLEL_SIZE is not None, "sequence parallel size is not initialized"
rank = get_rank()
cp_group_rank = rank // SEQ_PARALLEL_SIZE
return cp_group_rank
def get_sequence_parallel_proc_num():
return SEQ_PARALLEL_PROC_NUM
-377
View File
@@ -1,377 +0,0 @@
import os
import math
import time
import json
from collections import defaultdict, deque
import datetime
import numpy as np
import torch
from torch import optim as optim
import torch.distributed as dist
#from tensorboardX import SummaryWriter
def is_dist_avail_and_initialized():
if not dist.is_available():
return False
if not dist.is_initialized():
return False
return True
def get_world_size():
if not is_dist_avail_and_initialized():
return 1
return dist.get_world_size()
def get_rank():
if not is_dist_avail_and_initialized():
return 0
return dist.get_rank()
def is_main_process():
return get_rank() == 0
def save_on_master(*args, **kwargs):
if is_main_process():
torch.save(*args, **kwargs)
def setup_for_distributed(is_master):
"""
This function disables printing when not in master process
"""
import builtins as __builtin__
builtin_print = __builtin__.print
def print(*args, **kwargs):
force = kwargs.pop('force', False)
if is_master or force:
builtin_print(*args, **kwargs)
__builtin__.print = print
def init_distributed_mode(args):
if int(os.getenv('OMPI_COMM_WORLD_SIZE', '0')) > 0:
rank = int(os.environ['OMPI_COMM_WORLD_RANK'])
local_rank = int(os.environ['OMPI_COMM_WORLD_LOCAL_RANK'])
world_size = int(os.environ['OMPI_COMM_WORLD_SIZE'])
os.environ["LOCAL_RANK"] = os.environ['OMPI_COMM_WORLD_LOCAL_RANK']
os.environ["RANK"] = os.environ['OMPI_COMM_WORLD_RANK']
os.environ["WORLD_SIZE"] = os.environ['OMPI_COMM_WORLD_SIZE']
args.rank = int(os.environ["RANK"])
args.world_size = int(os.environ["WORLD_SIZE"])
args.gpu = int(os.environ["LOCAL_RANK"])
elif 'RANK' in os.environ and 'WORLD_SIZE' in os.environ:
args.rank = int(os.environ["RANK"])
args.world_size = int(os.environ['WORLD_SIZE'])
args.gpu = int(os.environ['LOCAL_RANK'])
else:
print('Not using distributed mode')
args.distributed = False
return
args.distributed = True
args.dist_backend = 'nccl'
args.dist_url = "env://"
print('| distributed init (rank {}): {}, gpu {}'.format(
args.rank, args.dist_url, args.gpu), flush=True)
def cosine_scheduler(base_value, final_value, epochs, niter_per_ep, warmup_epochs=0,
start_warmup_value=0, warmup_steps=-1):
warmup_schedule = np.array([])
warmup_iters = warmup_epochs * niter_per_ep
if warmup_steps > 0:
warmup_iters = warmup_steps
print("Set warmup steps = %d" % warmup_iters)
if warmup_epochs > 0:
warmup_schedule = np.linspace(start_warmup_value, base_value, warmup_iters)
iters = np.arange(epochs * niter_per_ep - warmup_iters)
schedule = np.array(
[final_value + 0.5 * (base_value - final_value) * (1 + math.cos(math.pi * i / (len(iters)))) for i in iters])
schedule = np.concatenate((warmup_schedule, schedule))
assert len(schedule) == epochs * niter_per_ep
return schedule
def constant_scheduler(base_value, epochs, niter_per_ep, warmup_epochs=0,
start_warmup_value=1e-6, warmup_steps=-1):
warmup_schedule = np.array([])
warmup_iters = warmup_epochs * niter_per_ep
if warmup_steps > 0:
warmup_iters = warmup_steps
print("Set warmup steps = %d" % warmup_iters)
if warmup_iters > 0:
warmup_schedule = np.linspace(start_warmup_value, base_value, warmup_iters)
iters = epochs * niter_per_ep - warmup_iters
schedule = np.array([base_value] * iters)
schedule = np.concatenate((warmup_schedule, schedule))
assert len(schedule) == epochs * niter_per_ep
return schedule
def get_parameter_groups(model, weight_decay=1e-5, base_lr=1e-4, skip_list=(), get_num_layer=None, get_layer_scale=None, **kwargs):
parameter_group_names = {}
parameter_group_vars = {}
for name, param in model.named_parameters():
if not param.requires_grad:
continue # frozen weights
if len(kwargs.get('filter_name', [])) > 0:
flag = False
for filter_n in kwargs.get('filter_name', []):
if filter_n in name:
print(f"filter {name} because of the pattern {filter_n}")
flag = True
if flag:
continue
default_scale=1.
if param.ndim <= 1 or name.endswith(".bias") or name in skip_list: # param.ndim <= 1 len(param.shape) == 1
group_name = "no_decay"
this_weight_decay = 0.
else:
group_name = "decay"
this_weight_decay = weight_decay
if get_num_layer is not None:
layer_id = get_num_layer(name)
group_name = "layer_%d_%s" % (layer_id, group_name)
else:
layer_id = None
if group_name not in parameter_group_names:
if get_layer_scale is not None:
scale = get_layer_scale(layer_id)
else:
scale = default_scale
parameter_group_names[group_name] = {
"weight_decay": this_weight_decay,
"params": [],
"lr": base_lr,
"lr_scale": scale,
}
parameter_group_vars[group_name] = {
"weight_decay": this_weight_decay,
"params": [],
"lr": base_lr,
"lr_scale": scale,
}
parameter_group_vars[group_name]["params"].append(param)
parameter_group_names[group_name]["params"].append(name)
print("Param groups = %s" % json.dumps(parameter_group_names, indent=2))
return list(parameter_group_vars.values())
def create_optimizer(args, model, get_num_layer=None, get_layer_scale=None, filter_bias_and_bn=True, skip_list=None, **kwargs):
opt_lower = args.opt.lower()
weight_decay = args.weight_decay
skip = {}
if skip_list is not None:
skip = skip_list
elif hasattr(model, 'no_weight_decay'):
skip = model.no_weight_decay()
print(f"Skip weight decay name marked in model: {skip}")
parameters = get_parameter_groups(model, weight_decay, args.lr, skip, get_num_layer, get_layer_scale, **kwargs)
weight_decay = 0.
if 'fused' in opt_lower:
assert has_apex and torch.cuda.is_available(), 'APEX and CUDA required for fused optimizers'
opt_args = dict(lr=args.lr, weight_decay=weight_decay)
if hasattr(args, 'opt_eps') and args.opt_eps is not None:
opt_args['eps'] = args.opt_eps
if hasattr(args, 'opt_beta1') and args.opt_beta1 is not None:
opt_args['betas'] = (args.opt_beta1, args.opt_beta2)
print('Optimizer config:', opt_args)
opt_split = opt_lower.split('_')
opt_lower = opt_split[-1]
if opt_lower == 'sgd' or opt_lower == 'nesterov':
opt_args.pop('eps', None)
optimizer = optim.SGD(parameters, momentum=args.momentum, nesterov=True, **opt_args)
elif opt_lower == 'momentum':
opt_args.pop('eps', None)
optimizer = optim.SGD(parameters, momentum=args.momentum, nesterov=False, **opt_args)
elif opt_lower == 'adam':
optimizer = optim.Adam(parameters, **opt_args)
elif opt_lower == 'adamw':
optimizer = optim.AdamW(parameters, **opt_args)
elif opt_lower == 'adadelta':
optimizer = optim.Adadelta(parameters, **opt_args)
elif opt_lower == 'rmsprop':
optimizer = optim.RMSprop(parameters, alpha=0.9, momentum=args.momentum, **opt_args)
else:
assert False and "Invalid optimizer"
raise ValueError
return optimizer
class SmoothedValue(object):
"""Track a series of values and provide access to smoothed values over a
window or the global series average.
"""
def __init__(self, window_size=20, fmt=None):
if fmt is None:
fmt = "{median:.4f} ({global_avg:.4f})"
self.deque = deque(maxlen=window_size)
self.total = 0.0
self.count = 0
self.fmt = fmt
def update(self, value, n=1):
self.deque.append(value)
self.count += n
self.total += value * n
def synchronize_between_processes(self):
"""
Warning: does not synchronize the deque!
"""
if not is_dist_avail_and_initialized():
return
t = torch.tensor([self.count, self.total], dtype=torch.float64, device='cuda')
dist.barrier()
dist.all_reduce(t)
t = t.tolist()
self.count = int(t[0])
self.total = t[1]
@property
def median(self):
d = torch.tensor(list(self.deque))
return d.median().item()
@property
def avg(self):
d = torch.tensor(list(self.deque), dtype=torch.float32)
return d.mean().item()
@property
def global_avg(self):
return self.total / self.count
@property
def max(self):
return max(self.deque)
@property
def value(self):
return self.deque[-1]
def __str__(self):
return self.fmt.format(
median=self.median,
avg=self.avg,
global_avg=self.global_avg,
max=self.max,
value=self.value)
class MetricLogger(object):
def __init__(self, delimiter="\t"):
self.meters = defaultdict(SmoothedValue)
self.delimiter = delimiter
def update(self, **kwargs):
for k, v in kwargs.items():
if v is None:
continue
if isinstance(v, torch.Tensor):
v = v.item()
assert isinstance(v, (float, int))
self.meters[k].update(v)
def __getattr__(self, attr):
if attr in self.meters:
return self.meters[attr]
if attr in self.__dict__:
return self.__dict__[attr]
raise AttributeError("'{}' object has no attribute '{}'".format(
type(self).__name__, attr))
def __str__(self):
loss_str = []
for name, meter in self.meters.items():
loss_str.append(
"{}: {}".format(name, str(meter))
)
return self.delimiter.join(loss_str)
def synchronize_between_processes(self):
for meter in self.meters.values():
meter.synchronize_between_processes()
def add_meter(self, name, meter):
self.meters[name] = meter
def log_every(self, iterable, print_freq, header=None):
i = 0
if not header:
header = ''
start_time = time.time()
end = time.time()
iter_time = SmoothedValue(fmt='{avg:.4f}')
data_time = SmoothedValue(fmt='{avg:.4f}')
space_fmt = ':' + str(len(str(len(iterable)))) + 'd'
log_msg = [
header,
'[{0' + space_fmt + '}/{1}]',
'eta: {eta}',
'{meters}',
'time: {time}',
'data: {data}'
]
if torch.cuda.is_available():
log_msg.append('max mem: {memory:.0f}')
log_msg = self.delimiter.join(log_msg)
MB = 1024.0 * 1024.0
for obj in iterable:
data_time.update(time.time() - end)
yield obj
iter_time.update(time.time() - end)
if i % print_freq == 0 or i == len(iterable) - 1:
eta_seconds = iter_time.global_avg * (len(iterable) - i)
eta_string = str(datetime.timedelta(seconds=int(eta_seconds)))
if torch.cuda.is_available():
print(log_msg.format(
i, len(iterable), eta=eta_string,
meters=str(self),
time=str(iter_time), data=str(data_time),
memory=torch.cuda.max_memory_allocated() / MB))
else:
print(log_msg.format(
i, len(iterable), eta=eta_string,
meters=str(self),
time=str(iter_time), data=str(data_time)))
i += 1
end = time.time()
total_time = time.time() - start_time
total_time_str = str(datetime.timedelta(seconds=int(total_time)))
print('{} Total time: {} ({:.4f} s / it)'.format(
header, total_time_str, total_time / len(iterable)))