diff --git a/examples/pyramidflow_miniflux_example_01.json b/examples/pyramidflow_miniflux_example_01.json new file mode 100644 index 0000000..909e3ff --- /dev/null +++ b/examples/pyramidflow_miniflux_example_01.json @@ -0,0 +1,333 @@ +{ + "last_node_id": 27, + "last_link_id": 38, + "nodes": [ + { + "id": 8, + "type": "PyramidFlowVAEDecode", + "pos": { + "0": 1161, + "1": 873 + }, + "size": { + "0": 315, + "1": 102 + }, + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "PYRAMIDFLOWMODEL", + "link": 8 + }, + { + "name": "samples", + "type": "LATENT", + "link": 9 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 38 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "PyramidFlowVAEDecode" + }, + "widgets_values": [ + 256, + 2 + ] + }, + { + "id": 22, + "type": "PyramidFlowTextEncode", + "pos": { + "0": 567, + "1": 757 + }, + "size": { + "0": 434.50982666015625, + "1": 227.74803161621094 + }, + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "PYRAMIDFLOWMODEL", + "link": 30 + }, + { + "name": "prev_prompt", + "type": "PYRAMIDFLOWPROMPT", + "link": null, + "shape": 7 + } + ], + "outputs": [ + { + "name": "prompt_embeds", + "type": "PYRAMIDFLOWPROMPT", + "links": [ + 31 + ] + } + ], + "properties": { + "Node name for S&R": "PyramidFlowTextEncode" + }, + "widgets_values": [ + "A campfire burning with flames and embers, gradually increasing in size and intensity before dying down towards the end, hyper quality, Ultra HD, 8K", + "cartoon style, worst quality, low quality, blurry, absolute black, absolute white, low res, extra limbs, extra digits, misplaced objects, mutated anatomy, monochrome, horror", + false + ] + }, + { + "id": 5, + "type": "DownloadAndLoadPyramidFlowModel", + "pos": { + "0": 576, + "1": 496 + }, + "size": { + "0": 385.7839050292969, + "1": 202 + }, + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "pyramidflow_model", + "type": "PYRAMIDFLOWMODEL", + "links": [ + 7, + 30 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "DownloadAndLoadPyramidFlowModel" + }, + "widgets_values": [ + "rain1011/pyramid-flow-miniflux", + "diffusion_transformer_384p", + "bf16", + "bf16", + "bf16", + false + ] + }, + { + "id": 9, + "type": "PyramidFlowSampler", + "pos": { + "0": 1059, + "1": 497 + }, + "size": { + "0": 411.5168151855469, + "1": 314 + }, + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "PYRAMIDFLOWMODEL", + "link": 7 + }, + { + "name": "prompt_embeds", + "type": "PYRAMIDFLOWPROMPT", + "link": 31 + }, + { + "name": "input_latent", + "type": "LATENT", + "link": null, + "shape": 7 + } + ], + "outputs": [ + { + "name": "model", + "type": "PYRAMIDFLOWMODEL", + "links": [ + 8 + ] + }, + { + "name": "samples", + "type": "LATENT", + "links": [ + 9 + ], + "slot_index": 1 + } + ], + "properties": { + "Node name for S&R": "PyramidFlowSampler" + }, + "widgets_values": [ + 640, + 384, + "20, 20, 20", + "10, 10, 10", + 16, + 9, + 5, + 44664248661394, + "fixed", + "" + ] + }, + { + "id": 14, + "type": "VHS_VideoCombine", + "pos": { + "0": 1534, + "1": 490 + }, + "size": [ + 1698.6201171875, + 1331.1720703125 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 38 + }, + { + "name": "audio", + "type": "AUDIO", + "link": null, + "shape": 7 + }, + { + "name": "meta_batch", + "type": "VHS_BatchManager", + "link": null, + "shape": 7 + }, + { + "name": "vae", + "type": "VAE", + "link": null, + "shape": 7 + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 24, + "loop_count": 0, + "filename_prefix": "PyramidFlow", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 19, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "PyramidFlow_00061.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 24 + }, + "muted": false + } + } + } + ], + "links": [ + [ + 7, + 5, + 0, + 9, + 0, + "PYRAMIDFLOWMODEL" + ], + [ + 8, + 9, + 0, + 8, + 0, + "PYRAMIDFLOWMODEL" + ], + [ + 9, + 9, + 1, + 8, + 1, + "LATENT" + ], + [ + 30, + 5, + 0, + 22, + 0, + "PYRAMIDFLOWMODEL" + ], + [ + 31, + 22, + 0, + 9, + 1, + "PYRAMIDFLOWPROMPT" + ], + [ + 38, + 8, + 0, + 14, + 0, + "IMAGE" + ] + ], + "groups": [], + "config": {}, + "extra": { + "ds": { + "scale": 0.6934334949442617, + "offset": [ + -267.34972182737584, + -351.34162515690946 + ] + } + }, + "version": 0.4 +} \ No newline at end of file diff --git a/nodes.py b/nodes.py index cfdd0e6..654ff94 100644 --- a/nodes.py +++ b/nodes.py @@ -25,6 +25,8 @@ class DownloadAndLoadPyramidFlowModel: "model": ( [ "rain1011/pyramid-flow-sd3", + "rain1011/pyramid-flow-miniflux" + ], ), "variant": ( @@ -79,10 +81,12 @@ class DownloadAndLoadPyramidFlowModel: local_dir=model_path, local_dir_use_symlinks=False, ) - + model_name = "pyramid_flux" if "flux" in model else "pyramid_mmdit" + print(model_name) model = PyramidDiTForVideoGeneration( model_path, model_dtype, + model_name, text_encoder_dtype, vae_dtype, model_variant=variant, diff --git a/pyramid_dit/__init__.py b/pyramid_dit/__init__.py index dd0d5ff..266c6a1 100644 --- a/pyramid_dit/__init__.py +++ b/pyramid_dit/__init__.py @@ -1,3 +1,3 @@ -from .modeling_pyramid_mmdit import PyramidDiffusionMMDiT from .pyramid_dit_for_video_gen_pipeline import PyramidDiTForVideoGeneration -from .modeling_text_encoder import SD3TextEncoderWithMask \ No newline at end of file +#from .flux_modules import FluxSingleTransformerBlock, FluxTransformerBlock, FluxTextEncoderWithMask +#from .mmdit_modules import JointTransformerBlock, SD3TextEncoderWithMask \ No newline at end of file diff --git a/pyramid_dit/flux_modules/__init__.py b/pyramid_dit/flux_modules/__init__.py new file mode 100644 index 0000000..9fedf56 --- /dev/null +++ b/pyramid_dit/flux_modules/__init__.py @@ -0,0 +1,3 @@ +from .modeling_pyramid_flux import PyramidFluxTransformer +from .modeling_text_encoder import FluxTextEncoderWithMask +from .modeling_flux_block import FluxSingleTransformerBlock, FluxTransformerBlock \ No newline at end of file diff --git a/pyramid_dit/flux_modules/modeling_embedding.py b/pyramid_dit/flux_modules/modeling_embedding.py new file mode 100644 index 0000000..11b5b03 --- /dev/null +++ b/pyramid_dit/flux_modules/modeling_embedding.py @@ -0,0 +1,201 @@ +import math +from typing import List, Optional, Tuple, Union + +import numpy as np +import torch +import torch.nn.functional as F +from torch import nn + +from diffusers.models.activations import get_activation, FP32SiLU + +def get_timestep_embedding( + timesteps: torch.Tensor, + embedding_dim: int, + flip_sin_to_cos: bool = False, + downscale_freq_shift: float = 1, + scale: float = 1, + max_period: int = 10000, +): + """ + This matches the implementation in Denoising Diffusion Probabilistic Models: Create sinusoidal timestep embeddings. + + Args + timesteps (torch.Tensor): + a 1-D Tensor of N indices, one per batch element. These may be fractional. + embedding_dim (int): + the dimension of the output. + flip_sin_to_cos (bool): + Whether the embedding order should be `cos, sin` (if True) or `sin, cos` (if False) + downscale_freq_shift (float): + Controls the delta between frequencies between dimensions + scale (float): + Scaling factor applied to the embeddings. + max_period (int): + Controls the maximum frequency of the embeddings + Returns + torch.Tensor: an [N x dim] Tensor of positional embeddings. + """ + assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array" + + half_dim = embedding_dim // 2 + exponent = -math.log(max_period) * torch.arange( + start=0, end=half_dim, dtype=torch.float32, device=timesteps.device + ) + exponent = exponent / (half_dim - downscale_freq_shift) + + emb = torch.exp(exponent) + emb = timesteps[:, None].float() * emb[None, :] + + # scale embeddings + emb = scale * emb + + # concat sine and cosine embeddings + emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1) + + # flip sine and cosine embeddings + if flip_sin_to_cos: + emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1) + + # zero pad + if embedding_dim % 2 == 1: + emb = torch.nn.functional.pad(emb, (0, 1, 0, 0)) + return emb + + +class Timesteps(nn.Module): + def __init__(self, num_channels: int, flip_sin_to_cos: bool, downscale_freq_shift: float, scale: int = 1): + super().__init__() + self.num_channels = num_channels + self.flip_sin_to_cos = flip_sin_to_cos + self.downscale_freq_shift = downscale_freq_shift + self.scale = scale + + def forward(self, timesteps): + t_emb = get_timestep_embedding( + timesteps, + self.num_channels, + flip_sin_to_cos=self.flip_sin_to_cos, + downscale_freq_shift=self.downscale_freq_shift, + scale=self.scale, + ) + return t_emb + + +class TimestepEmbedding(nn.Module): + def __init__( + self, + in_channels: int, + time_embed_dim: int, + act_fn: str = "silu", + out_dim: int = None, + post_act_fn: Optional[str] = None, + cond_proj_dim=None, + sample_proj_bias=True, + ): + super().__init__() + + self.linear_1 = nn.Linear(in_channels, time_embed_dim, sample_proj_bias) + + if cond_proj_dim is not None: + self.cond_proj = nn.Linear(cond_proj_dim, in_channels, bias=False) + else: + self.cond_proj = None + + self.act = get_activation(act_fn) + + if out_dim is not None: + time_embed_dim_out = out_dim + else: + time_embed_dim_out = time_embed_dim + self.linear_2 = nn.Linear(time_embed_dim, time_embed_dim_out, sample_proj_bias) + + if post_act_fn is None: + self.post_act = None + else: + self.post_act = get_activation(post_act_fn) + + def forward(self, sample, condition=None): + if condition is not None: + sample = sample + self.cond_proj(condition) + sample = self.linear_1(sample) + + if self.act is not None: + sample = self.act(sample) + + sample = self.linear_2(sample) + + if self.post_act is not None: + sample = self.post_act(sample) + return sample + + +class PixArtAlphaTextProjection(nn.Module): + """ + Projects caption embeddings. Also handles dropout for classifier-free guidance. + + Adapted from https://github.com/PixArt-alpha/PixArt-alpha/blob/master/diffusion/model/nets/PixArt_blocks.py + """ + + def __init__(self, in_features, hidden_size, out_features=None, act_fn="gelu_tanh"): + super().__init__() + if out_features is None: + out_features = hidden_size + self.linear_1 = nn.Linear(in_features=in_features, out_features=hidden_size, bias=True) + if act_fn == "gelu_tanh": + self.act_1 = nn.GELU(approximate="tanh") + elif act_fn == "silu": + self.act_1 = nn.SiLU() + elif act_fn == "silu_fp32": + self.act_1 = FP32SiLU() + else: + raise ValueError(f"Unknown activation function: {act_fn}") + self.linear_2 = nn.Linear(in_features=hidden_size, out_features=out_features, bias=True) + + def forward(self, caption): + hidden_states = self.linear_1(caption) + hidden_states = self.act_1(hidden_states) + hidden_states = self.linear_2(hidden_states) + return hidden_states + + +class CombinedTimestepGuidanceTextProjEmbeddings(nn.Module): + def __init__(self, embedding_dim, pooled_projection_dim): + super().__init__() + + self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) + self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) + self.guidance_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) + self.text_embedder = PixArtAlphaTextProjection(pooled_projection_dim, embedding_dim, act_fn="silu") + + def forward(self, timestep, guidance, pooled_projection): + timesteps_proj = self.time_proj(timestep) + timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=pooled_projection.dtype)) # (N, D) + + guidance_proj = self.time_proj(guidance) + guidance_emb = self.guidance_embedder(guidance_proj.to(dtype=pooled_projection.dtype)) # (N, D) + + time_guidance_emb = timesteps_emb + guidance_emb + + pooled_projections = self.text_embedder(pooled_projection) + conditioning = time_guidance_emb + pooled_projections + + return conditioning + + +class CombinedTimestepTextProjEmbeddings(nn.Module): + def __init__(self, embedding_dim, pooled_projection_dim): + super().__init__() + + self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0) + self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim) + self.text_embedder = PixArtAlphaTextProjection(pooled_projection_dim, embedding_dim, act_fn="silu") + + def forward(self, timestep, pooled_projection): + timesteps_proj = self.time_proj(timestep) + timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=pooled_projection.dtype)) # (N, D) + + pooled_projections = self.text_embedder(pooled_projection) + + conditioning = timesteps_emb + pooled_projections + + return conditioning \ No newline at end of file diff --git a/pyramid_dit/flux_modules/modeling_flux_block.py b/pyramid_dit/flux_modules/modeling_flux_block.py new file mode 100644 index 0000000..379975f --- /dev/null +++ b/pyramid_dit/flux_modules/modeling_flux_block.py @@ -0,0 +1,1044 @@ +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 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) + 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) + + +class FeedForward(nn.Module): + r""" + A feed-forward layer. + + Parameters: + dim (`int`): The number of channels in the input. + dim_out (`int`, *optional*): The number of channels in the output. If not given, defaults to `dim`. + mult (`int`, *optional*, defaults to 4): The multiplier to use for the hidden dimension. + dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. + activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward. + final_dropout (`bool` *optional*, defaults to False): Apply a final dropout. + bias (`bool`, defaults to True): Whether to use a bias in the linear layer. + """ + + def __init__( + self, + dim: int, + dim_out: Optional[int] = None, + mult: int = 4, + dropout: float = 0.0, + activation_fn: str = "geglu", + final_dropout: bool = False, + inner_dim=None, + bias: bool = True, + ): + super().__init__() + if inner_dim is None: + inner_dim = int(dim * mult) + dim_out = dim_out if dim_out is not None else dim + + if activation_fn == "gelu": + act_fn = GELU(dim, inner_dim, bias=bias) + if activation_fn == "gelu-approximate": + act_fn = GELU(dim, inner_dim, approximate="tanh", bias=bias) + elif activation_fn == "geglu": + act_fn = GEGLU(dim, inner_dim, bias=bias) + elif activation_fn == "geglu-approximate": + act_fn = ApproximateGELU(dim, inner_dim, bias=bias) + elif activation_fn == "swiglu": + act_fn = SwiGLU(dim, inner_dim, bias=bias) + + self.net = nn.ModuleList([]) + # project in + self.net.append(act_fn) + # project dropout + self.net.append(nn.Dropout(dropout)) + # project out + self.net.append(nn.Linear(inner_dim, dim_out, bias=bias)) + # FF as used in Vision Transformer, MLP-Mixer, etc. have a final dropout + if final_dropout: + self.net.append(nn.Dropout(dropout)) + + def forward(self, hidden_states: torch.Tensor, *args, **kwargs) -> torch.Tensor: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + for module in self.net: + hidden_states = module(hidden_states) + 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): + 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] + output_hidden = torch.zeros_like(query) + output_encoder_hidden = torch.zeros_like(encoder_query) + encoder_length = encoder_query.shape[1] + + 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] + + i_sum = 0 + 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] + 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]) + + 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) + stage_encoder_hidden_output = stage_output[:, :encoder_length] + stage_hidden_output = stage_output[:, encoder_length:] + 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_hidden = output_hidden.flatten(2, 3) + output_encoder_hidden = output_encoder_hidden.flatten(2, 3) + + 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): + 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" + + encoder_length = encoder_query.shape[1] + 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] + + 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] + 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) + + # with torch.backends.cuda.sdp_kernel(enable_math=False, enable_flash=False, enable_mem_efficient=True): + 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).flatten(2, 3) # [bs, tot_seq, dim] + + output_encoder_hidden_list.append(stage_hidden_states[:, :encoder_length]) + output_hidden_list.append(stage_hidden_states[:, encoder_length:]) + i_sum += length + + output_encoder_hidden = torch.stack(output_encoder_hidden_list, dim=1) # [b n s d] + output_encoder_hidden = rearrange(output_encoder_hidden, 'b n s d -> (b n) s d') + output_hidden = torch.cat(output_hidden_list, dim=1) + + 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): + 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] + output_hidden = torch.zeros_like(query) + + qkv_list = [] + num_stages = len(hidden_length) + qkv = torch.stack([query, key, value], dim=2) # [bs, sub_seq, 3, head, head_dim] + + i_sum = 0 + for i_p, length in enumerate(hidden_length): + qkv_tokens = qkv[:, i_sum:i_sum+length] + + 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) + output_hidden[:, i_sum:i_sum+length] = stage_output + token_sum += tot_token_num + i_sum += length + + output_hidden = output_hidden.flatten(2, 3) + + 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): + 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] + + i_sum = 0 + output_hidden_list = [] + + for i_p, length in enumerate(hidden_length): + qkv_tokens = qkv[:, i_sum:i_sum+length] + + 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) + 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).flatten(2, 3) # [bs, tot_seq, dim] + + output_hidden_list.append(stage_hidden_states) + i_sum += length + + output_hidden = torch.cat(output_hidden_list, dim=1) + + return output_hidden + + +class Attention(nn.Module): + + def __init__( + self, + query_dim: int, + cross_attention_dim: Optional[int] = None, + heads: int = 8, + dim_head: int = 64, + dropout: float = 0.0, + bias: bool = False, + qk_norm: Optional[str] = None, + added_kv_proj_dim: Optional[int] = None, + added_proj_bias: Optional[bool] = True, + out_bias: bool = True, + only_cross_attention: bool = False, + eps: float = 1e-5, + processor: Optional["AttnProcessor"] = None, + out_dim: int = None, + context_pre_only=None, + pre_only=False, + ): + super().__init__() + + self.inner_dim = out_dim if out_dim is not None else dim_head * heads + self.inner_kv_dim = self.inner_dim + self.query_dim = query_dim + self.use_bias = bias + self.cross_attention_dim = cross_attention_dim if cross_attention_dim is not None else query_dim + + self.dropout = dropout + self.out_dim = out_dim if out_dim is not None else query_dim + self.context_pre_only = context_pre_only + self.pre_only = pre_only + + self.scale = dim_head**-0.5 + self.heads = out_dim // dim_head if out_dim is not None else heads + + + self.added_kv_proj_dim = added_kv_proj_dim + self.only_cross_attention = only_cross_attention + + if self.added_kv_proj_dim is None and self.only_cross_attention: + raise ValueError( + "`only_cross_attention` can only be set to True if `added_kv_proj_dim` is not None. Make sure to set either `only_cross_attention=False` or define `added_kv_proj_dim`." + ) + + if qk_norm is None: + self.norm_q = None + self.norm_k = None + elif qk_norm == "rms_norm": + self.norm_q = RMSNorm(dim_head, eps=eps) + self.norm_k = RMSNorm(dim_head, eps=eps) + else: + raise ValueError(f"unknown qk_norm: {qk_norm}. Should be None or 'layer_norm'") + + self.to_q = nn.Linear(query_dim, self.inner_dim, bias=bias) + + if not self.only_cross_attention: + # only relevant for the `AddedKVProcessor` classes + self.to_k = nn.Linear(self.cross_attention_dim, self.inner_kv_dim, bias=bias) + self.to_v = nn.Linear(self.cross_attention_dim, self.inner_kv_dim, bias=bias) + else: + self.to_k = None + self.to_v = None + + self.added_proj_bias = added_proj_bias + if self.added_kv_proj_dim is not None: + self.add_k_proj = nn.Linear(added_kv_proj_dim, self.inner_kv_dim, bias=added_proj_bias) + self.add_v_proj = nn.Linear(added_kv_proj_dim, self.inner_kv_dim, bias=added_proj_bias) + if self.context_pre_only is not None: + self.add_q_proj = nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias) + + if not self.pre_only: + self.to_out = nn.ModuleList([]) + self.to_out.append(nn.Linear(self.inner_dim, self.out_dim, bias=out_bias)) + self.to_out.append(nn.Dropout(dropout)) + + if self.context_pre_only is not None and not self.context_pre_only: + self.to_add_out = nn.Linear(self.inner_dim, self.out_dim, bias=out_bias) + + if qk_norm is not None and added_kv_proj_dim is not None: + if qk_norm == "fp32_layer_norm": + self.norm_added_q = FP32LayerNorm(dim_head, elementwise_affine=False, bias=False, eps=eps) + self.norm_added_k = FP32LayerNorm(dim_head, elementwise_affine=False, bias=False, eps=eps) + elif qk_norm == "rms_norm": + self.norm_added_q = RMSNorm(dim_head, eps=eps) + self.norm_added_k = RMSNorm(dim_head, eps=eps) + else: + self.norm_added_q = None + self.norm_added_k = None + + # set attention processor + self.set_processor(processor) + + def set_processor(self, processor: "AttnProcessor") -> None: + self.processor = processor + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: Optional[torch.Tensor] = None, + encoder_attention_mask: Optional[torch.Tensor] = None, + attention_mask: Optional[torch.Tensor] = None, + hidden_length: List = None, + image_rotary_emb: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + + return self.processor( + self, + hidden_states, + encoder_hidden_states=encoder_hidden_states, + encoder_attention_mask=encoder_attention_mask, + attention_mask=attention_mask, + hidden_length=hidden_length, + image_rotary_emb=image_rotary_emb, + ) + + +class FluxSingleAttnProcessor2_0: + r""" + Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). + """ + def __init__(self, use_flash_attn=False): + 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() + else: + if is_sequence_parallel_initialized(): + self.varlen_attn = SequenceParallelVarlenAttnSingle() + else: + self.varlen_attn = VarlenSelfAttnSingle() + + def __call__( + self, + attn: Attention, + hidden_states: torch.Tensor, + encoder_hidden_states: Optional[torch.Tensor] = None, + encoder_attention_mask: Optional[torch.Tensor] = None, + attention_mask: Optional[torch.FloatTensor] = None, + hidden_length: List = None, + image_rotary_emb: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + + query = attn.to_q(hidden_states) + key = attn.to_k(hidden_states) + value = attn.to_v(hidden_states) + + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + + query = query.view(query.shape[0], -1, attn.heads, head_dim) + key = key.view(key.shape[0], -1, attn.heads, head_dim) + value = value.view(value.shape[0], -1, attn.heads, head_dim) + + if attn.norm_q is not None: + query = attn.norm_q(query) + if attn.norm_k is not None: + key = attn.norm_k(key) + + if self.use_flash_attn: + hidden_states = self.varlen_flash_attn( + query, key, value, + attn.heads, attn.scale, hidden_length, + image_rotary_emb, encoder_attention_mask, + ) + else: + hidden_states = self.varlen_attn( + query, key, value, + attn.heads, attn.scale, hidden_length, + image_rotary_emb, attention_mask, + ) + + return hidden_states + + +class FluxAttnProcessor2_0: + """Attention processor used typically in processing the SD3-like self-attention projections.""" + + def __init__(self, use_flash_attn=False): + 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() + else: + if is_sequence_parallel_initialized(): + self.varlen_attn = SequenceParallelVarlenSelfAttentionWithT5Mask() + else: + self.varlen_attn = VarlenSelfAttentionWithT5Mask() + + def __call__( + self, + attn: Attention, + hidden_states: torch.FloatTensor, + encoder_hidden_states: torch.FloatTensor = None, + encoder_attention_mask: Optional[torch.Tensor] = None, + attention_mask: Optional[torch.FloatTensor] = None, + hidden_length: List = None, + 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) + + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + + query = query.view(query.shape[0], -1, attn.heads, head_dim) + key = key.view(key.shape[0], -1, attn.heads, head_dim) + value = value.view(value.shape[0], -1, attn.heads, head_dim) + + if attn.norm_q is not None: + query = attn.norm_q(query) + if attn.norm_k is not None: + key = attn.norm_k(key) + + # `context` projections. + encoder_hidden_states_query_proj = attn.add_q_proj(encoder_hidden_states) + encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states) + encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states) + + encoder_hidden_states_query_proj = encoder_hidden_states_query_proj.view( + encoder_hidden_states_query_proj.shape[0], -1, attn.heads, head_dim + ) + encoder_hidden_states_key_proj = encoder_hidden_states_key_proj.view( + encoder_hidden_states_key_proj.shape[0], -1, attn.heads, head_dim + ) + encoder_hidden_states_value_proj = encoder_hidden_states_value_proj.view( + encoder_hidden_states_value_proj.shape[0], -1, attn.heads, head_dim + ) + + if attn.norm_added_q is not None: + encoder_hidden_states_query_proj = attn.norm_added_q(encoder_hidden_states_query_proj) + if attn.norm_added_k is not None: + encoder_hidden_states_key_proj = attn.norm_added_k(encoder_hidden_states_key_proj) + + if self.use_flash_attn: + hidden_states, encoder_hidden_states = self.varlen_flash_attn( + query, key, value, + encoder_hidden_states_query_proj, encoder_hidden_states_key_proj, + encoder_hidden_states_value_proj, attn.heads, attn.scale, hidden_length, + image_rotary_emb, encoder_attention_mask, + ) + else: + hidden_states, encoder_hidden_states = self.varlen_attn( + query, key, value, + encoder_hidden_states_query_proj, encoder_hidden_states_key_proj, + encoder_hidden_states_value_proj, attn.heads, attn.scale, hidden_length, + image_rotary_emb, attention_mask, + ) + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + encoder_hidden_states = attn.to_add_out(encoder_hidden_states) + + return hidden_states, encoder_hidden_states + + +class FluxSingleTransformerBlock(nn.Module): + r""" + A Transformer block following the MMDiT architecture, introduced in Stable Diffusion 3. + + Reference: https://arxiv.org/abs/2403.03206 + + Parameters: + dim (`int`): The number of channels in the input and output. + num_attention_heads (`int`): The number of heads to use for multi-head attention. + attention_head_dim (`int`): The number of channels in each head. + context_pre_only (`bool`): Boolean to determine if we should add some blocks associated with the + processing of `context` conditions. + """ + + def __init__(self, dim, num_attention_heads, attention_head_dim, mlp_ratio=4.0, use_flash_attn=False): + super().__init__() + self.mlp_hidden_dim = int(dim * mlp_ratio) + + self.norm = AdaLayerNormZeroSingle(dim) + self.proj_mlp = nn.Linear(dim, self.mlp_hidden_dim) + self.act_mlp = nn.GELU(approximate="tanh") + self.proj_out = nn.Linear(dim + self.mlp_hidden_dim, dim) + + processor = FluxSingleAttnProcessor2_0(use_flash_attn) + self.attn = Attention( + query_dim=dim, + cross_attention_dim=None, + dim_head=attention_head_dim, + heads=num_attention_heads, + out_dim=dim, + bias=True, + processor=processor, + qk_norm="rms_norm", + eps=1e-6, + pre_only=True, + ) + + def forward( + self, + hidden_states: torch.FloatTensor, + temb: torch.FloatTensor, + encoder_attention_mask=None, + attention_mask=None, + hidden_length=None, + image_rotary_emb=None, + ): + residual = hidden_states + norm_hidden_states, gate = self.norm(hidden_states, emb=temb, hidden_length=hidden_length) + mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states)) + + attn_output = self.attn( + hidden_states=norm_hidden_states, + encoder_hidden_states=None, + encoder_attention_mask=encoder_attention_mask, + attention_mask=attention_mask, + hidden_length=hidden_length, + image_rotary_emb=image_rotary_emb, + ) + + hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) + hidden_states = gate * self.proj_out(hidden_states) + hidden_states = residual + hidden_states + if hidden_states.dtype == torch.float16: + hidden_states = hidden_states.clip(-65504, 65504) + + return hidden_states + + +class FluxTransformerBlock(nn.Module): + r""" + A Transformer block following the MMDiT architecture, introduced in Stable Diffusion 3. + + Reference: https://arxiv.org/abs/2403.03206 + + Parameters: + dim (`int`): The number of channels in the input and output. + num_attention_heads (`int`): The number of heads to use for multi-head attention. + attention_head_dim (`int`): The number of channels in each head. + context_pre_only (`bool`): Boolean to determine if we should add some blocks associated with the + processing of `context` conditions. + """ + + def __init__(self, dim, num_attention_heads, attention_head_dim, qk_norm="rms_norm", eps=1e-6, use_flash_attn=False): + super().__init__() + + self.norm1 = AdaLayerNormZero(dim) + + self.norm1_context = AdaLayerNormZero(dim) + + if hasattr(F, "scaled_dot_product_attention"): + processor = FluxAttnProcessor2_0(use_flash_attn) + else: + raise ValueError( + "The current PyTorch version does not support the `scaled_dot_product_attention` function." + ) + self.attn = Attention( + query_dim=dim, + cross_attention_dim=None, + added_kv_proj_dim=dim, + dim_head=attention_head_dim, + heads=num_attention_heads, + out_dim=dim, + context_pre_only=False, + bias=True, + processor=processor, + qk_norm=qk_norm, + eps=eps, + ) + + self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) + self.ff = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") + + self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) + self.ff_context = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") + + def forward( + self, + hidden_states: torch.FloatTensor, + encoder_hidden_states: torch.FloatTensor, + encoder_attention_mask: torch.FloatTensor, + temb: torch.FloatTensor, + attention_mask: torch.FloatTensor = None, + hidden_length: List = None, + image_rotary_emb=None, + ): + norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb, hidden_length=hidden_length) + + norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( + encoder_hidden_states, emb=temb + ) + + # Attention. + attn_output, context_attn_output = self.attn( + hidden_states=norm_hidden_states, + encoder_hidden_states=norm_encoder_hidden_states, + encoder_attention_mask=encoder_attention_mask, + attention_mask=attention_mask, + hidden_length=hidden_length, + image_rotary_emb=image_rotary_emb, + ) + + # Process attention outputs for the `hidden_states`. + attn_output = gate_msa * attn_output + hidden_states = hidden_states + attn_output + + norm_hidden_states = self.norm2(hidden_states) + norm_hidden_states = norm_hidden_states * (1 + scale_mlp) + shift_mlp + + ff_output = self.ff(norm_hidden_states) + ff_output = gate_mlp * ff_output + + hidden_states = hidden_states + ff_output + + # Process attention outputs for the `encoder_hidden_states`. + + context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output + encoder_hidden_states = encoder_hidden_states + context_attn_output + + norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) + norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] + + context_ff_output = self.ff_context(norm_encoder_hidden_states) + encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output + + if encoder_hidden_states.dtype == torch.float16: + encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504) + + return encoder_hidden_states, hidden_states diff --git a/pyramid_dit/flux_modules/modeling_normalization.py b/pyramid_dit/flux_modules/modeling_normalization.py new file mode 100644 index 0000000..76eb409 --- /dev/null +++ b/pyramid_dit/flux_modules/modeling_normalization.py @@ -0,0 +1,249 @@ +import numbers +from typing import Dict, Optional, Tuple + +import torch +import torch.nn as nn +import torch.nn.functional as F +from einops import rearrange +from diffusers.utils import is_torch_version + + +if is_torch_version(">=", "2.1.0"): + LayerNorm = nn.LayerNorm +else: + # Has optional bias parameter compared to torch layer norm + # TODO: replace with torch layernorm once min required torch version >= 2.1 + class LayerNorm(nn.Module): + def __init__(self, dim, eps: float = 1e-5, elementwise_affine: bool = True, bias: bool = True): + super().__init__() + + self.eps = eps + + if isinstance(dim, numbers.Integral): + dim = (dim,) + + self.dim = torch.Size(dim) + + if elementwise_affine: + self.weight = nn.Parameter(torch.ones(dim)) + self.bias = nn.Parameter(torch.zeros(dim)) if bias else None + else: + self.weight = None + self.bias = None + + def forward(self, input): + return F.layer_norm(input, self.dim, self.weight, self.bias, self.eps) + + +class FP32LayerNorm(nn.LayerNorm): + def forward(self, inputs: torch.Tensor) -> torch.Tensor: + origin_dtype = inputs.dtype + return F.layer_norm( + inputs.float(), + self.normalized_shape, + self.weight.float() if self.weight is not None else None, + self.bias.float() if self.bias is not None else None, + self.eps, + ).to(origin_dtype) + + +class RMSNorm(nn.Module): + def __init__(self, dim, eps: float, elementwise_affine: bool = True): + super().__init__() + + self.eps = eps + + if isinstance(dim, numbers.Integral): + dim = (dim,) + + self.dim = torch.Size(dim) + + if elementwise_affine: + self.weight = nn.Parameter(torch.ones(dim)) + else: + self.weight = None + + def forward(self, hidden_states): + input_dtype = hidden_states.dtype + variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True) + hidden_states = hidden_states * torch.rsqrt(variance + self.eps) + + if self.weight is not None: + # convert into half-precision if necessary + if self.weight.dtype in [torch.float16, torch.bfloat16]: + hidden_states = hidden_states.to(self.weight.dtype) + hidden_states = hidden_states * self.weight + else: + hidden_states = hidden_states.to(input_dtype) + + return hidden_states + + +class AdaLayerNormContinuous(nn.Module): + def __init__( + self, + embedding_dim: int, + conditioning_embedding_dim: int, + # NOTE: It is a bit weird that the norm layer can be configured to have scale and shift parameters + # because the output is immediately scaled and shifted by the projected conditioning embeddings. + # Note that AdaLayerNorm does not let the norm layer have scale and shift parameters. + # However, this is how it was implemented in the original code, and it's rather likely you should + # set `elementwise_affine` to False. + elementwise_affine=True, + eps=1e-5, + bias=True, + norm_type="layer_norm", + ): + super().__init__() + self.silu = nn.SiLU() + self.linear = nn.Linear(conditioning_embedding_dim, embedding_dim * 2, bias=bias) + if norm_type == "layer_norm": + self.norm = LayerNorm(embedding_dim, eps, elementwise_affine, bias) + elif norm_type == "rms_norm": + self.norm = RMSNorm(embedding_dim, eps, elementwise_affine) + else: + raise ValueError(f"unknown norm_type {norm_type}") + + def forward_with_pad(self, x: torch.Tensor, conditioning_embedding: torch.Tensor, hidden_length=None) -> torch.Tensor: + assert hidden_length is not None + + emb = self.linear(self.silu(conditioning_embedding).to(x.dtype)) + batch_emb = torch.zeros_like(x).repeat(1, 1, 2) + + i_sum = 0 + num_stages = len(hidden_length) + for i_p, length in enumerate(hidden_length): + batch_emb[:, i_sum:i_sum+length] = emb[i_p::num_stages][:,None] + i_sum += length + + batch_scale, batch_shift = torch.chunk(batch_emb, 2, dim=2) + x = self.norm(x) * (1 + batch_scale) + batch_shift + return x + + def forward(self, x: torch.Tensor, conditioning_embedding: torch.Tensor, hidden_length=None) -> torch.Tensor: + # convert back to the original dtype in case `conditioning_embedding`` is upcasted to float32 (needed for hunyuanDiT) + if hidden_length is not None: + return self.forward_with_pad(x, conditioning_embedding, hidden_length) + emb = self.linear(self.silu(conditioning_embedding).to(x.dtype)) + scale, shift = torch.chunk(emb, 2, dim=1) + x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :] + return x + + +class AdaLayerNormZero(nn.Module): + r""" + Norm layer adaptive layer norm zero (adaLN-Zero). + + Parameters: + embedding_dim (`int`): The size of each embedding vector. + num_embeddings (`int`): The size of the embeddings dictionary. + """ + + def __init__(self, embedding_dim: int, num_embeddings: Optional[int] = None): + super().__init__() + self.emb = None + + self.silu = nn.SiLU() + self.linear = nn.Linear(embedding_dim, 6 * embedding_dim, bias=True) + self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=False, eps=1e-6) + + def forward_with_pad( + self, + x: torch.Tensor, + timestep: Optional[torch.Tensor] = None, + class_labels: Optional[torch.LongTensor] = None, + hidden_dtype: Optional[torch.dtype] = None, + emb: Optional[torch.Tensor] = None, + hidden_length: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + # hidden_length: [[20, 30], [30, 40], [50, 60]] + # x: [bs, seq_len, dim] + if self.emb is not None: + emb = self.emb(timestep, class_labels, hidden_dtype=hidden_dtype) + + emb = self.linear(self.silu(emb)) + batch_emb = torch.zeros_like(x).repeat(1, 1, 6) + + i_sum = 0 + num_stages = len(hidden_length) + for i_p, length in enumerate(hidden_length): + batch_emb[:, i_sum:i_sum+length] = emb[i_p::num_stages][:,None] + i_sum += length + + batch_shift_msa, batch_scale_msa, batch_gate_msa, batch_shift_mlp, batch_scale_mlp, batch_gate_mlp = batch_emb.chunk(6, dim=2) + x = self.norm(x) * (1 + batch_scale_msa) + batch_shift_msa + return x, batch_gate_msa, batch_shift_mlp, batch_scale_mlp, batch_gate_mlp + + def forward( + self, + x: torch.Tensor, + timestep: Optional[torch.Tensor] = None, + class_labels: Optional[torch.LongTensor] = None, + hidden_dtype: Optional[torch.dtype] = None, + emb: Optional[torch.Tensor] = None, + hidden_length: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + if hidden_length is not None: + return self.forward_with_pad(x, timestep, class_labels, hidden_dtype, emb, hidden_length) + if self.emb is not None: + emb = self.emb(timestep, class_labels, hidden_dtype=hidden_dtype) + emb = self.linear(self.silu(emb)) + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = emb.chunk(6, dim=1) + x = self.norm(x) * (1 + scale_msa[:, None]) + shift_msa[:, None] + return x, gate_msa, shift_mlp, scale_mlp, gate_mlp + + +class AdaLayerNormZeroSingle(nn.Module): + r""" + Norm layer adaptive layer norm zero (adaLN-Zero). + + Parameters: + embedding_dim (`int`): The size of each embedding vector. + num_embeddings (`int`): The size of the embeddings dictionary. + """ + + def __init__(self, embedding_dim: int, norm_type="layer_norm", bias=True): + super().__init__() + + self.silu = nn.SiLU() + self.linear = nn.Linear(embedding_dim, 3 * embedding_dim, bias=bias) + if norm_type == "layer_norm": + self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=False, eps=1e-6) + else: + raise ValueError( + f"Unsupported `norm_type` ({norm_type}) provided. Supported ones are: 'layer_norm', 'fp32_layer_norm'." + ) + + def forward_with_pad( + self, + x: torch.Tensor, + emb: Optional[torch.Tensor] = None, + hidden_length: Optional[torch.Tensor] = None, + ): + emb = self.linear(self.silu(emb)) + batch_emb = torch.zeros_like(x).repeat(1, 1, 3) + + i_sum = 0 + num_stages = len(hidden_length) + for i_p, length in enumerate(hidden_length): + batch_emb[:, i_sum:i_sum+length] = emb[i_p::num_stages][:,None] + i_sum += length + + batch_shift_msa, batch_scale_msa, batch_gate_msa = batch_emb.chunk(3, dim=2) + + x = self.norm(x) * (1 + batch_scale_msa) + batch_shift_msa + + return x, batch_gate_msa + + def forward( + self, + x: torch.Tensor, + emb: Optional[torch.Tensor] = None, + hidden_length: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + if hidden_length is not None: + return self.forward_with_pad(x, emb, hidden_length) + emb = self.linear(self.silu(emb)) + shift_msa, scale_msa, gate_msa = emb.chunk(3, dim=1) + x = self.norm(x) * (1 + scale_msa[:, None]) + shift_msa[:, None] + return x, gate_msa diff --git a/pyramid_dit/flux_modules/modeling_pyramid_flux.py b/pyramid_dit/flux_modules/modeling_pyramid_flux.py new file mode 100644 index 0000000..b73c562 --- /dev/null +++ b/pyramid_dit/flux_modules/modeling_pyramid_flux.py @@ -0,0 +1,538 @@ +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_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." + + scale = torch.arange(0, dim, 2, dtype=torch.float64, device=pos.device) / dim + omega = 1.0 / (theta**scale) + + batch_size, seq_length = pos.shape + out = torch.einsum("...n,d->...nd", pos, omega) + cos_out = torch.cos(out) + sin_out = torch.sin(out) + + stacked_out = torch.stack([cos_out, -sin_out, sin_out, cos_out], dim=-1) + out = stacked_out.view(batch_size, -1, dim // 2, 2, 2) + return out.float() + + +class EmbedND(nn.Module): + def __init__(self, dim: int, theta: int, axes_dim: List[int]): + super().__init__() + self.dim = dim + self.theta = theta + self.axes_dim = axes_dim + + def forward(self, ids: torch.Tensor) -> torch.Tensor: + n_axes = ids.shape[-1] + emb = torch.cat( + [rope(ids[..., i], self.axes_dim[i], self.theta) for i in range(n_axes)], + dim=-3, + ) + return emb.unsqueeze(2) + + +class PyramidFluxTransformer(ModelMixin, ConfigMixin): + """ + The Transformer model introduced in Flux. + + Reference: https://blackforestlabs.ai/announcing-black-forest-labs/ + + Parameters: + patch_size (`int`): Patch size to turn the input data into small patches. + in_channels (`int`, *optional*, defaults to 16): The number of channels in the input. + num_layers (`int`, *optional*, defaults to 18): The number of layers of MMDiT blocks to use. + num_single_layers (`int`, *optional*, defaults to 18): The number of layers of single DiT blocks to use. + attention_head_dim (`int`, *optional*, defaults to 64): The number of channels in each head. + num_attention_heads (`int`, *optional*, defaults to 18): The number of heads to use for multi-head attention. + joint_attention_dim (`int`, *optional*): The number of `encoder_hidden_states` dimensions to use. + pooled_projection_dim (`int`): Number of dimensions to use when projecting the `pooled_projections`. + """ + + _supports_gradient_checkpointing = True + + @register_to_config + def __init__( + self, + patch_size: int = 1, + in_channels: int = 64, + num_layers: int = 19, + num_single_layers: int = 38, + attention_head_dim: int = 64, + num_attention_heads: int = 24, + joint_attention_dim: int = 4096, + pooled_projection_dim: int = 768, + axes_dims_rope: List[int] = [16, 24, 24], + use_flash_attn: bool = False, + use_temporal_causal: bool = True, + interp_condition_pos: bool = True, + use_gradient_checkpointing: bool = False, + gradient_checkpointing_ratio: float = 0.6, + ): + super().__init__() + self.out_channels = in_channels + self.inner_dim = self.config.num_attention_heads * self.config.attention_head_dim + + self.pos_embed = EmbedND(dim=self.inner_dim, theta=10000, axes_dim=axes_dims_rope) + self.time_text_embed = CombinedTimestepTextProjEmbeddings( + embedding_dim=self.inner_dim, pooled_projection_dim=self.config.pooled_projection_dim + ) + + self.context_embedder = nn.Linear(self.config.joint_attention_dim, self.inner_dim) + self.x_embedder = torch.nn.Linear(self.config.in_channels, self.inner_dim) + + self.transformer_blocks = nn.ModuleList( + [ + FluxTransformerBlock( + dim=self.inner_dim, + num_attention_heads=self.config.num_attention_heads, + attention_head_dim=self.config.attention_head_dim, + use_flash_attn=use_flash_attn, + ) + for i in range(self.config.num_layers) + ] + ) + + self.single_transformer_blocks = nn.ModuleList( + [ + FluxSingleTransformerBlock( + dim=self.inner_dim, + num_attention_heads=self.config.num_attention_heads, + attention_head_dim=self.config.attention_head_dim, + use_flash_attn=use_flash_attn, + ) + for i in range(self.config.num_single_layers) + ] + ) + + self.norm_out = AdaLayerNormContinuous(self.inner_dim, self.inner_dim, elementwise_affine=False, eps=1e-6) + self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=True) + + self.gradient_checkpointing = use_gradient_checkpointing + self.gradient_checkpointing_ratio = gradient_checkpointing_ratio + + self.use_temporal_causal = use_temporal_causal + if self.use_temporal_causal: + print("Using temporal causal attention") + + self.use_flash_attn = use_flash_attn + if self.use_flash_attn: + print("Using Flash attention") + + self.patch_size = 2 # hard-code for now + + # init weights + self.initialize_weights() + + def initialize_weights(self): + # Initialize transformer layers: + def _basic_init(module): + if isinstance(module, (nn.Linear, nn.Conv2d, nn.Conv3d)): + torch.nn.init.xavier_uniform_(module.weight) + if module.bias is not None: + nn.init.constant_(module.bias, 0) + self.apply(_basic_init) + + # Initialize all the conditioning to normal init + nn.init.normal_(self.time_text_embed.timestep_embedder.linear_1.weight, std=0.02) + nn.init.normal_(self.time_text_embed.timestep_embedder.linear_2.weight, std=0.02) + nn.init.normal_(self.time_text_embed.text_embedder.linear_1.weight, std=0.02) + nn.init.normal_(self.time_text_embed.text_embedder.linear_2.weight, std=0.02) + nn.init.normal_(self.context_embedder.weight, std=0.02) + + # Zero-out adaLN modulation layers in DiT blocks: + for block in self.transformer_blocks: + nn.init.constant_(block.norm1.linear.weight, 0) + nn.init.constant_(block.norm1.linear.bias, 0) + nn.init.constant_(block.norm1_context.linear.weight, 0) + nn.init.constant_(block.norm1_context.linear.bias, 0) + + for block in self.single_transformer_blocks: + nn.init.constant_(block.norm.linear.weight, 0) + nn.init.constant_(block.norm.linear.bias, 0) + + # Zero-out output layers: + nn.init.constant_(self.norm_out.linear.weight, 0) + nn.init.constant_(self.norm_out.linear.bias, 0) + nn.init.constant_(self.proj_out.weight, 0) + nn.init.constant_(self.proj_out.bias, 0) + + @torch.no_grad() + def _prepare_image_ids(self, batch_size, temp, height, width, train_height, train_width, device, start_time_stamp=0): + latent_image_ids = torch.zeros(temp, height, width, 3) + + # Temporal Rope + latent_image_ids[..., 0] = latent_image_ids[..., 0] + torch.arange(start_time_stamp, start_time_stamp + temp)[:, None, None] + + # height Rope + if height != train_height: + height_pos = F.interpolate(torch.arange(train_height)[None, None, :].float(), height, mode='linear').squeeze(0, 1) + else: + height_pos = torch.arange(train_height).float() + + latent_image_ids[..., 1] = latent_image_ids[..., 1] + height_pos[None, :, None] + + # width rope + if width != train_width: + width_pos = F.interpolate(torch.arange(train_width)[None, None, :].float(), width, mode='linear').squeeze(0, 1) + else: + width_pos = torch.arange(train_width).float() + + latent_image_ids[..., 2] = latent_image_ids[..., 2] + width_pos[None, None, :] + + latent_image_ids = latent_image_ids[None, :].repeat(batch_size, 1, 1, 1, 1) + latent_image_ids = rearrange(latent_image_ids, 'b t h w c -> b (t h w) c') + + return latent_image_ids.to(device=device) + + @torch.no_grad() + def _prepare_pyramid_image_ids(self, sample, batch_size, device): + image_ids_list = [] + + for i_b, sample_ in enumerate(sample): + if not isinstance(sample_, list): + sample_ = [sample_] + + cur_image_ids = [] + start_time_stamp = 0 + + train_height = sample_[-1].shape[-2] // self.patch_size + train_width = sample_[-1].shape[-1] // self.patch_size + + for clip_ in sample_: + _, _, temp, height, width = clip_.shape + height = height // self.patch_size + width = width // self.patch_size + cur_image_ids.append(self._prepare_image_ids(batch_size, temp, height, width, train_height, train_width, device, start_time_stamp=start_time_stamp)) + start_time_stamp += temp + + cur_image_ids = torch.cat(cur_image_ids, dim=1) + image_ids_list.append(cur_image_ids) + + return image_ids_list + + def merge_input(self, sample, encoder_hidden_length, encoder_attention_mask): + """ + Merge the input video with different resolutions into one sequence + Sample: From low resolution to high resolution + """ + if isinstance(sample[0], list): + device = sample[0][-1].device + pad_batch_size = sample[0][-1].shape[0] + else: + device = sample[0].device + pad_batch_size = sample[0].shape[0] + + num_stages = len(sample) + height_list = [];width_list = [];temp_list = [] + trainable_token_list = [] + + for i_b, sample_ in enumerate(sample): + if isinstance(sample_, list): + sample_ = sample_[-1] + _, _, temp, height, width = sample_.shape + height = height // self.patch_size + width = width // self.patch_size + temp_list.append(temp) + height_list.append(height) + width_list.append(width) + trainable_token_list.append(height * width * temp) + + # prepare the RoPE IDs, + image_ids_list = self._prepare_pyramid_image_ids(sample, pad_batch_size, device) + text_ids = torch.zeros(pad_batch_size, encoder_attention_mask.shape[1], 3).to(device=device) + input_ids_list = [torch.cat([text_ids, image_ids], dim=1) for image_ids in image_ids_list] + image_rotary_emb = [self.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: + video_tokens = [] + + for each_latent in sample_: + each_latent = rearrange(each_latent, 'b c t h w -> b t h w c') + each_latent = rearrange(each_latent, 'b t (h p1) (w p2) c -> b (t h w) (p1 p2 c)', p1=self.patch_size, p2=self.patch_size) + video_tokens.append(each_latent) + + video_tokens = torch.cat(video_tokens, dim=1) + video_tokens = self.x_embedder(video_tokens) + hidden_states.append(video_tokens) + hidden_length.append(video_tokens.shape[1]) + + # prepare the attention mask + if self.use_flash_attn: + attention_mask = None + indices_list = [] + 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() + + indices_list.append( + { + 'indices': indices, + 'seqlens_in_batch': seqlens_in_batch, + } + ) + encoder_attention_mask = indices_list + else: + assert encoder_attention_mask.shape[1] == encoder_hidden_length + real_batch_size = encoder_attention_mask.shape[0] + + # prepare text ids + text_ids = torch.arange(1, real_batch_size + 1, dtype=encoder_attention_mask.dtype).unsqueeze(1).repeat(1, encoder_hidden_length) + text_ids = text_ids.to(device) + text_ids[encoder_attention_mask == 0] = 0 + + # prepare image ids + image_ids = torch.arange(1, real_batch_size + 1, dtype=encoder_attention_mask.dtype).unsqueeze(1).repeat(1, max(hidden_length)) + image_ids = image_ids.to(device) + 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)): + image_ids = image_ids_list[i_p] + token_ids = torch.cat([text_ids[i_p::num_stages], image_ids], dim=1) + stage_attention_mask = rearrange(token_ids, 'b i -> b 1 i 1') == rearrange(token_ids, 'b j -> b 1 1 j') # [bs, 1, q_len, k_len] + if self.use_temporal_causal: + input_order_ids = input_ids_list[i_p][:,:,0] + temporal_causal_mask = rearrange(input_order_ids, 'b i -> b 1 i 1') >= rearrange(input_order_ids, 'b j -> b 1 1 j') + stage_attention_mask = stage_attention_mask & temporal_causal_mask + attention_mask.append(stage_attention_mask) + + return hidden_states, hidden_length, temp_list, height_list, width_list, trainable_token_list, encoder_attention_mask, attention_mask, image_rotary_emb + + def split_output(self, batch_hidden_states, hidden_length, temps, heights, widths, trainable_token_list): + # To split the hidden states + batch_size = batch_hidden_states.shape[0] + 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:] + + # unpatchify + hidden_states = hidden_states.reshape( + shape=(batch_size, temp, height, width, self.patch_size, self.patch_size, self.out_channels // 4) + ) + hidden_states = rearrange(hidden_states, "b t h w p1 p2 c -> b t (h p1) (w p2) c") + hidden_states = rearrange(hidden_states, "b t h w c -> b c t h w") + output_hidden_list.append(hidden_states) + + return output_hidden_list + + def forward( + self, + sample: torch.FloatTensor, # [num_stages] + encoder_hidden_states: torch.Tensor = None, + encoder_attention_mask: torch.FloatTensor = None, + pooled_projections: torch.Tensor = None, + timestep_ratio: torch.LongTensor = None, + ): + temb = self.time_text_embed(timestep_ratio, pooled_projections) + encoder_hidden_states = self.context_embedder(encoder_hidden_states) + encoder_hidden_length = encoder_hidden_states.shape[1] + + # Get the input sequence + 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) + + 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, + ) + + # 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, + ) + + 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) + + 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) + hidden_states = self.proj_out(hidden_states) + + output = self.split_output(hidden_states, hidden_length, temps, heights, widths, trainable_token_list) + + return output \ No newline at end of file diff --git a/pyramid_dit/flux_modules/modeling_text_encoder.py b/pyramid_dit/flux_modules/modeling_text_encoder.py new file mode 100644 index 0000000..aaffbb2 --- /dev/null +++ b/pyramid_dit/flux_modules/modeling_text_encoder.py @@ -0,0 +1,134 @@ +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, + ) + + 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 \ No newline at end of file diff --git a/pyramid_dit/mmdit_modules/__init__.py b/pyramid_dit/mmdit_modules/__init__.py new file mode 100644 index 0000000..9ac4925 --- /dev/null +++ b/pyramid_dit/mmdit_modules/__init__.py @@ -0,0 +1,2 @@ +from .modeling_pyramid_mmdit import PyramidDiffusionMMDiT +from .modeling_text_encoder import SD3TextEncoderWithMask \ No newline at end of file diff --git a/pyramid_dit/modeling_embedding.py b/pyramid_dit/mmdit_modules/modeling_embedding.py similarity index 100% rename from pyramid_dit/modeling_embedding.py rename to pyramid_dit/mmdit_modules/modeling_embedding.py diff --git a/pyramid_dit/modeling_mmdit_block.py b/pyramid_dit/mmdit_modules/modeling_mmdit_block.py similarity index 99% rename from pyramid_dit/modeling_mmdit_block.py rename to pyramid_dit/mmdit_modules/modeling_mmdit_block.py index 0160f47..5f87237 100644 --- a/pyramid_dit/modeling_mmdit_block.py +++ b/pyramid_dit/mmdit_modules/modeling_mmdit_block.py @@ -15,7 +15,7 @@ except: flash_attn_varlen_func = None print("Please install flash attention") -from ..trainer_misc import ( +from ...trainer_misc import ( is_sequence_parallel_initialized, get_sequence_parallel_group, get_sequence_parallel_world_size, diff --git a/pyramid_dit/modeling_normalization.py b/pyramid_dit/mmdit_modules/modeling_normalization.py similarity index 100% rename from pyramid_dit/modeling_normalization.py rename to pyramid_dit/mmdit_modules/modeling_normalization.py diff --git a/pyramid_dit/modeling_pyramid_mmdit.py b/pyramid_dit/mmdit_modules/modeling_pyramid_mmdit.py similarity index 99% rename from pyramid_dit/modeling_pyramid_mmdit.py rename to pyramid_dit/mmdit_modules/modeling_pyramid_mmdit.py index 80765c5..cf0209a 100644 --- a/pyramid_dit/modeling_pyramid_mmdit.py +++ b/pyramid_dit/mmdit_modules/modeling_pyramid_mmdit.py @@ -13,7 +13,7 @@ from .modeling_embedding import PatchEmbed3D, CombinedTimestepConditionEmbedding from .modeling_normalization import AdaLayerNormContinuous from .modeling_mmdit_block import JointTransformerBlock -from ..trainer_misc import ( +from ...trainer_misc import ( is_sequence_parallel_initialized, get_sequence_parallel_group, get_sequence_parallel_world_size, diff --git a/pyramid_dit/modeling_text_encoder.py b/pyramid_dit/mmdit_modules/modeling_text_encoder.py similarity index 100% rename from pyramid_dit/modeling_text_encoder.py rename to pyramid_dit/mmdit_modules/modeling_text_encoder.py diff --git a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py index 7708804..d77f7ee 100644 --- a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py +++ b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py @@ -14,9 +14,16 @@ from typing import List, Optional, Union from ..diffusion_schedulers import PyramidFlowMatchEulerDiscreteScheduler from ..video_vae.modeling_causal_vae import CausalVideoVAE -from .modeling_pyramid_mmdit import PyramidDiffusionMMDiT -from .modeling_text_encoder import SD3TextEncoderWithMask +from .mmdit_modules import ( + PyramidDiffusionMMDiT, + SD3TextEncoderWithMask, +) + +from .flux_modules import ( + PyramidFluxTransformer, + FluxTextEncoderWithMask, +) from comfy.utils import ProgressBar def compute_density_for_timestep_sampling( @@ -34,12 +41,66 @@ def compute_density_for_timestep_sampling( return u +def build_pyramid_dit( + model_name : str, + model_path : str, + torch_dtype, + use_flash_attn : bool, + #use_mixed_training: bool, + interp_condition_pos: bool = True, + use_gradient_checkpointing: bool = False, + use_temporal_causal: bool = True, + gradient_checkpointing_ratio: float = 0.6, +): + #model_dtype = torch.float32 if use_mixed_training else torch_dtype + if model_name == "pyramid_flux": + dit = PyramidFluxTransformer.from_pretrained( + model_path, torch_dtype=torch_dtype, + use_gradient_checkpointing=use_gradient_checkpointing, + gradient_checkpointing_ratio=gradient_checkpointing_ratio, + use_flash_attn=use_flash_attn, use_temporal_causal=use_temporal_causal, + interp_condition_pos=interp_condition_pos, axes_dims_rope=[16, 24, 24], + ) + elif model_name == "pyramid_mmdit": + dit = PyramidDiffusionMMDiT.from_pretrained( + model_path, torch_dtype=torch_dtype, use_gradient_checkpointing=use_gradient_checkpointing, + gradient_checkpointing_ratio=gradient_checkpointing_ratio, + use_flash_attn=use_flash_attn, use_t5_mask=True, + add_temp_pos_embed=True, temp_pos_embed_type='rope', + use_temporal_causal=use_temporal_causal, interp_condition_pos=interp_condition_pos, + ) + else: + raise NotImplementedError(f"Unsupported DiT architecture, please set the model_name to `pyramid_flux` or `pyramid_mmdit`") + + return dit + + +def build_text_encoder( + model_name : str, + model_path : str, + torch_dtype, + load_text_encoder: bool = True, +): + # The text encoder + if load_text_encoder: + if model_name == "pyramid_flux": + text_encoder = FluxTextEncoderWithMask(model_path, torch_dtype=torch_dtype) + elif model_name == "pyramid_mmdit": + text_encoder = SD3TextEncoderWithMask(model_path, torch_dtype=torch_dtype) + else: + raise NotImplementedError(f"Unsupported Text Encoder architecture, please set the model_name to `pyramid_flux` or `pyramid_mmdit`") + else: + text_encoder = None + + return text_encoder + + class PyramidDiTForVideoGeneration: """ The pyramid dit for both image and video generation, The running class wrapper This class is mainly for fixed unit implementation: 1 + n + n + n """ - def __init__(self, model_path, model_dtype, text_encoder_dtype, vae_dtype, use_gradient_checkpointing=False, return_log=True, + def __init__(self, model_path, model_dtype, model_name, text_encoder_dtype, vae_dtype, use_gradient_checkpointing=False, return_log=True, model_variant="diffusion_transformer_768p", timestep_shift=1.0, stage_range=[0, 1/3, 2/3, 1], sample_ratios=[1, 1, 1], scheduler_gamma=1/3, use_flash_attn=False, load_text_encoder=True, load_vae=True, max_temporal_length=31, frame_per_unit=1, use_temporal_causal=True, @@ -55,23 +116,22 @@ class PyramidDiTForVideoGeneration: self.stages = stages self.sample_ratios = sample_ratios self.corrupt_ratio = corrupt_ratio + self.model_name = model_name dit_path = os.path.join(model_path, model_variant) - self.dit = PyramidDiffusionMMDiT.from_pretrained( - dit_path, - torch_dtype=self.dtype, - use_gradient_checkpointing=use_gradient_checkpointing, + # The dit + self.dit = build_pyramid_dit( + model_name, dit_path, self.dtype, use_flash_attn=use_flash_attn, - use_t5_mask=True, - add_temp_pos_embed=True, - temp_pos_embed_type='rope', - use_temporal_causal=True if not use_flash_attn else False, - interp_condition_pos=interp_condition_pos, - ) + interp_condition_pos=interp_condition_pos, use_gradient_checkpointing=use_gradient_checkpointing, + use_temporal_causal=use_temporal_causal, + ) + + if model_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]: for name, param in self.dit.named_parameters(): if name != "pos_embedding": @@ -82,10 +142,10 @@ class PyramidDiTForVideoGeneration: convert_fp8_linear(self.dit, torch.bfloat16) # The text encoder - if load_text_encoder: - self.text_encoder = SD3TextEncoderWithMask(model_path, torch_dtype=text_encoder_dtype) - else: - self.text_encoder = None + self.text_encoder = build_text_encoder( + model_name, model_path, text_encoder_dtype, load_text_encoder=load_text_encoder, + ) + self.load_text_encoder = load_text_encoder # The base video vae decoder if load_vae: @@ -280,9 +340,9 @@ class PyramidDiTForVideoGeneration: sample=latents, generator=generator, ).prev_sample - nan_mask = torch.isnan(latents) - if torch.any(nan_mask): - raise ValueError("nan in latents") + #nan_mask = torch.isnan(latents) + #if torch.any(nan_mask): + # raise ValueError("nan in latents") intermed_latents.append(latents) @@ -358,7 +418,7 @@ class PyramidDiTForVideoGeneration: # Create the initial random noise - num_channels_latents = self.dit.config.in_channels + 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( batch_size * num_images_per_prompt, num_channels_latents, @@ -537,7 +597,7 @@ class PyramidDiTForVideoGeneration: prompt_attention_mask = prompt_attention_mask.to(self.dtype) # Create the initial random noise - num_channels_latents = self.dit.config.in_channels + 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( batch_size * num_images_per_prompt, num_channels_latents,