Add latent preview, cleanup
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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,3 +1,2 @@
|
||||
from .modeling_pyramid_flux import PyramidFluxTransformer
|
||||
from .modeling_text_encoder import FluxTextEncoderWithMask
|
||||
from .modeling_flux_block import FluxSingleTransformerBlock, FluxTransformerBlock
|
||||
@@ -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)
|
||||
|
||||
@@ -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 +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):
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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)))
|
||||
Reference in New Issue
Block a user