From 861a378edfc21bdf40318ada1f4bc415a7735ace Mon Sep 17 00:00:00 2001 From: junsong Date: Sat, 30 Nov 2024 05:11:01 -0800 Subject: [PATCH 1/4] sucess add dcae into the repo; --- VAE/conf.py | 20 + VAE/loader.py | 3 + VAE/models/dcae.py | 1018 ++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 1041 insertions(+) create mode 100644 VAE/models/dcae.py diff --git a/VAE/conf.py b/VAE/conf.py index 87112dd..19710de 100644 --- a/VAE/conf.py +++ b/VAE/conf.py @@ -155,5 +155,25 @@ vae_conf = { "ch_mult" : [1, 2, 2, 4], "num_res_blocks" : 2, "attn_resolutions" : [32], + }, + # DCAE configs + "dcae-f32c32-sana-1.0": { + "type" : "DCAE", + "in_channels" : 3, + "embed_scale" : 32, + "embed_dim" : 32, + "encoder_block_type" : ["ResBlock", "ResBlock", "ResBlock", "EViTS5GLU", "EViTS5GLU", "EViTS5GLU"], + "encoder_width_list" : [128, 256, 512, 512, 1024, 1024], + "encoder_depth_list" : [2, 2, 2, 3, 3, 3], + "encoder_norm" : "rms2d", + "encoder_act" : "silu", + "downsample_block_type" : "Conv", + "decoder_block_type" : ["ResBlock", "ResBlock", "ResBlock", "EViTS5GLU", "EViTS5GLU", "EViTS5GLU"], + "decoder_width_list" : [128, 256, 512, 512, 1024, 1024], + "decoder_depth_list" : [3, 3, 3, 3, 3, 3], + "decoder_norm" : "rms2d", + "decoder_act" : "silu", + "upsample_block_type" : "InterpolateConv", + "scaling_factor" : 0.41407 } } diff --git a/VAE/loader.py b/VAE/loader.py index 88de4a8..748e3ab 100644 --- a/VAE/loader.py +++ b/VAE/loader.py @@ -32,6 +32,9 @@ class EXVAE(comfy.sd.VAE): elif model_conf["type"] == "MoVQ3": from .models.movq3 import MoVQ model = MoVQ(model_conf) + elif model_conf["type"] == "DCAE": + from .models.dcae import DCAE + model = DCAE(**model_conf) else: raise NotImplementedError(f"Unknown VAE type '{model_conf['type']}'") diff --git a/VAE/models/dcae.py b/VAE/models/dcae.py new file mode 100644 index 0000000..ecfdde9 --- /dev/null +++ b/VAE/models/dcae.py @@ -0,0 +1,1018 @@ +# Copyright 2024 MIT, Tsinghua University, NVIDIA CORPORATION and The HuggingFace Team. +# All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from typing import Optional, Tuple + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +ACTIVATION_FUNCTIONS = { + "swish": nn.SiLU(), + "silu": nn.SiLU(), + "mish": nn.Mish(), + "gelu": nn.GELU(), + "relu": nn.ReLU(), +} + + +def get_activation(act_fn: str) -> nn.Module: + + """Helper function to get activation function from string. + + Args: + act_fn (str): Name of activation function. + + Returns: + nn.Module: Activation function. + """ + + act_fn = act_fn.lower() + if act_fn in ACTIVATION_FUNCTIONS: + return ACTIVATION_FUNCTIONS[act_fn] + else: + raise ValueError(f"Unsupported activation function: {act_fn}") + + +class ConvPixelShuffleUpsample2D(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size: int, + factor: int, + ): + super().__init__() + self.factor = factor + out_ratio = factor**2 + self.conv = nn.Conv2d( + in_channels=in_channels, + out_channels=out_channels * out_ratio, + kernel_size=kernel_size, + padding=kernel_size // 2, + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.conv(x) + x = F.pixel_shuffle(x, self.factor) + return x + + +class ChannelDuplicatingPixelUnshuffleUpsample2D(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + factor: int, + ): + super().__init__() + self.in_channels = in_channels + self.out_channels = out_channels + self.factor = factor + assert out_channels * factor**2 % in_channels == 0 + self.repeats = out_channels * factor**2 // in_channels + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = x.repeat_interleave(self.repeats, dim=1) + x = F.pixel_shuffle(x, self.factor) + return x + + +class Upsample2D(nn.Module): + """A 2D upsampling layer with an optional convolution. + + Parameters: + channels (`int`): + number of channels in the inputs and outputs. + use_conv (`bool`, default `False`): + option to use a convolution. + use_conv_transpose (`bool`, default `False`): + option to use a convolution transpose. + out_channels (`int`, optional): + number of output channels. Defaults to `channels`. + name (`str`, default `conv`): + name of the upsampling 2D layer. + """ + + def __init__( + self, + channels: int, + use_conv: bool = False, + use_conv_transpose: bool = False, + out_channels: Optional[int] = None, + name: str = "conv", + kernel_size: Optional[int] = None, + padding=1, + norm_type=None, + eps=None, + elementwise_affine=None, + bias=True, + interpolate=True, + ): + super().__init__() + self.channels = channels + self.out_channels = out_channels or channels + self.use_conv = use_conv + self.use_conv_transpose = use_conv_transpose + self.name = name + self.interpolate = interpolate + + if norm_type == "ln_norm": + self.norm = nn.LayerNorm(channels, eps, elementwise_affine) + elif norm_type is None: + self.norm = None + else: + raise ValueError(f"unknown norm_type: {norm_type}") + + conv = None + if use_conv_transpose: + if kernel_size is None: + kernel_size = 4 + conv = nn.ConvTranspose2d( + channels, self.out_channels, kernel_size=kernel_size, stride=2, padding=padding, bias=bias + ) + elif use_conv: + if kernel_size is None: + kernel_size = 3 + conv = nn.Conv2d(self.channels, self.out_channels, kernel_size=kernel_size, padding=padding, bias=bias) + + # TODO(Suraj, Patrick) - clean up after weight dicts are correctly renamed + if name == "conv": + self.conv = conv + else: + self.Conv2d_0 = conv + + def forward(self, hidden_states: torch.Tensor, output_size: Optional[int] = None, *args, **kwargs) -> torch.Tensor: + + assert hidden_states.shape[1] == self.channels + + if self.norm is not None: + hidden_states = self.norm(hidden_states.permute(0, 2, 3, 1)).permute(0, 3, 1, 2) + + if self.use_conv_transpose: + return self.conv(hidden_states) + + # Cast to float32 to as 'upsample_nearest2d_out_frame' op does not support bfloat16 until PyTorch 2.1 + # https://github.com/pytorch/pytorch/issues/86679#issuecomment-1783978767 + dtype = hidden_states.dtype + if dtype == torch.bfloat16: + hidden_states = hidden_states.to(torch.float32) + + # upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984 + if hidden_states.shape[0] >= 64: + hidden_states = hidden_states.contiguous() + + # if `output_size` is passed we force the interpolation output + # size and do not make use of `scale_factor=2` + if self.interpolate: + if output_size is None: + hidden_states = F.interpolate(hidden_states, scale_factor=2.0, mode="nearest") + else: + hidden_states = F.interpolate(hidden_states, size=output_size, mode="nearest") + + # Cast back to original dtype + if dtype == torch.bfloat16: + hidden_states = hidden_states.to(dtype) + + # TODO(Suraj, Patrick) - clean up after weight dicts are correctly renamed + if self.use_conv: + if self.name == "conv": + hidden_states = self.conv(hidden_states) + else: + hidden_states = self.Conv2d_0(hidden_states) + + return hidden_states + + +class RMSNorm2d(nn.Module): + def __init__(self, num_features: int, eps: float = 1e-5, elementwise_affine: bool = True, bias: bool = True) -> None: + super().__init__() + self.num_features = num_features + self.eps = eps + self.elementwise_affine = elementwise_affine + if self.elementwise_affine: + self.weight = torch.nn.parameter.Parameter(torch.empty(self.num_features)) + if bias: + self.bias = torch.nn.parameter.Parameter(torch.empty(self.num_features)) + else: + self.register_parameter('bias', None) + else: + self.register_parameter('weight', None) + self.register_parameter('bias', None) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = (x / torch.sqrt(torch.square(x.float()).mean(dim=1, keepdim=True) + self.eps)).to(x.dtype) + if self.elementwise_affine: + x = x * self.weight.view(1, -1, 1, 1) + self.bias.view(1, -1, 1, 1) + return x + + +class DCAELiteMLA(nn.Module): + r"""Lightweight multi-scale linear attention used in DC-AE""" + + def __init__( + self, + in_channels: int, + out_channels: int, + heads: Optional[int] = None, + heads_ratio: float = 1.0, + dim=8, + use_bias=(False, False), + norm=(None, "bn2d"), + act_func=(None, None), + kernel_func="relu", + scales: Tuple[int, ...] = (5,), + eps=1.0e-15, + ): + super().__init__() + self.eps = eps + heads = int(in_channels // dim * heads_ratio) if heads is None else heads + + total_dim = heads * dim + + self.dim = dim + + qkv = [nn.Conv2d(in_channels=in_channels, out_channels=3 * total_dim, kernel_size=1, bias=use_bias[0])] + if norm[0] is None: + pass + elif norm[0] == "rms2d": + qkv.append(RMSNorm2d(num_features=3 * total_dim)) + else: + raise ValueError(f"norm {norm[0]} is not supported") + if act_func[0] is not None: + qkv.append(get_activation(act_func[0])) + self.qkv = nn.Sequential(*qkv) + + self.aggreg = nn.ModuleList( + [ + nn.Sequential( + nn.Conv2d( + 3 * total_dim, + 3 * total_dim, + scale, + padding=scale // 2, + groups=3 * total_dim, + bias=use_bias[0], + ), + nn.Conv2d(3 * total_dim, 3 * total_dim, 1, groups=3 * heads, bias=use_bias[0]), + ) + for scale in scales + ] + ) + self.kernel_func = get_activation(kernel_func) + + proj = [nn.Conv2d(in_channels=total_dim * (1 + len(scales)), out_channels=out_channels, kernel_size=1, bias=use_bias[1])] + if norm[1] is None: + pass + elif norm[1] == "rms2d": + proj.append(RMSNorm2d(num_features=out_channels)) + else: + raise ValueError(f"norm {norm[1]} is not supported") + if act_func[1] is not None: + proj.append(get_activation(act_func[1])) + self.proj = nn.Sequential(*proj) + + def relu_linear_att(self, qkv: torch.Tensor) -> torch.Tensor: + B, _, H, W = list(qkv.size()) + + if qkv.dtype == torch.float16: + qkv = qkv.float() + + qkv = torch.reshape( + qkv, + ( + B, + -1, + 3 * self.dim, + H * W, + ), + ) + q, k, v = ( + qkv[:, :, 0 : self.dim], + qkv[:, :, self.dim : 2 * self.dim], + qkv[:, :, 2 * self.dim :], + ) + + # lightweight linear attention + q = self.kernel_func(q) + k = self.kernel_func(k) + + # linear matmul + trans_k = k.transpose(-1, -2) + + v = F.pad(v, (0, 0, 0, 1), mode="constant", value=1) + vk = torch.matmul(v, trans_k) + out = torch.matmul(vk, q) + if out.dtype == torch.bfloat16: + out = out.float() + out = out[:, :, :-1] / (out[:, :, -1:] + self.eps) + + out = torch.reshape(out, (B, -1, H, W)) + return out + + def relu_quadratic_att(self, qkv: torch.Tensor) -> torch.Tensor: + B, _, H, W = list(qkv.size()) + + qkv = torch.reshape( + qkv, + ( + B, + -1, + 3 * self.dim, + H * W, + ), + ) + q, k, v = ( + qkv[:, :, 0 : self.dim], + qkv[:, :, self.dim : 2 * self.dim], + qkv[:, :, 2 * self.dim :], + ) + + q = self.kernel_func(q) + k = self.kernel_func(k) + + att_map = torch.matmul(k.transpose(-1, -2), q) # b h n n + original_dtype = att_map.dtype + if original_dtype in [torch.float16, torch.bfloat16]: + att_map = att_map.float() + att_map = att_map / (torch.sum(att_map, dim=2, keepdim=True) + self.eps) # b h n n + att_map = att_map.to(original_dtype) + out = torch.matmul(v, att_map) # b h d n + + out = torch.reshape(out, (B, -1, H, W)) + return out + + def forward(self, x: torch.Tensor) -> torch.Tensor: + # generate multi-scale q, k, v + qkv = self.qkv(x) + multi_scale_qkv = [qkv] + for op in self.aggreg: + multi_scale_qkv.append(op(qkv)) + qkv = torch.cat(multi_scale_qkv, dim=1) + + H, W = list(qkv.size())[-2:] + if H * W > self.dim: + out = self.relu_linear_att(qkv).to(qkv.dtype) + else: + out = self.relu_quadratic_att(qkv) + out = self.proj(out) + + return x + out + + +class ConvPixelUnshuffleDownsample2D(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size: int, + factor: int, + ): + super().__init__() + self.factor = factor + out_ratio = factor**2 + assert out_channels % out_ratio == 0 + self.conv = nn.Conv2d( + in_channels=in_channels, + out_channels=out_channels // out_ratio, + kernel_size=kernel_size, + padding=kernel_size // 2, + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.conv(x) + x = F.pixel_unshuffle(x, self.factor) + return x + + +class PixelUnshuffleChannelAveragingDownsample2D(nn.Module): + + def __init__( + self, + in_channels: int, + out_channels: int, + factor: int, + ): + super().__init__() + self.in_channels = in_channels + self.out_channels = out_channels + self.factor = factor + assert in_channels * factor**2 % out_channels == 0 + self.group_size = in_channels * factor**2 // out_channels + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = F.pixel_unshuffle(x, self.factor) + B, C, H, W = x.shape + x = x.view(B, self.out_channels, self.group_size, H, W) + x = x.mean(dim=2) + return x + + +class ConvLayer(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size=3, + stride=1, + dilation=1, + groups=1, + use_bias=False, + dropout=0, + norm="bn2d", + act_func="relu", + ): + super().__init__() + + padding = kernel_size // 2 + padding *= dilation + + self.dropout = nn.Dropout2d(dropout, inplace=False) if dropout > 0 else None + self.conv = nn.Conv2d( + in_channels, + out_channels, + kernel_size=(kernel_size, kernel_size), + stride=(stride, stride), + padding=padding, + dilation=(dilation, dilation), + groups=groups, + bias=use_bias, + ) + if norm is None: + self.norm = None + elif norm == "rms2d": + self.norm = RMSNorm2d(num_features=out_channels) + else: + raise ValueError(f"norm {norm} is not supported") + self.act = get_activation(act_func) if act_func is not None else None + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if self.dropout is not None: + x = self.dropout(x) + x = self.conv(x) + if self.norm: + x = self.norm(x) + if self.act: + x = self.act(x) + return x + + +class GLUMBConv(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size=3, + stride=1, + mid_channels=None, + expand_ratio=6, + use_bias=(False, False, False), + norm=(None, None, "ln2d"), + act_func=("silu", "silu", None), + ): + super().__init__() + + mid_channels = round(in_channels * expand_ratio) if mid_channels is None else mid_channels + + self.glu_act = get_activation(act_func[1]) + self.inverted_conv = ConvLayer( + in_channels, + mid_channels * 2, + 1, + use_bias=use_bias[0], + norm=norm[0], + act_func=act_func[0], + ) + self.depth_conv = ConvLayer( + mid_channels * 2, + mid_channels * 2, + kernel_size, + stride=stride, + groups=mid_channels * 2, + use_bias=use_bias[1], + norm=norm[1], + act_func=None, + ) + self.point_conv = ConvLayer( + mid_channels, + out_channels, + 1, + use_bias=use_bias[2], + norm=norm[2], + act_func=act_func[2], + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + y = self.inverted_conv(x) + y = self.depth_conv(y) + + y, gate = torch.chunk(y, 2, dim=1) + gate = self.glu_act(gate) + y = y * gate + + y = self.point_conv(y) + return x + y + + +class ResBlock(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size=3, + stride=1, + mid_channels=None, + expand_ratio=1, + use_bias=(False, False), + norm=("bn2d", "bn2d"), + act_func=("relu6", None), + ): + super().__init__() + mid_channels = round(in_channels * expand_ratio) if mid_channels is None else mid_channels + + self.conv1 = ConvLayer( + in_channels, + mid_channels, + kernel_size, + stride, + use_bias=use_bias[0], + norm=norm[0], + act_func=act_func[0], + ) + self.conv2 = ConvLayer( + mid_channels, + out_channels, + kernel_size, + 1, + use_bias=use_bias[1], + norm=norm[1], + act_func=act_func[1], + ) + self.shortcut = nn.Identity() + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.conv2(self.conv1(x)) + x + return x + + +class EfficientViTBlock(nn.Module): + def __init__( + self, + in_channels: int, + heads_ratio: float = 1.0, + dim=32, + expand_ratio: float = 4, + scales: tuple[int, ...] = (5,), + norm: str = "bn2d", + act_func: str = "hswish", + context_module: str = "LiteMLA", + local_module: str = "MBConv", + ): + super().__init__() + if context_module == "LiteMLA": + self.context_module = DCAELiteMLA( + in_channels=in_channels, + out_channels=in_channels, + heads_ratio=heads_ratio, + dim=dim, + norm=(None, norm), + scales=scales, + ) + else: + raise ValueError(f"context_module {context_module} is not supported") + if local_module == "GLUMBConv": + self.local_module = GLUMBConv( + in_channels=in_channels, + out_channels=in_channels, + expand_ratio=expand_ratio, + use_bias=(True, True, False), + norm=(None, None, norm), + act_func=(act_func, act_func, None), + ) + else: + raise NotImplementedError(f"local_module {local_module} is not supported") + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.context_module(x) + x = self.local_module(x) + return x + + +################################################################################# +# Functional Blocks # +################################################################################# + + +class ResidualBlock(nn.Module): + def __init__( + self, + main: Optional[nn.Module], + shortcut: Optional[nn.Module], + post_act=None, + pre_norm: Optional[nn.Module] = None, + ): + super().__init__() + + self.pre_norm = pre_norm + self.main = main + self.shortcut = shortcut + self.post_act = get_activation(post_act) if post_act is not None else None + + def forward_main(self, x: torch.Tensor) -> torch.Tensor: + if self.pre_norm is None: + return self.main(x) + else: + return self.main(self.pre_norm(x)) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if self.main is None: + res = x + elif self.shortcut is None: + res = self.forward_main(x) + else: + res = self.forward_main(x) + self.shortcut(x) + if self.post_act: + res = self.post_act(res) + return res + + +class Encoder(nn.Module): + def __init__( + self, + in_channels: int, + latent_channels: int, + width_list: list[int] = [128, 256, 512, 512, 1024, 1024], + depth_list: list[int] = [2, 2, 2, 2, 2, 2], + block_type: str | list[str] = "ResBlock", + norm: str = "rms2d", + act: str = "silu", + downsample_block_type: str = "ConvPixelUnshuffle", + downsample_shortcut: Optional[str] = "averaging", + out_norm: Optional[str] = None, + out_act: Optional[str] = None, + out_shortcut: Optional[str] = "averaging", + double_latent: bool = False, + ): + super().__init__() + num_stages = len(width_list) + self.num_stages = num_stages + + # validate config + if len(depth_list) != num_stages or len(width_list) != num_stages: + raise ValueError(f"len(depth_list) {len(depth_list)} and len(width_list) {len(width_list)} should be equal to num_stages {num_stages}") + if not isinstance(block_type, (str, list)) or (isinstance(block_type, list) and len(block_type) != num_stages): + raise ValueError(f"block_type should be either a str or a list of str with length {num_stages}, but got {block_type}") + + # project in + if depth_list[0] > 0: + project_in_block = nn.Conv2d( + in_channels=in_channels, + out_channels=width_list[0], + kernel_size=3, + padding=1, + ) + elif depth_list[1] > 0: + if downsample_block_type == "Conv": + project_in_block = nn.Conv2d( + in_channels=in_channels, + out_channels=width_list[1], + kernel_size=3, + stride=2, + padding=1, + ) + elif downsample_block_type == "ConvPixelUnshuffle": + project_in_block = ConvPixelUnshuffleDownsample2D( + in_channels=in_channels, out_channels=width_list[1], kernel_size=3, factor=2 + ) + else: + raise ValueError(f"block_type {downsample_block_type} is not supported for downsampling") + else: + raise ValueError(f"depth list {depth_list} is not supported for encoder project in") + self.project_in = project_in_block + + # stages + self.stages: list[nn.Module] = [] + for stage_id, (width, depth) in enumerate(zip(width_list, depth_list)): + stage_block_type = block_type[stage_id] if isinstance(block_type, list) else block_type + if not (isinstance(stage_block_type, str) or (isinstance(stage_block_type, list) and depth == len(stage_block_type))): + raise ValueError(f"block type {stage_block_type} is not supported for encoder stage {stage_id} with depth {depth}") + stage = [] + # stage main + for d in range(depth): + current_block_type = stage_block_type[d] if isinstance(stage_block_type, list) else stage_block_type + if current_block_type == "ResBlock": + block = ResBlock( + in_channels=width, + out_channels=width, + kernel_size=3, + stride=1, + use_bias=(True, False), + norm=(None, norm), + act_func=(act, None), + ) + elif current_block_type == "EViTGLU": + block = EfficientViTBlock(width, norm=norm, act_func=act, local_module="GLUMBConv", scales=()) + elif current_block_type == "EViTS5GLU": + block = EfficientViTBlock(width, norm=norm, act_func=act, local_module="GLUMBConv", scales=(5,)) + else: + raise ValueError(f"block type {current_block_type} is not supported") + stage.append(block) + # downsample + if stage_id < num_stages - 1 and depth > 0: + downsample_out_channels = width_list[stage_id + 1] + if downsample_block_type == "Conv": + downsample_block = nn.Conv2d( + in_channels=width, + out_channels=downsample_out_channels, + kernel_size=3, + stride=2, + padding=1, + ) + elif downsample_block_type == "ConvPixelUnshuffle": + downsample_block = ConvPixelUnshuffleDownsample2D( + in_channels=width, out_channels=downsample_out_channels, kernel_size=3, factor=2 + ) + else: + raise ValueError(f"downsample_block_type {downsample_block_type} is not supported for downsampling") + if downsample_shortcut is None: + pass + elif downsample_shortcut == "averaging": + shortcut_block = PixelUnshuffleChannelAveragingDownsample2D( + in_channels=width, out_channels=downsample_out_channels, factor=2 + ) + downsample_block = ResidualBlock(downsample_block, shortcut_block) + else: + raise ValueError(f"shortcut {downsample_shortcut} is not supported for downsample") + stage.append(downsample_block) + self.stages.append(nn.Sequential(*stage)) + self.stages = nn.ModuleList(self.stages) + + # project out + project_out_layers: list[nn.Module] = [] + if out_norm is None: + pass + elif out_norm == "rms2d": + project_out_layers.append(RMSNorm2d(num_features=width_list[-1])) + else: + raise ValueError(f"norm {out_norm} is not supported for encoder project out") + if out_act is not None: + project_out_layers.append(get_activation(out_act)) + project_out_out_channels = 2 * latent_channels if double_latent else latent_channels + project_out_layers.append(ConvLayer( + in_channels=width_list[-1], + out_channels=project_out_out_channels, + kernel_size=3, + stride=1, + use_bias=True, + norm=None, + act_func=None, + )) + project_out_block = nn.Sequential(*project_out_layers) + if out_shortcut is None: + pass + elif out_shortcut == "averaging": + shortcut_block = PixelUnshuffleChannelAveragingDownsample2D( + in_channels=width_list[-1], out_channels=project_out_out_channels, factor=1 + ) + project_out_block = ResidualBlock(project_out_block, shortcut_block) + else: + raise ValueError(f"shortcut {out_shortcut} is not supported for encoder project out") + self.project_out = project_out_block + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.project_in(x) + for stage in self.stages: + if len(stage) == 0: + continue + x = stage(x) + x = self.project_out(x) + return x + + +class Decoder(nn.Module): + def __init__( + self, + in_channels: int, + latent_channels: int, + in_shortcut: Optional[str] = "duplicating", + width_list: list[int] = [128, 256, 512, 512, 1024, 1024], + depth_list: list[int] = [2, 2, 2, 2, 2, 2], + block_type: str | list[str] = "ResBlock", + norm: str | list[str] = "rms2d", + act: str | list[str] = "silu", + upsample_block_type: str = "ConvPixelShuffle", + upsample_shortcut: str = "duplicating", + out_norm: str = "rms2d", + out_act: str = "relu", + ): + super().__init__() + num_stages = len(width_list) + self.num_stages = num_stages + + # validate config + if len(depth_list) != num_stages or len(width_list) != num_stages: + raise ValueError(f"len(depth_list) {len(depth_list)} and len(width_list) {len(width_list)} should be equal to num_stages {num_stages}") + if not isinstance(block_type, (str, list)) or (isinstance(block_type, list) and len(block_type) != num_stages): + raise ValueError(f"block_type should be either a str or a list of str with length {num_stages}, but got {block_type}") + if not isinstance(norm, (str, list)) or (isinstance(norm, list) and len(norm) != num_stages): + raise ValueError(f"norm should be either a str or a list of str with length {num_stages}, but got {norm}") + if not isinstance(act, (str, list)) or (isinstance(act, list) and len(act) != num_stages): + raise ValueError(f"act should be either a str or a list of str with length {num_stages}, but got {act}") + + # project in + project_in_block = ConvLayer( + in_channels=latent_channels, + out_channels=width_list[-1], + kernel_size=3, + stride=1, + use_bias=True, + norm=None, + act_func=None, + ) + if in_shortcut is None: + pass + elif in_shortcut == "duplicating": + shortcut_block = ChannelDuplicatingPixelUnshuffleUpsample2D( + in_channels=latent_channels, out_channels=width_list[-1], factor=1 + ) + project_in_block = ResidualBlock(project_in_block, shortcut_block) + else: + raise ValueError(f"shortcut {in_shortcut} is not supported for decoder project in") + self.project_in = project_in_block + + # stages + self.stages: list[nn.Module] = [] + for stage_id, (width, depth) in reversed(list(enumerate(zip(width_list, depth_list)))): + stage = [] + # upsample + if stage_id < num_stages - 1 and depth > 0: + upsample_out_channels = width + if upsample_block_type == "ConvPixelShuffle": + upsample_block = ConvPixelShuffleUpsample2D( + in_channels=width_list[stage_id + 1], out_channels=upsample_out_channels, kernel_size=3, factor=2 + ) + elif upsample_block_type == "InterpolateConv": + upsample_block = Upsample2D(channels=width_list[stage_id + 1], use_conv=True, out_channels=upsample_out_channels) + else: + raise ValueError(f"upsample_block_type {upsample_block_type} is not supported") + if upsample_shortcut is None: + pass + elif upsample_shortcut == "duplicating": + shortcut_block = ChannelDuplicatingPixelUnshuffleUpsample2D( + in_channels=width_list[stage_id + 1], out_channels=upsample_out_channels, factor=2 + ) + upsample_block = ResidualBlock(upsample_block, shortcut_block) + else: + raise ValueError(f"shortcut {upsample_shortcut} is not supported for upsample") + stage.append(upsample_block) + # stage main + stage_block_type = block_type[stage_id] if isinstance(block_type, list) else block_type + stage_norm = norm[stage_id] if isinstance(norm, list) else norm + stage_act = act[stage_id] if isinstance(act, list) else act + for d in range(depth): + current_block_type = stage_block_type[d] if isinstance(stage_block_type, list) else stage_block_type + if current_block_type == "ResBlock": + block = ResBlock( + in_channels=width, + out_channels=width, + kernel_size=3, + stride=1, + use_bias=(True, False), + norm=(None, stage_norm), + act_func=(stage_act, None), + ) + elif current_block_type == "EViTGLU": + block = EfficientViTBlock(width, norm=stage_norm, act_func=stage_act, local_module="GLUMBConv", scales=()) + elif current_block_type == "EViTS5GLU": + block = EfficientViTBlock(width, norm=stage_norm, act_func=stage_act, local_module="GLUMBConv", scales=(5,)) + else: + raise ValueError(f"block type {current_block_type} is not supported") + stage.append(block) + + self.stages.insert(0, nn.Sequential(*stage)) + self.stages = nn.ModuleList(self.stages) + + # project out + project_out_layers: list[nn.Module] = [] + if depth_list[0] > 0: + project_out_in_channels = width_list[0] + elif depth_list[1] > 0: + project_out_in_channels = width_list[1] + else: + raise ValueError(f"depth list {depth_list} is not supported for decoder project out") + if out_norm is None: + pass + elif out_norm == "rms2d": + project_out_layers.append(RMSNorm2d(num_features=project_out_in_channels)) + else: + raise ValueError(f"norm {out_norm} is not supported for decoder project out") + project_out_layers.append(get_activation(out_act)) + if depth_list[0] > 0: + project_out_layers.append( + ConvLayer( + in_channels=project_out_in_channels, + out_channels=in_channels, + kernel_size=3, + stride=1, + use_bias=True, + norm=None, + act_func=None, + ) + ) + elif depth_list[1] > 0: + if upsample_block_type == "ConvPixelShuffle": + project_out_conv = ConvPixelShuffleUpsample2D( + in_channels=project_out_in_channels, out_channels=in_channels, kernel_size=3, factor=2 + ) + elif upsample_block_type == "InterpolateConv": + project_out_conv = Upsample2D(channels=project_out_in_channels, use_conv=True, out_channels=in_channels) + else: + raise ValueError(f"upsample_block_type {upsample_block_type} is not supported for upsampling") + + project_out_layers.append(project_out_conv) + else: + raise ValueError(f"depth list {depth_list} is not supported for decoder project out") + self.project_out = nn.Sequential(*project_out_layers) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.project_in(x) + for stage in reversed(self.stages): + if len(stage) == 0: + continue + x = stage(x) + x = self.project_out(x) + return x + + +class DCAE(nn.Module): + def __init__( + self, + in_channels: int = 3, + embed_dim: int = 32, + encoder_block_type: str | list[str] = "ResBlock", + encoder_width_list: list[int] = [128, 256, 512, 512, 1024, 1024], + encoder_depth_list: list[int] = [2, 2, 2, 2, 2, 2], + encoder_norm: str = "rms2d", + encoder_act: str = "silu", + downsample_block_type: str = "ConvPixelUnshuffle", + decoder_block_type: str | list[str] = "ResBlock", + decoder_width_list: list[int] = [128, 256, 512, 512, 1024, 1024], + decoder_depth_list: list[int] = [2, 2, 2, 2, 2, 2], + decoder_norm: str = "rms2d", + decoder_act: str = "silu", + upsample_block_type: str = "ConvPixelShuffle", + scaling_factor: Optional[float] = None, + **ignore_kwargs, + ): + super().__init__() + self.scaling_factor = scaling_factor + self.encoder = Encoder( + in_channels=in_channels, + latent_channels=embed_dim, + width_list=encoder_width_list, + depth_list=encoder_depth_list, + block_type=encoder_block_type, + norm=encoder_norm, + act=encoder_act, + downsample_block_type=downsample_block_type, + ) + self.decoder = Decoder( + in_channels=in_channels, + latent_channels=embed_dim, + width_list=decoder_width_list, + depth_list=decoder_depth_list, + block_type=decoder_block_type, + norm=decoder_norm, + act=decoder_act, + upsample_block_type=upsample_block_type, + ) + + @property + def spatial_compression_ratio(self) -> int: + return 2 ** (self.decoder.num_stages - 1) + + def encode(self, x: torch.Tensor) -> torch.Tensor: + x = self.encoder(x) + return x + + def decode(self, x: torch.Tensor, return_dict: bool = True) -> torch.Tensor: + x = self.decoder(x) + return x + + def forward(self, x: torch.Tensor, global_step: int) -> torch.Tensor: + x = self.encoder(x) + x = self.decoder(x) + return x, torch.tensor(0), {} From 9ec31c864fab6c77a4be3f589c68c88270160a65 Mon Sep 17 00:00:00 2001 From: junsong Date: Sat, 30 Nov 2024 11:50:15 -0800 Subject: [PATCH 2/4] first run sucessfull with text encoder mask bug not fix; --- Gemma/nodes.py | 128 +++++ Sana/conf.py | 98 ++++ Sana/diffusers_convert.py | 223 +++++++++ Sana/loader.py | 100 ++++ Sana/lora.py | 146 ++++++ Sana/models/act.py | 59 +++ Sana/models/basic_modules.py | 361 +++++++++++++++ Sana/models/norms.py | 225 +++++++++ Sana/models/sana.py | 379 +++++++++++++++ Sana/models/sana_blocks.py | 798 ++++++++++++++++++++++++++++++++ Sana/models/sana_multi_scale.py | 374 +++++++++++++++ Sana/models/utils.py | 591 +++++++++++++++++++++++ Sana/nodes.py | 223 +++++++++ VAE/nodes.py | 32 ++ __init__.py | 9 + 15 files changed, 3746 insertions(+) create mode 100644 Gemma/nodes.py create mode 100644 Sana/conf.py create mode 100644 Sana/diffusers_convert.py create mode 100644 Sana/loader.py create mode 100644 Sana/lora.py create mode 100644 Sana/models/act.py create mode 100644 Sana/models/basic_modules.py create mode 100644 Sana/models/norms.py create mode 100644 Sana/models/sana.py create mode 100644 Sana/models/sana_blocks.py create mode 100644 Sana/models/sana_multi_scale.py create mode 100644 Sana/models/utils.py create mode 100644 Sana/nodes.py diff --git a/Gemma/nodes.py b/Gemma/nodes.py new file mode 100644 index 0000000..55a8a03 --- /dev/null +++ b/Gemma/nodes.py @@ -0,0 +1,128 @@ +import os +import torch +import folder_paths +from transformers import AutoTokenizer, AutoModelForCausalLM +from ..utils.dtype import string_to_dtype +from huggingface_hub import snapshot_download + + +# 初始化自定义文件夹路径 +os.makedirs( + os.path.join(folder_paths.models_dir, "text_encoders"), + exist_ok=True +) +folder_paths.folder_names_and_paths["text_encoders"] = ( + [ + os.path.join(folder_paths.models_dir, "text_encoders"), + *folder_paths.folder_names_and_paths.get("text_encoders", [[],set()])[0] + ], + folder_paths.supported_pt_extensions +) + +dtypes = [ + "default", + "auto (comfy)", + "BF16", + "FP32", + "FP16", +] +try: torch.float8_e5m2 +except AttributeError: print("Torch版本过旧,不支持FP8") +else: dtypes += ["FP8 E4M3", "FP8 E5M2"] + +class GemmaLoader: + @classmethod + def INPUT_TYPES(s): + devices = ["auto", "cpu", "cuda"] + # 支持多GPU + for k in range(1, torch.cuda.device_count()): + devices.append(f"cuda:{k}") + return { + "required": { + "model_name": (["google/gemma-2-2b-it", "unsloth/gemma-2-2b-it-bnb-4bit"],), + "device": (devices, {"default":"cpu"}), + "dtype": (dtypes,), + } + } + RETURN_TYPES = ("GEMMA",) + FUNCTION = "load_model" + CATEGORY = "ExtraModels/Gemma" + TITLE = "Gemma Loader" + + def load_model(self, model_name, device, dtype): + dtype = string_to_dtype(dtype, "text_encoder") + if device == "cpu": + assert dtype in [None, torch.float32], f"Can't use dtype '{dtype}' with CPU! Set dtype to 'default'." + + if model_name == 'google/gemma-2-2b-it': + text_encoder_dir = os.path.join(folder_paths.models_dir, 'text_encoders', 'models--google--gemma-2-2b-it') + if not os.path.exists(os.path.join(text_encoder_dir, 'model.safetensors')): + snapshot_download('google/gemma-2-2b-it', local_dir=text_encoder_dir) + elif model_name == 'unsloth/gemma-2-2b-it-bnb-4bit': + text_encoder_dir = os.path.join(folder_paths.models_dir, 'text_encoders', 'models--unsloth--gemma-2-2b-it-bnb-4bit') + if not os.path.exists(os.path.join(text_encoder_dir, 'model.safetensors')): + snapshot_download('unsloth/gemma-2-2b-it-bnb-4bit', local_dir=text_encoder_dir) + else: + raise ValueError('Not implemented!') + + tokenizer = AutoTokenizer.from_pretrained(model_name) + text_encoder_model = AutoModelForCausalLM.from_pretrained(model_name, torch_dtype=dtype) + tokenizer.padding_side = "right" + text_encoder = text_encoder_model.get_decoder() + + if device != "cpu": + text_encoder = text_encoder.to(device) + + return ({ + "tokenizer": tokenizer, + "text_encoder": text_encoder, + "text_encoder_model": text_encoder_model + },) + + +class GemmaTextEncode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "text": ("STRING", {"multiline": True}), + "GEMMA": ("GEMMA",), + } + } + + RETURN_TYPES = ("CONDITIONING",) + FUNCTION = "encode" + CATEGORY = "ExtraModels/Gemma" + TITLE = "Gemma Text Encode" + + def encode(self, text, GEMMA=None): + print(text) + tokenizer = GEMMA["tokenizer"] + text_encoder = GEMMA["text_encoder"] + + with torch.no_grad(): + tokens = tokenizer( + text, + max_length=300, + padding="max_length", + truncation=True, + return_tensors="pt" + ).to(text_encoder.device) + + cond = text_encoder(tokens.input_ids, tokens.attention_mask)[0][:, None] + emb_masks = tokens.attention_mask + + # 利用emb_masks将有效的cond选出来,其他置零 + # cond = cond * emb_masks.unsqueeze(-1) + + return ([[cond, {}]], ) + +NODE_CLASS_MAPPINGS = { + "GemmaLoader": GemmaLoader, + "GemmaTextEncode": GemmaTextEncode, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "GemmaLoader": "Gemma Loader", + "GemmaTextEncode": "Gemma Text Encode", +} diff --git a/Sana/conf.py b/Sana/conf.py new file mode 100644 index 0000000..7f5a046 --- /dev/null +++ b/Sana/conf.py @@ -0,0 +1,98 @@ +""" +List of all Sana model types / settings +""" + +sampling_settings = { + "shift": 3.0, +} + +sana_conf = { + "SanaMS_600M_P1_D28": { + "target": "SanaMS", + "unet_config": { + "in_channels": 32, + "depth": 28, + "hidden_size": 1152, + "patch_size": 1, + "num_heads": 36, + "linear_head_dim": 32, + "model_max_length": 300, + "y_norm": True, + "attn_type": "linear", + "ffn_type": "glumbconv", + "mlp_ratio": 2.5, + "mlp_acts": ["silu", "silu", None], + "use_pe": False, + "pred_sigma": False, + "learn_sigma": False, + "fp32_attention": True, + }, + "sampling_settings" : sampling_settings, + }, + "SanaMS_1600M_P1_D20": { + "target": "SanaMS", + "unet_config": { + "in_channels": 32, + "depth": 20, + "hidden_size": 2240, + "patch_size": 1, + "num_heads": 70, + "linear_head_dim": 32, + "model_max_length": 300, + "y_norm": True, + "attn_type": "linear", + "ffn_type": "glumbconv", + "mlp_ratio": 2.5, + "mlp_acts": ["silu", "silu", None], + "use_pe": False, + "pred_sigma": False, + "learn_sigma": False, + "fp32_attention": True, + }, + "sampling_settings" : sampling_settings, + }, +} + +sana_res = { + "1024px": { # models/SanaMS 1024x1024 + '0.25': [512, 2048], '0.26': [512, 1984], '0.27': [512, 1920], '0.28': [512, 1856], + '0.32': [576, 1792], '0.33': [576, 1728], '0.35': [576, 1664], '0.40': [640, 1600], + '0.42': [640, 1536], '0.48': [704, 1472], '0.50': [704, 1408], '0.52': [704, 1344], + '0.57': [768, 1344], '0.60': [768, 1280], '0.68': [832, 1216], '0.72': [832, 1152], + '0.78': [896, 1152], '0.82': [896, 1088], '0.88': [960, 1088], '0.94': [960, 1024], + '1.00': [1024,1024], '1.07': [1024, 960], '1.13': [1088, 960], '1.21': [1088, 896], + '1.29': [1152, 896], '1.38': [1152, 832], '1.46': [1216, 832], '1.67': [1280, 768], + '1.75': [1344, 768], '2.00': [1408, 704], '2.09': [1472, 704], '2.40': [1536, 640], + '2.50': [1600, 640], '2.89': [1664, 576], '3.00': [1728, 576], '3.11': [1792, 576], + '3.62': [1856, 512], '3.75': [1920, 512], '3.88': [1984, 512], '4.00': [2048, 512], + }, + "512px": { # models/SanaMS 512x512 + '0.25': [256,1024], '0.26': [256, 992], '0.27': [256, 960], '0.28': [256, 928], + '0.32': [288, 896], '0.33': [288, 864], '0.35': [288, 832], '0.40': [320, 800], + '0.42': [320, 768], '0.48': [352, 736], '0.50': [352, 704], '0.52': [352, 672], + '0.57': [384, 672], '0.60': [384, 640], '0.68': [416, 608], '0.72': [416, 576], + '0.78': [448, 576], '0.82': [448, 544], '0.88': [480, 544], '0.94': [480, 512], + '1.00': [512, 512], '1.07': [512, 480], '1.13': [544, 480], '1.21': [544, 448], + '1.29': [576, 448], '1.38': [576, 416], '1.46': [608, 416], '1.67': [640, 384], + '1.75': [672, 384], '2.00': [704, 352], '2.09': [736, 352], '2.40': [768, 320], + '2.50': [800, 320], '2.89': [832, 288], '3.00': [864, 288], '3.11': [896, 288], + '3.62': [928, 256], '3.75': [960, 256], '3.88': [992, 256], '4.00': [1024,256] + }, + "2K": { + '0.25': [1024, 4096], '0.26': [1024, 3968], '0.27': [1024, 3840], '0.28': [1024, 3712], + '0.32': [1152, 3584], '0.33': [1152, 3456], '0.35': [1152, 3328], '0.40': [1280, 3200], + '0.42': [1280, 3072], '0.48': [1408, 2944], '0.50': [1408, 2816], '0.52': [1408, 2688], + '0.57': [1536, 2688], '0.60': [1536, 2560], '0.68': [1664, 2432], '0.72': [1664, 2304], + '0.78': [1792, 2304], '0.82': [1792, 2176], '0.88': [1920, 2176], '0.94': [1920, 2048], + '1.00': [2048, 2048], '1.07': [2048, 1920], '1.13': [2176, 1920], '1.21': [2176, 1792], + '1.29': [2304, 1792], '1.38': [2304, 1664], '1.46': [2432, 1664], '1.67': [2560, 1536], + '1.75': [2688, 1536], '2.00': [2816, 1408], '2.09': [2944, 1408], '2.40': [3072, 1280], + '2.50': [3200, 1280], '2.89': [3328, 1152], '3.00': [3456, 1152], '3.11': [3584, 1152], + '3.62': [3712, 1024], '3.75': [3840, 1024], '3.88': [3968, 1024], '4.00': [4096, 1024] + } +} +# These should be the same +sana_res.update({ + "SanaMS_600M_P1_D28": sana_res["1024px"], + "SanaMS_1600M_P1_D20": sana_res["1024px"], +}) diff --git a/Sana/diffusers_convert.py b/Sana/diffusers_convert.py new file mode 100644 index 0000000..312ea9d --- /dev/null +++ b/Sana/diffusers_convert.py @@ -0,0 +1,223 @@ +# For using the diffusers format weights +# Based on the original ComfyUI function + +# https://github.com/PixArt-alpha/PixArt-alpha/blob/master/tools/convert_pixart_alpha_to_diffusers.py +import torch + +conversion_map_ms = [ # for multi_scale_train (MS) + # Resolution + ("csize_embedder.mlp.0.weight", "adaln_single.emb.resolution_embedder.linear_1.weight"), + ("csize_embedder.mlp.0.bias", "adaln_single.emb.resolution_embedder.linear_1.bias"), + ("csize_embedder.mlp.2.weight", "adaln_single.emb.resolution_embedder.linear_2.weight"), + ("csize_embedder.mlp.2.bias", "adaln_single.emb.resolution_embedder.linear_2.bias"), + # Aspect ratio + ("ar_embedder.mlp.0.weight", "adaln_single.emb.aspect_ratio_embedder.linear_1.weight"), + ("ar_embedder.mlp.0.bias", "adaln_single.emb.aspect_ratio_embedder.linear_1.bias"), + ("ar_embedder.mlp.2.weight", "adaln_single.emb.aspect_ratio_embedder.linear_2.weight"), + ("ar_embedder.mlp.2.bias", "adaln_single.emb.aspect_ratio_embedder.linear_2.bias"), +] + +def get_depth(state_dict): + return sum(key.endswith('.attn1.to_k.bias') for key in state_dict.keys()) + +def get_lora_depth(state_dict): + cnt = max([ + sum(key.endswith('.attn1.to_k.lora_A.weight') for key in state_dict.keys()), + sum(key.endswith('_attn1_to_k.lora_A.weight') for key in state_dict.keys()), + sum(key.endswith('.attn1.to_k.lora_up.weight') for key in state_dict.keys()), + sum(key.endswith('_attn1_to_k.lora_up.weight') for key in state_dict.keys()), + ]) + assert cnt > 0, "Unable to detect model depth!" + return cnt + +def get_conversion_map(state_dict): + conversion_map = [ # main SD conversion map (PixArt reference, HF Diffusers) + # Patch embeddings + ("x_embedder.proj.weight", "pos_embed.proj.weight"), + ("x_embedder.proj.bias", "pos_embed.proj.bias"), + # Caption projection + ("y_embedder.y_embedding", "caption_projection.y_embedding"), + ("y_embedder.y_proj.fc1.weight", "caption_projection.linear_1.weight"), + ("y_embedder.y_proj.fc1.bias", "caption_projection.linear_1.bias"), + ("y_embedder.y_proj.fc2.weight", "caption_projection.linear_2.weight"), + ("y_embedder.y_proj.fc2.bias", "caption_projection.linear_2.bias"), + # AdaLN-single LN + ("t_embedder.mlp.0.weight", "adaln_single.emb.timestep_embedder.linear_1.weight"), + ("t_embedder.mlp.0.bias", "adaln_single.emb.timestep_embedder.linear_1.bias"), + ("t_embedder.mlp.2.weight", "adaln_single.emb.timestep_embedder.linear_2.weight"), + ("t_embedder.mlp.2.bias", "adaln_single.emb.timestep_embedder.linear_2.bias"), + # Shared norm + ("t_block.1.weight", "adaln_single.linear.weight"), + ("t_block.1.bias", "adaln_single.linear.bias"), + # Final block + ("final_layer.linear.weight", "proj_out.weight"), + ("final_layer.linear.bias", "proj_out.bias"), + ("final_layer.scale_shift_table", "scale_shift_table"), + ] + + # Add actual transformer blocks + for depth in range(get_depth(state_dict)): + # Transformer blocks + conversion_map += [ + (f"blocks.{depth}.scale_shift_table", f"transformer_blocks.{depth}.scale_shift_table"), + # Projection + (f"blocks.{depth}.attn.proj.weight", f"transformer_blocks.{depth}.attn1.to_out.0.weight"), + (f"blocks.{depth}.attn.proj.bias", f"transformer_blocks.{depth}.attn1.to_out.0.bias"), + # Feed-forward + (f"blocks.{depth}.mlp.fc1.weight", f"transformer_blocks.{depth}.ff.net.0.proj.weight"), + (f"blocks.{depth}.mlp.fc1.bias", f"transformer_blocks.{depth}.ff.net.0.proj.bias"), + (f"blocks.{depth}.mlp.fc2.weight", f"transformer_blocks.{depth}.ff.net.2.weight"), + (f"blocks.{depth}.mlp.fc2.bias", f"transformer_blocks.{depth}.ff.net.2.bias"), + # Cross-attention (proj) + (f"blocks.{depth}.cross_attn.proj.weight" ,f"transformer_blocks.{depth}.attn2.to_out.0.weight"), + (f"blocks.{depth}.cross_attn.proj.bias" ,f"transformer_blocks.{depth}.attn2.to_out.0.bias"), + ] + return conversion_map + +def find_prefix(state_dict, target_key): + prefix = "" + for k in state_dict.keys(): + if k.endswith(target_key): + prefix = k.split(target_key)[0] + break + return prefix + +def convert_state_dict(state_dict): + if "adaln_single.emb.resolution_embedder.linear_1.weight" in state_dict.keys(): + cmap = get_conversion_map(state_dict) + conversion_map_ms + else: + cmap = get_conversion_map(state_dict) + + missing = [k for k,v in cmap if v not in state_dict] + new_state_dict = {k: state_dict[v] for k,v in cmap if k not in missing} + matched = list(v for k,v in cmap if v in state_dict.keys()) + + for depth in range(get_depth(state_dict)): + for wb in ["weight", "bias"]: + # Self Attention + key = lambda a: f"transformer_blocks.{depth}.attn1.to_{a}.{wb}" + new_state_dict[f"blocks.{depth}.attn.qkv.{wb}"] = torch.cat(( + state_dict[key('q')], state_dict[key('k')], state_dict[key('v')] + ), dim=0) + matched += [key('q'), key('k'), key('v')] + + # Cross-attention (linear) + key = lambda a: f"transformer_blocks.{depth}.attn2.to_{a}.{wb}" + new_state_dict[f"blocks.{depth}.cross_attn.q_linear.{wb}"] = state_dict[key('q')] + new_state_dict[f"blocks.{depth}.cross_attn.kv_linear.{wb}"] = torch.cat(( + state_dict[key('k')], state_dict[key('v')] + ), dim=0) + matched += [key('q'), key('k'), key('v')] + + if len(matched) < len(state_dict): + print(f"PixArt: UNET conversion has leftover keys! ({len(matched)} vs {len(state_dict)})") + print(list( set(state_dict.keys()) - set(matched) )) + + if len(missing) > 0: + print(f"PixArt: UNET conversion has missing keys!") + print(missing) + + return new_state_dict + +# Same as above but for LoRA weights: +def convert_lora_state_dict(state_dict, peft=True): + # koyha + rep_ak = lambda x: x.replace(".weight", ".lora_down.weight") + rep_bk = lambda x: x.replace(".weight", ".lora_up.weight") + rep_pk = lambda x: x.replace(".weight", ".alpha") + if peft: # peft + rep_ap = lambda x: x.replace(".weight", ".lora_A.weight") + rep_bp = lambda x: x.replace(".weight", ".lora_B.weight") + rep_pp = lambda x: x.replace(".weight", ".alpha") + + prefix = find_prefix(state_dict, "adaln_single.linear.lora_A.weight") + state_dict = {k[len(prefix):]:v for k,v in state_dict.items()} + else: # OneTrainer + rep_ap = lambda x: x.replace(".", "_")[:-7] + ".lora_down.weight" + rep_bp = lambda x: x.replace(".", "_")[:-7] + ".lora_up.weight" + rep_pp = lambda x: x.replace(".", "_")[:-7] + ".alpha" + + prefix = "lora_transformer_" + t5_marker = "lora_te_encoder" + t5_keys = [] + for key in list(state_dict.keys()): + if key.startswith(prefix): + state_dict[key[len(prefix):]] = state_dict.pop(key) + elif t5_marker in key: + t5_keys.append(state_dict.pop(key)) + if len(t5_keys) > 0: + print(f"Text Encoder not supported for PixArt LoRA, ignoring {len(t5_keys)} keys") + + cmap = [] + cmap_unet = get_conversion_map(state_dict) + conversion_map_ms # todo: 512 model + for k, v in cmap_unet: + if v.endswith(".weight"): + cmap.append((rep_ak(k), rep_ap(v))) + cmap.append((rep_bk(k), rep_bp(v))) + if not peft: + cmap.append((rep_pk(k), rep_pp(v))) + + missing = [k for k,v in cmap if v not in state_dict] + new_state_dict = {k: state_dict[v] for k,v in cmap if k not in missing} + matched = list(v for k,v in cmap if v in state_dict.keys()) + + lora_depth = get_lora_depth(state_dict) + for fp, fk in ((rep_ap, rep_ak),(rep_bp, rep_bk)): + for depth in range(lora_depth): + # Self Attention + key = lambda a: fp(f"transformer_blocks.{depth}.attn1.to_{a}.weight") + new_state_dict[fk(f"blocks.{depth}.attn.qkv.weight")] = torch.cat(( + state_dict[key('q')], state_dict[key('k')], state_dict[key('v')] + ), dim=0) + + matched += [key('q'), key('k'), key('v')] + if not peft: + akey = lambda a: rep_pp(f"transformer_blocks.{depth}.attn1.to_{a}.weight") + new_state_dict[rep_pk((f"blocks.{depth}.attn.qkv.weight"))] = state_dict[akey("q")] + matched += [akey('q'), akey('k'), akey('v')] + + # Self Attention projection? + key = lambda a: fp(f"transformer_blocks.{depth}.attn1.to_{a}.weight") + new_state_dict[fk(f"blocks.{depth}.attn.proj.weight")] = state_dict[key('out.0')] + matched += [key('out.0')] + + # Cross-attention (linear) + key = lambda a: fp(f"transformer_blocks.{depth}.attn2.to_{a}.weight") + new_state_dict[fk(f"blocks.{depth}.cross_attn.q_linear.weight")] = state_dict[key('q')] + new_state_dict[fk(f"blocks.{depth}.cross_attn.kv_linear.weight")] = torch.cat(( + state_dict[key('k')], state_dict[key('v')] + ), dim=0) + matched += [key('q'), key('k'), key('v')] + if not peft: + akey = lambda a: rep_pp(f"transformer_blocks.{depth}.attn2.to_{a}.weight") + new_state_dict[rep_pk((f"blocks.{depth}.cross_attn.q_linear.weight"))] = state_dict[akey("q")] + new_state_dict[rep_pk((f"blocks.{depth}.cross_attn.kv_linear.weight"))] = state_dict[akey("k")] + matched += [akey('q'), akey('k'), akey('v')] + + # Cross Attention projection? + key = lambda a: fp(f"transformer_blocks.{depth}.attn2.to_{a}.weight") + new_state_dict[fk(f"blocks.{depth}.cross_attn.proj.weight")] = state_dict[key('out.0')] + matched += [key('out.0')] + + try: + key = fp(f"transformer_blocks.{depth}.ff.net.0.proj.weight") + new_state_dict[fk(f"blocks.{depth}.mlp.fc1.weight")] = state_dict[key] + matched += [key] + except KeyError: + pass + + try: + key = fp(f"transformer_blocks.{depth}.ff.net.2.weight") + new_state_dict[fk(f"blocks.{depth}.mlp.fc2.weight")] = state_dict[key] + matched += [key] + except KeyError: + pass + + if len(matched) < len(state_dict): + print(f"PixArt: LoRA conversion has leftover keys! ({len(matched)} vs {len(state_dict)})") + print(list( set(state_dict.keys()) - set(matched) )) + + if len(missing) > 0: + print(f"PixArt: LoRA conversion has missing keys! (probably)") + print(missing) + + return new_state_dict diff --git a/Sana/loader.py b/Sana/loader.py new file mode 100644 index 0000000..34ba85f --- /dev/null +++ b/Sana/loader.py @@ -0,0 +1,100 @@ +import comfy.supported_models_base +import comfy.latent_formats +import comfy.model_patcher +import comfy.model_base +import comfy.utils +import comfy.conds +import torch +import math +from comfy import model_management +from comfy.latent_formats import LatentFormat +from .diffusers_convert import convert_state_dict + + +class SanaLatent(LatentFormat): + latent_channels = 32 + def __init__(self): + self.scale_factor = 0.41407 + + +class EXM_Sana(comfy.supported_models_base.BASE): + unet_config = {} + unet_extra_config = {} + latent_format = SanaLatent + + def __init__(self, model_conf): + self.model_target = model_conf.get("target") + self.unet_config = model_conf.get("unet_config", {}) + self.sampling_settings = model_conf.get("sampling_settings", {}) + self.latent_format = self.latent_format() + # UNET is handled by extension + self.unet_config["disable_unet_model_creation"] = True + + def model_type(self, state_dict, prefix=""): + return comfy.model_base.ModelType.FLOW + + +class EXM_Sana_Model(comfy.model_base.BaseModel): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + def extra_conds(self, **kwargs): + out = super().extra_conds(**kwargs) + + cn_hint = kwargs.get("cn_hint", None) + if cn_hint is not None: + out["cn_hint"] = comfy.conds.CONDRegular(cn_hint) + + return out + + +def load_sana(model_path, model_conf, dtype): + state_dict = comfy.utils.load_torch_file(model_path) + state_dict = state_dict.get("model", state_dict) + + # prefix + for prefix in ["model.diffusion_model.",]: + if any(True for x in state_dict if x.startswith(prefix)): + state_dict = {k[len(prefix):]:v for k,v in state_dict.items()} + + # diffusers + if "adaln_single.linear.weight" in state_dict: + state_dict = convert_state_dict(state_dict) # Diffusers + + parameters = comfy.utils.calculate_parameters(state_dict) + unet_dtype = dtype + load_device = comfy.model_management.get_torch_device() + offload_device = comfy.model_management.unet_offload_device() + + # ignore fp8/etc and use directly for now + manual_cast_dtype = model_management.unet_manual_cast(unet_dtype, load_device) + if manual_cast_dtype: + print(f"Sana: falling back to {manual_cast_dtype}") + unet_dtype = manual_cast_dtype + + model_conf = EXM_Sana(model_conf) # convert to object + model = EXM_Sana_Model( # same as comfy.model_base.BaseModel + model_conf, + model_type=comfy.model_base.ModelType.FLOW, + device=model_management.get_torch_device() + ) + + if model_conf.model_target == "SanaMS": + from .models.sana_multi_scale import SanaMS + model.diffusion_model = SanaMS(**model_conf.unet_config) + else: + raise NotImplementedError(f"Unknown model target '{model_conf.model_target}'") + + m, u = model.diffusion_model.load_state_dict(state_dict, strict=False) + if len(m) > 0: print("Missing UNET keys", m) + if len(u) > 0: print("Leftover UNET keys", u) + model.diffusion_model.dtype = unet_dtype + model.diffusion_model.eval() + model.diffusion_model.to(unet_dtype) + + model_patcher = comfy.model_patcher.ModelPatcher( + model, + load_device = load_device, + offload_device = offload_device, + ) + return model_patcher diff --git a/Sana/lora.py b/Sana/lora.py new file mode 100644 index 0000000..fca5931 --- /dev/null +++ b/Sana/lora.py @@ -0,0 +1,146 @@ +import os +import copy +import json +import torch +import comfy.lora +import comfy.model_management +from comfy.model_patcher import ModelPatcher +from .diffusers_convert import convert_lora_state_dict + +class EXM_PixArt_ModelPatcher(ModelPatcher): + def calculate_weight(self, patches, weight, key): + """ + This is almost the same as the comfy function, but stripped down to just the LoRA patch code. + The problem with the original code is the q/k/v keys being combined into one for the attention. + In the diffusers code, they're treated as separate keys, but in the reference code they're recombined (q+kv|qkv). + This means, for example, that the [1152,1152] weights become [3456,1152] in the state dict. + The issue with this is that the LoRA weights are [128,1152],[1152,128] and become [384,1162],[3456,128] instead. + + This is the best thing I could think of that would fix that, but it's very fragile. + - Check key shape to determine if it needs the fallback logic + - Cut the input into parts based on the shape (undoing the torch.cat) + - Do the matrix multiplication logic + - Recombine them to match the expected shape + """ + for p in patches: + alpha = p[0] + v = p[1] + strength_model = p[2] + if strength_model != 1.0: + weight *= strength_model + + if isinstance(v, list): + v = (self.calculate_weight(v[1:], v[0].clone(), key), ) + + if len(v) == 2: + patch_type = v[0] + v = v[1] + + if patch_type == "lora": + mat1 = comfy.model_management.cast_to_device(v[0], weight.device, torch.float32) + mat2 = comfy.model_management.cast_to_device(v[1], weight.device, torch.float32) + if v[2] is not None: + alpha *= v[2] / mat2.shape[0] + try: + mat1 = mat1.flatten(start_dim=1) + mat2 = mat2.flatten(start_dim=1) + + ch1 = mat1.shape[0] // mat2.shape[1] + ch2 = mat2.shape[0] // mat1.shape[1] + ### Fallback logic for shape mismatch ### + if mat1.shape[0] != mat2.shape[1] and ch1 == ch2 and (mat1.shape[0]/mat2.shape[1])%1 == 0: + mat1 = mat1.chunk(ch1, dim=0) + mat2 = mat2.chunk(ch1, dim=0) + weight += torch.cat( + [alpha * torch.mm(mat1[x], mat2[x]) for x in range(ch1)], + dim=0, + ).reshape(weight.shape).type(weight.dtype) + else: + weight += (alpha * torch.mm(mat1, mat2)).reshape(weight.shape).type(weight.dtype) + except Exception as e: + print("ERROR", key, e) + return weight + + def clone(self): + n = EXM_PixArt_ModelPatcher(self.model, self.load_device, self.offload_device, self.size, self.current_device, weight_inplace_update=self.weight_inplace_update) + n.patches = {} + for k in self.patches: + n.patches[k] = self.patches[k][:] + + n.object_patches = self.object_patches.copy() + n.model_options = copy.deepcopy(self.model_options) + n.model_keys = self.model_keys + return n + +def replace_model_patcher(model): + n = EXM_PixArt_ModelPatcher( + model = model.model, + size = model.size, + load_device = model.load_device, + offload_device = model.offload_device, + weight_inplace_update = model.weight_inplace_update, + ) + n.patches = {} + for k in model.patches: + n.patches[k] = model.patches[k][:] + + n.object_patches = model.object_patches.copy() + n.model_options = copy.deepcopy(model.model_options) + return n + +def find_peft_alpha(path): + def load_json(json_path): + with open(json_path) as f: + data = json.load(f) + alpha = data.get("lora_alpha") + alpha = alpha or data.get("alpha") + if not alpha: + print(" Found config but `lora_alpha` is missing!") + else: + print(f" Found config at {json_path} [alpha:{alpha}]") + return alpha + + # For some weird reason peft doesn't include the alpha in the actual model + print("PixArt: Warning! This is a PEFT LoRA. Trying to find config...") + files = [ + f"{os.path.splitext(path)[0]}.json", + f"{os.path.splitext(path)[0]}.config.json", + os.path.join(os.path.dirname(path),"adapter_config.json"), + ] + for file in files: + if os.path.isfile(file): + return load_json(file) + + print(" Missing config/alpha! assuming alpha of 8. Consider converting it/adding a config json to it.") + return 8.0 + +def load_pixart_lora(model, lora, lora_path, strength): + k_back = lambda x: x.replace(".lora_up.weight", "") + # need to convert the actual weights for this to work. + if any(True for x in lora.keys() if x.endswith("adaln_single.linear.lora_A.weight")): + lora = convert_lora_state_dict(lora, peft=True) + alpha = find_peft_alpha(lora_path) + lora.update({f"{k_back(x)}.alpha":torch.tensor(alpha) for x in lora.keys() if "lora_up" in x}) + else: # OneTrainer + lora = convert_lora_state_dict(lora, peft=False) + + key_map = {k_back(x):f"diffusion_model.{k_back(x)}.weight" for x in lora.keys() if "lora_up" in x} # fake + + loaded = comfy.lora.load_lora(lora, key_map) + if model is not None: + # switch to custom model patcher when using LoRAs + if isinstance(model, EXM_PixArt_ModelPatcher): + new_modelpatcher = model.clone() + else: + new_modelpatcher = replace_model_patcher(model) + k = new_modelpatcher.add_patches(loaded, strength) + else: + k = () + new_modelpatcher = None + + k = set(k) + for x in loaded: + if (x not in k): + print("NOT LOADED", x) + + return new_modelpatcher diff --git a/Sana/models/act.py b/Sana/models/act.py new file mode 100644 index 0000000..9df6a7a --- /dev/null +++ b/Sana/models/act.py @@ -0,0 +1,59 @@ +# Copyright 2024 NVIDIA CORPORATION & AFFILIATES +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# SPDX-License-Identifier: Apache-2.0 + +import copy + +import torch.nn as nn + +__all__ = ["build_act", "get_act_name"] + +# register activation function here +# name: module, kwargs with default values +REGISTERED_ACT_DICT: dict[str, tuple[type, dict[str, any]]] = { + "relu": (nn.ReLU, {"inplace": True}), + "relu6": (nn.ReLU6, {"inplace": True}), + "hswish": (nn.Hardswish, {"inplace": True}), + "hsigmoid": (nn.Hardsigmoid, {"inplace": True}), + "swish": (nn.SiLU, {"inplace": True}), + "silu": (nn.SiLU, {"inplace": True}), + "tanh": (nn.Tanh, {}), + "sigmoid": (nn.Sigmoid, {}), + "gelu": (nn.GELU, {"approximate": "tanh"}), + "mish": (nn.Mish, {"inplace": True}), + "identity": (nn.Identity, {}), +} + + +def build_act(name: str or None, **kwargs) -> nn.Module or None: + if name in REGISTERED_ACT_DICT: + act_cls, default_args = copy.deepcopy(REGISTERED_ACT_DICT[name]) + for key in default_args: + if key in kwargs: + default_args[key] = kwargs[key] + return act_cls(**default_args) + elif name is None or name.lower() == "none": + return None + else: + raise ValueError(f"do not support: {name}") + + +def get_act_name(act: nn.Module or None) -> str or None: + if act is None: + return None + module2name = {} + for key, config in REGISTERED_ACT_DICT.items(): + module2name[config[0].__name__] = key + return module2name.get(type(act).__name__, "unknown") diff --git a/Sana/models/basic_modules.py b/Sana/models/basic_modules.py new file mode 100644 index 0000000..ece579a --- /dev/null +++ b/Sana/models/basic_modules.py @@ -0,0 +1,361 @@ +# Copyright 2024 NVIDIA CORPORATION & AFFILIATES +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# SPDX-License-Identifier: Apache-2.0 + +# This file is modified from https://github.com/PixArt-alpha/PixArt-sigma +import torch +import torch.nn as nn +from timm.models.vision_transformer import Mlp + +from .act import build_act, get_act_name +from .norms import build_norm, get_norm_name +from .utils import get_same_padding, val2tuple + + +class ConvLayer(nn.Module): + def __init__( + self, + in_dim: int, + out_dim: int, + kernel_size=3, + stride=1, + dilation=1, + groups=1, + padding: int or None = None, + use_bias=False, + dropout=0.0, + norm="bn2d", + act="relu", + ): + super().__init__() + if padding is None: + padding = get_same_padding(kernel_size) + padding *= dilation + + self.in_dim = in_dim + self.out_dim = out_dim + self.kernel_size = kernel_size + self.stride = stride + self.dilation = dilation + self.groups = groups + self.padding = padding + self.use_bias = use_bias + + self.dropout = nn.Dropout2d(dropout, inplace=False) if dropout > 0 else None + self.conv = nn.Conv2d( + in_dim, + out_dim, + kernel_size=(kernel_size, kernel_size), + stride=(stride, stride), + padding=padding, + dilation=(dilation, dilation), + groups=groups, + bias=use_bias, + ) + self.norm = build_norm(norm, num_features=out_dim) + self.act = build_act(act) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if self.dropout is not None: + x = self.dropout(x) + x = self.conv(x) + if self.norm: + x = self.norm(x) + if self.act: + x = self.act(x) + return x + + +class GLUMBConv(nn.Module): + def __init__( + self, + in_features: int, + hidden_features: int, + out_feature=None, + kernel_size=3, + stride=1, + padding: int or None = None, + use_bias=False, + norm=(None, None, None), + act=("silu", "silu", None), + dilation=1, + ): + out_feature = out_feature or in_features + super().__init__() + use_bias = val2tuple(use_bias, 3) + norm = val2tuple(norm, 3) + act = val2tuple(act, 3) + + self.glu_act = build_act(act[1], inplace=False) + self.inverted_conv = ConvLayer( + in_features, + hidden_features * 2, + 1, + use_bias=use_bias[0], + norm=norm[0], + act=act[0], + ) + self.depth_conv = ConvLayer( + hidden_features * 2, + hidden_features * 2, + kernel_size, + stride=stride, + groups=hidden_features * 2, + padding=padding, + use_bias=use_bias[1], + norm=norm[1], + act=None, + dilation=dilation, + ) + self.point_conv = ConvLayer( + hidden_features, + out_feature, + 1, + use_bias=use_bias[2], + norm=norm[2], + act=act[2], + ) + # from IPython import embed; embed(header='debug dilate conv') + + def forward(self, x: torch.Tensor, HW=None) -> torch.Tensor: + B, N, C = x.shape + if HW is None: + H = W = int(N**0.5) + else: + H, W = HW + + x = x.reshape(B, H, W, C).permute(0, 3, 1, 2) + x = self.inverted_conv(x) + x = self.depth_conv(x) + + x, gate = torch.chunk(x, 2, dim=1) + gate = self.glu_act(gate) + x = x * gate + + x = self.point_conv(x) + x = x.reshape(B, C, N).permute(0, 2, 1) + + return x + + +class SlimGLUMBConv(GLUMBConv): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + # 移除 self.inverted_conv 层 + del self.inverted_conv + self.out_dim = self.point_conv.out_dim + + def forward(self, x: torch.Tensor, HW=None) -> torch.Tensor: + B, N, C = x.shape + if HW is None: + H = W = int(N**0.5) + else: + H, W = HW + + # 直接使用 x,跳过 self.inverted_conv 层的调用 + x = x.reshape(B, H, W, C).permute(0, 3, 1, 2) + # x = self.inverted_conv(x) + x = self.depth_conv(x) + + x, gate = torch.chunk(x, 2, dim=1) + gate = self.glu_act(gate) + x = x * gate + + x = self.point_conv(x) + x = x.reshape(B, self.out_dim, N).permute(0, 2, 1) + + return x + + +class MBConvPreGLU(nn.Module): + def __init__( + self, + in_dim: int, + out_dim: int, + kernel_size=3, + stride=1, + mid_dim=None, + expand=6, + padding: int or None = None, + use_bias=False, + norm=(None, None, "ln2d"), + act=("silu", "silu", None), + ): + super().__init__() + use_bias = val2tuple(use_bias, 3) + norm = val2tuple(norm, 3) + act = val2tuple(act, 3) + + mid_dim = mid_dim or round(in_dim * expand) + + self.inverted_conv = ConvLayer( + in_dim, + mid_dim * 2, + 1, + use_bias=use_bias[0], + norm=norm[0], + act=None, + ) + self.glu_act = build_act(act[0], inplace=False) + self.depth_conv = ConvLayer( + mid_dim, + mid_dim, + kernel_size, + stride=stride, + groups=mid_dim, + padding=padding, + use_bias=use_bias[1], + norm=norm[1], + act=act[1], + ) + self.point_conv = ConvLayer( + mid_dim, + out_dim, + 1, + use_bias=use_bias[2], + norm=norm[2], + act=act[2], + ) + + def forward(self, x: torch.Tensor, HW=None) -> torch.Tensor: + B, N, C = x.shape + if HW is None: + H = W = int(N**0.5) + else: + H, W = HW + + x = x.reshape(B, H, W, C).permute(0, 3, 1, 2) + + x = self.inverted_conv(x) + x, gate = torch.chunk(x, 2, dim=1) + gate = self.glu_act(gate) + x = x * gate + + x = self.depth_conv(x) + x = self.point_conv(x) + + x = x.reshape(B, C, N).permute(0, 2, 1) + return x + + @property + def module_str(self) -> str: + _str = f"{self.depth_conv.kernel_size}{type(self).__name__}(" + _str += f"in={self.inverted_conv.in_dim},mid={self.depth_conv.in_dim},out={self.point_conv.out_dim},s={self.depth_conv.stride}" + _str += ( + f",norm={get_norm_name(self.inverted_conv.norm)}" + f"+{get_norm_name(self.depth_conv.norm)}" + f"+{get_norm_name(self.point_conv.norm)}" + ) + _str += ( + f",act={get_act_name(self.inverted_conv.act)}" + f"+{get_act_name(self.depth_conv.act)}" + f"+{get_act_name(self.point_conv.act)}" + ) + _str += f",glu_act={get_act_name(self.glu_act)})" + return _str + + +class DWMlp(Mlp): + """MLP as used in Vision Transformer, MLP-Mixer and related networks""" + + def __init__( + self, + in_features, + hidden_features=None, + out_features=None, + act_layer=nn.GELU, + bias=True, + drop=0.0, + kernel_size=3, + stride=1, + dilation=1, + padding=None, + ): + super().__init__( + in_features=in_features, + hidden_features=hidden_features, + out_features=out_features, + act_layer=act_layer, + bias=bias, + drop=drop, + ) + hidden_features = hidden_features or in_features + self.hidden_features = hidden_features + if padding is None: + padding = get_same_padding(kernel_size) + padding *= dilation + + self.conv = nn.Conv2d( + hidden_features, + hidden_features, + kernel_size=(kernel_size, kernel_size), + stride=(stride, stride), + padding=padding, + dilation=(dilation, dilation), + groups=hidden_features, + bias=bias, + ) + + def forward(self, x, HW=None): + B, N, C = x.shape + if HW is None: + H = W = int(N**0.5) + else: + H, W = HW + x = self.fc1(x) + x = self.act(x) + x = self.drop1(x) + x = x.reshape(B, H, W, self.hidden_features).permute(0, 3, 1, 2) + x = self.conv(x) + x = x.reshape(B, self.hidden_features, N).permute(0, 2, 1) + x = self.fc2(x) + x = self.drop2(x) + return x + + +class Mlp(Mlp): + """MLP as used in Vision Transformer, MLP-Mixer and related networks""" + + def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, bias=True, drop=0.0): + super().__init__( + in_features=in_features, + hidden_features=hidden_features, + out_features=out_features, + act_layer=act_layer, + bias=bias, + drop=drop, + ) + + def forward(self, x, HW=None): + x = self.fc1(x) + x = self.act(x) + x = self.drop1(x) + x = self.fc2(x) + x = self.drop2(x) + return x + + +if __name__ == "__main__": + model = GLUMBConv( + 1152, + 1152 * 4, + 1152, + use_bias=(True, True, False), + norm=(None, None, None), + act=("silu", "silu", None), + ).cuda() + input = torch.randn(4, 256, 1152).cuda() + output = model(input) diff --git a/Sana/models/norms.py b/Sana/models/norms.py new file mode 100644 index 0000000..4731d69 --- /dev/null +++ b/Sana/models/norms.py @@ -0,0 +1,225 @@ +# Copyright 2024 NVIDIA CORPORATION & AFFILIATES +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# SPDX-License-Identifier: Apache-2.0 + +import copy +import warnings + +import torch +import torch.nn as nn +from torch.nn.modules.batchnorm import _BatchNorm + +__all__ = ["LayerNorm2d", "build_norm", "get_norm_name", "reset_bn", "remove_bn", "set_norm_eps"] + + +class LayerNorm2d(nn.LayerNorm): + rmsnorm = False + + def forward(self, x: torch.Tensor) -> torch.Tensor: + out = x if LayerNorm2d.rmsnorm else x - torch.mean(x, dim=1, keepdim=True) + out = out / torch.sqrt(torch.square(out).mean(dim=1, keepdim=True) + self.eps) + if self.elementwise_affine: + out = out * self.weight.view(1, -1, 1, 1) + self.bias.view(1, -1, 1, 1) + return out + + def extra_repr(self) -> str: + return f"{self.normalized_shape}, eps={self.eps}, elementwise_affine={self.elementwise_affine}, rmsnorm={self.rmsnorm}" + + +# register normalization function here +# name: module, kwargs with default values +REGISTERED_NORMALIZATION_DICT: dict[str, tuple[type, dict[str, any]]] = { + "bn2d": (nn.BatchNorm2d, {"num_features": None, "eps": 1e-5, "momentum": 0.1, "affine": True}), + "syncbn": (nn.SyncBatchNorm, {"num_features": None, "eps": 1e-5, "momentum": 0.1, "affine": True}), + "ln": (nn.LayerNorm, {"normalized_shape": None, "eps": 1e-5, "elementwise_affine": True}), + "ln2d": (LayerNorm2d, {"normalized_shape": None, "eps": 1e-5, "elementwise_affine": True}), +} + + +def build_norm(name="bn2d", num_features=None, affine=True, **kwargs) -> nn.Module or None: + if name in ["ln", "ln2d"]: + kwargs["normalized_shape"] = num_features + kwargs["elementwise_affine"] = affine + else: + kwargs["num_features"] = num_features + kwargs["affine"] = affine + if name in REGISTERED_NORMALIZATION_DICT: + norm_cls, default_args = copy.deepcopy(REGISTERED_NORMALIZATION_DICT[name]) + for key in default_args: + if key in kwargs: + default_args[key] = kwargs[key] + return norm_cls(**default_args) + elif name is None or name.lower() == "none": + return None + else: + raise ValueError("do not support: %s" % name) + + +def get_norm_name(norm: nn.Module or None) -> str or None: + if norm is None: + return None + module2name = {} + for key, config in REGISTERED_NORMALIZATION_DICT.items(): + module2name[config[0].__name__] = key + return module2name.get(type(norm).__name__, "unknown") + + +def reset_bn( + model: nn.Module, + data_loader: list, + sync=True, + progress_bar=False, +) -> None: + import copy + + import torch.nn.functional as F + from packages.apps.utils import AverageMeter, is_master, sync_tensor + from packages.models.utils import get_device, list_join + from tqdm import tqdm + + bn_mean = {} + bn_var = {} + + tmp_model = copy.deepcopy(model) + for name, m in tmp_model.named_modules(): + if isinstance(m, _BatchNorm): + bn_mean[name] = AverageMeter(is_distributed=False) + bn_var[name] = AverageMeter(is_distributed=False) + + def new_forward(bn, mean_est, var_est): + def lambda_forward(x): + x = x.contiguous() + if sync: + batch_mean = x.mean(0, keepdim=True).mean(2, keepdim=True).mean(3, keepdim=True) # 1, C, 1, 1 + batch_mean = sync_tensor(batch_mean, reduce="cat") + batch_mean = torch.mean(batch_mean, dim=0, keepdim=True) + + batch_var = (x - batch_mean) * (x - batch_mean) + batch_var = batch_var.mean(0, keepdim=True).mean(2, keepdim=True).mean(3, keepdim=True) + batch_var = sync_tensor(batch_var, reduce="cat") + batch_var = torch.mean(batch_var, dim=0, keepdim=True) + else: + batch_mean = x.mean(0, keepdim=True).mean(2, keepdim=True).mean(3, keepdim=True) # 1, C, 1, 1 + batch_var = (x - batch_mean) * (x - batch_mean) + batch_var = batch_var.mean(0, keepdim=True).mean(2, keepdim=True).mean(3, keepdim=True) + + batch_mean = torch.squeeze(batch_mean) + batch_var = torch.squeeze(batch_var) + + mean_est.update(batch_mean.data, x.size(0)) + var_est.update(batch_var.data, x.size(0)) + + # bn forward using calculated mean & var + _feature_dim = batch_mean.shape[0] + return F.batch_norm( + x, + batch_mean, + batch_var, + bn.weight[:_feature_dim], + bn.bias[:_feature_dim], + False, + 0.0, + bn.eps, + ) + + return lambda_forward + + m.forward = new_forward(m, bn_mean[name], bn_var[name]) + + # skip if there is no batch normalization layers in the network + if len(bn_mean) == 0: + return + + tmp_model.eval() + with torch.inference_mode(): + with tqdm(total=len(data_loader), desc="reset bn", disable=not progress_bar or not is_master()) as t: + for images in data_loader: + images = images.to(get_device(tmp_model)) + tmp_model(images) + t.set_postfix( + { + "bs": images.size(0), + "res": list_join(images.shape[-2:], "x"), + } + ) + t.update() + + for name, m in model.named_modules(): + if name in bn_mean and bn_mean[name].count > 0: + feature_dim = bn_mean[name].avg.size(0) + assert isinstance(m, _BatchNorm) + m.running_mean.data[:feature_dim].copy_(bn_mean[name].avg) + m.running_var.data[:feature_dim].copy_(bn_var[name].avg) + + +def remove_bn(model: nn.Module) -> None: + for m in model.modules(): + if isinstance(m, _BatchNorm): + m.weight = m.bias = None + m.forward = lambda x: x + + +def set_norm_eps(model: nn.Module, eps: float or None = None, momentum: float or None = None) -> None: + for m in model.modules(): + if isinstance(m, (nn.GroupNorm, nn.LayerNorm, _BatchNorm)): + if eps is not None: + m.eps = eps + if momentum is not None: + m.momentum = momentum + + +class RMSNorm(torch.nn.Module): + def __init__(self, dim: int, scale_factor=1.0, eps: float = 1e-6): + """ + Initialize the RMSNorm normalization layer. + + Args: + dim (int): The dimension of the input tensor. + eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6. + + Attributes: + eps (float): A small value added to the denominator for numerical stability. + weight (nn.Parameter): Learnable scaling parameter. + + """ + super().__init__() + self.eps = eps + self.weight = nn.Parameter(torch.ones(dim) * scale_factor) + + def _norm(self, x): + """ + Apply the RMSNorm normalization to the input tensor. + + Args: + x (torch.Tensor): The input tensor. + + Returns: + torch.Tensor: The normalized tensor. + + """ + return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) + + def forward(self, x): + """ + Forward pass through the RMSNorm layer. + + Args: + x (torch.Tensor): The input tensor. + + Returns: + torch.Tensor: The output tensor after applying RMSNorm. + + """ + return (self.weight * self._norm(x.float())).type_as(x) diff --git a/Sana/models/sana.py b/Sana/models/sana.py new file mode 100644 index 0000000..0dd6551 --- /dev/null +++ b/Sana/models/sana.py @@ -0,0 +1,379 @@ +# Copyright 2024 NVIDIA CORPORATION & AFFILIATES +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# SPDX-License-Identifier: Apache-2.0 + +# This file is modified from https://github.com/PixArt-alpha/PixArt-sigma +import os + +import numpy as np +import torch +import torch.nn as nn +from timm.models.layers import DropPath + +from .basic_modules import DWMlp, GLUMBConv, MBConvPreGLU, Mlp +from .sana_blocks import ( + Attention, + CaptionEmbedder, + FlashAttention, + LiteLA, + MultiHeadCrossAttention, + PatchEmbed, + T2IFinalLayer, + TimestepEmbedder, + t2i_modulate, +) +from .norms import RMSNorm +from .utils import auto_grad_checkpoint, to_2tuple + + +class SanaBlock(nn.Module): + """ + A Sana block with global shared adaptive layer norm (adaLN-single) conditioning. + """ + + def __init__( + self, + hidden_size, + num_heads, + mlp_ratio=4.0, + drop_path=0, + input_size=None, + qk_norm=False, + attn_type="flash", + ffn_type="mlp", + mlp_acts=("silu", "silu", None), + linear_head_dim=32, + **block_kwargs, + ): + super().__init__() + self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + if attn_type == "flash": + # flash self attention + self.attn = FlashAttention( + hidden_size, + num_heads=num_heads, + qkv_bias=True, + qk_norm=qk_norm, + **block_kwargs, + ) + elif attn_type == "linear": + # linear self attention + # TODO: Here the num_heads set to 36 for tmp used + self_num_heads = hidden_size // linear_head_dim + self.attn = LiteLA(hidden_size, hidden_size, heads=self_num_heads, eps=1e-8, qk_norm=qk_norm) + elif attn_type == "vanilla": + # vanilla self attention + self.attn = Attention(hidden_size, num_heads=num_heads, qkv_bias=True) + else: + raise ValueError(f"{attn_type} type is not defined.") + + self.cross_attn = MultiHeadCrossAttention(hidden_size, num_heads, **block_kwargs) + self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + # to be compatible with lower version pytorch + if ffn_type == "dwmlp": + approx_gelu = lambda: nn.GELU(approximate="tanh") + self.mlp = DWMlp( + in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu, drop=0 + ) + elif ffn_type == "glumbconv": + self.mlp = GLUMBConv( + in_features=hidden_size, + hidden_features=int(hidden_size * mlp_ratio), + use_bias=(True, True, False), + norm=(None, None, None), + act=mlp_acts, + ) + elif ffn_type == "glumbconv_dilate": + self.mlp = GLUMBConv( + in_features=hidden_size, + hidden_features=int(hidden_size * mlp_ratio), + use_bias=(True, True, False), + norm=(None, None, None), + act=mlp_acts, + dilation=2, + ) + elif ffn_type == "mbconvpreglu": + self.mlp = MBConvPreGLU( + in_dim=hidden_size, + out_dim=hidden_size, + mid_dim=int(hidden_size * mlp_ratio), + use_bias=(True, True, False), + norm=None, + act=("silu", "silu", None), + ) + elif ffn_type == "mlp": + approx_gelu = lambda: nn.GELU(approximate="tanh") + self.mlp = Mlp( + in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu, drop=0 + ) + else: + raise ValueError(f"{ffn_type} type is not defined.") + self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() + self.scale_shift_table = nn.Parameter(torch.randn(6, hidden_size) / hidden_size**0.5) + + def forward(self, x, y, t, mask=None, **kwargs): + B, N, C = x.shape + + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( + self.scale_shift_table[None] + t.reshape(B, 6, -1) + ).chunk(6, dim=1) + x = x + self.drop_path(gate_msa * self.attn(t2i_modulate(self.norm1(x), shift_msa, scale_msa)).reshape(B, N, C)) + x = x + self.cross_attn(x, y, mask) + x = x + self.drop_path(gate_mlp * self.mlp(t2i_modulate(self.norm2(x), shift_mlp, scale_mlp))) + + return x + + +############################################################################# +# Core Sana Model # +################################################################################# +class Sana(nn.Module): + """ + Diffusion model with a Transformer backbone. + """ + + def __init__( + self, + input_size=32, + patch_size=1, + in_channels=32, + hidden_size=1152, + depth=28, + num_heads=36, + mlp_ratio=2.5, + class_dropout_prob=0.1, + pred_sigma=False, + drop_path: float = 0.0, + caption_channels=2304, + pe_interpolation=1.0, + config=None, + model_max_length=120, + qk_norm=False, + y_norm=False, + norm_eps=1e-5, + attn_type="flash", + ffn_type="mlp", + use_pe=False, + y_norm_scale_factor=1.0, + patch_embed_kernel=None, + mlp_acts=("silu", "silu", None), + linear_head_dim=32, + **kwargs, + ): + super().__init__() + self.pred_sigma = pred_sigma + self.in_channels = in_channels + self.out_channels = in_channels * 2 if pred_sigma else in_channels + self.patch_size = patch_size + self.num_heads = num_heads + self.pe_interpolation = pe_interpolation + self.depth = depth + self.use_pe = use_pe + self.y_norm = y_norm + self.fp32_attention = kwargs.get("use_fp32_attention", False) + + kernel_size = patch_embed_kernel or patch_size + self.x_embedder = PatchEmbed( + input_size, patch_size, in_channels, hidden_size, kernel_size=kernel_size, bias=True + ) + self.t_embedder = TimestepEmbedder(hidden_size) + num_patches = self.x_embedder.num_patches + self.base_size = input_size // self.patch_size + # Will use fixed sin-cos embedding: + self.register_buffer("pos_embed", torch.zeros(1, num_patches, hidden_size)) + + approx_gelu = lambda: nn.GELU(approximate="tanh") + self.t_block = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size, bias=True)) + self.y_embedder = CaptionEmbedder( + in_channels=caption_channels, + hidden_size=hidden_size, + uncond_prob=class_dropout_prob, + act_layer=approx_gelu, + token_num=model_max_length, + ) + if self.y_norm: + self.attention_y_norm = RMSNorm(hidden_size, scale_factor=y_norm_scale_factor, eps=norm_eps) + drop_path = [x.item() for x in torch.linspace(0, drop_path, depth)] # stochastic depth decay rule + self.blocks = nn.ModuleList( + [ + SanaBlock( + hidden_size, + num_heads, + mlp_ratio=mlp_ratio, + drop_path=drop_path[i], + input_size=(input_size // patch_size, input_size // patch_size), + qk_norm=qk_norm, + attn_type=attn_type, + ffn_type=ffn_type, + mlp_acts=mlp_acts, + linear_head_dim=linear_head_dim, + ) + for i in range(depth) + ] + ) + self.final_layer = T2IFinalLayer(hidden_size, patch_size, self.out_channels) + + self.initialize_weights() + + def forward(self, x, timestep, y, mask=None, data_info=None, **kwargs): + """ + Forward pass of Sana. + x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images) + t: (N,) tensor of diffusion timesteps + y: (N, 1, 120, C) tensor of class labels + """ + x = x.to(self.dtype) + timestep = timestep.to(self.dtype) + y = y.to(self.dtype) + pos_embed = self.pos_embed.to(self.dtype) + self.h, self.w = x.shape[-2] // self.patch_size, x.shape[-1] // self.patch_size + if self.use_pe: + x = self.x_embedder(x) + pos_embed # (N, T, D), where T = H * W / patch_size ** 2 + else: + x = self.x_embedder(x) + t = self.t_embedder(timestep.to(x.dtype)) # (N, D) + t0 = self.t_block(t) + y = self.y_embedder(y, self.training) # (N, 1, L, D) + if self.y_norm: + y = self.attention_y_norm(y) + if mask is not None: + if mask.shape[0] != y.shape[0]: + mask = mask.repeat(y.shape[0] // mask.shape[0], 1) + mask = mask.squeeze(1).squeeze(1) + y = y.squeeze(1).masked_select(mask.unsqueeze(-1) != 0).view(1, -1, x.shape[-1]) + y_lens = mask.sum(dim=1).tolist() + else: + y_lens = [y.shape[2]] * y.shape[0] + y = y.squeeze(1).view(1, -1, x.shape[-1]) + for block in self.blocks: + x = auto_grad_checkpoint(block, x, y, t0, y_lens) # (N, T, D) #support grad checkpoint + x = self.final_layer(x, t) # (N, T, patch_size ** 2 * out_channels) + x = self.unpatchify(x) # (N, out_channels, H, W) + return x + + def __call__(self, *args, **kwargs): + """ + This method allows the object to be called like a function. + It simply calls the forward method. + """ + return self.forward(*args, **kwargs) + + def forward_with_dpmsolver(self, x, timestep, y, mask=None, **kwargs): + """ + dpm solver donnot need variance prediction + """ + # https://github.com/openai/glide-text2im/blob/main/notebooks/text2im.ipynb + model_out = self.forward(x, timestep, y, mask) + return model_out.chunk(2, dim=1)[0] if self.pred_sigma else model_out + + def unpatchify(self, x): + """ + x: (N, T, patch_size**2 * C) + imgs: (N, H, W, C) + """ + c = self.out_channels + p = self.x_embedder.patch_size[0] + h = w = int(x.shape[1] ** 0.5) + assert h * w == x.shape[1] + + x = x.reshape(shape=(x.shape[0], h, w, p, p, c)) + x = torch.einsum("nhwpqc->nchpwq", x) + imgs = x.reshape(shape=(x.shape[0], c, h * p, h * p)) + return imgs + + def initialize_weights(self): + # Initialize transformer layers: + def _basic_init(module): + if isinstance(module, nn.Linear): + torch.nn.init.xavier_uniform_(module.weight) + if module.bias is not None: + nn.init.constant_(module.bias, 0) + + self.apply(_basic_init) + + if self.use_pe: + # Initialize (and freeze) pos_embed by sin-cos embedding: + pos_embed = get_2d_sincos_pos_embed( + self.pos_embed.shape[-1], + int(self.x_embedder.num_patches**0.5), + pe_interpolation=self.pe_interpolation, + base_size=self.base_size, + ) + self.pos_embed.data.copy_(torch.from_numpy(pos_embed).float().unsqueeze(0)) + + # Initialize patch_embed like nn.Linear (instead of nn.Conv2d): + w = self.x_embedder.proj.weight.data + nn.init.xavier_uniform_(w.view([w.shape[0], -1])) + + # Initialize timestep embedding MLP: + nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02) + nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02) + nn.init.normal_(self.t_block[1].weight, std=0.02) + + # Initialize caption embedding MLP: + nn.init.normal_(self.y_embedder.y_proj.fc1.weight, std=0.02) + nn.init.normal_(self.y_embedder.y_proj.fc2.weight, std=0.02) + + +def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False, extra_tokens=0, pe_interpolation=1.0, base_size=16): + """ + grid_size: int of the grid height and width + return: + pos_embed: [grid_size*grid_size, embed_dim] or [1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token) + """ + if isinstance(grid_size, int): + grid_size = to_2tuple(grid_size) + grid_h = np.arange(grid_size[0], dtype=np.float32) / (grid_size[0] / base_size) / pe_interpolation + grid_w = np.arange(grid_size[1], dtype=np.float32) / (grid_size[1] / base_size) / pe_interpolation + grid = np.meshgrid(grid_w, grid_h) # here w goes first + grid = np.stack(grid, axis=0) + grid = grid.reshape([2, 1, grid_size[1], grid_size[0]]) + + pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid) + if cls_token and extra_tokens > 0: + pos_embed = np.concatenate([np.zeros([extra_tokens, embed_dim]), pos_embed], axis=0) + return pos_embed + + +def get_2d_sincos_pos_embed_from_grid(embed_dim, grid): + assert embed_dim % 2 == 0 + + # use half of dimensions to encode grid_h + emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2) + emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2) + + emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D) + return emb + + +def get_1d_sincos_pos_embed_from_grid(embed_dim, pos): + """ + embed_dim: output dimension for each position + pos: a list of positions to be encoded: size (M,) + out: (M, D) + """ + assert embed_dim % 2 == 0 + omega = np.arange(embed_dim // 2, dtype=np.float64) + omega /= embed_dim / 2.0 + omega = 1.0 / 10000**omega # (D/2,) + + pos = pos.reshape(-1) # (M,) + out = np.einsum("m,d->md", pos, omega) # (M, D/2), outer product + + emb_sin = np.sin(out) # (M, D/2) + emb_cos = np.cos(out) # (M, D/2) + + emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D) + return emb \ No newline at end of file diff --git a/Sana/models/sana_blocks.py b/Sana/models/sana_blocks.py new file mode 100644 index 0000000..31ac821 --- /dev/null +++ b/Sana/models/sana_blocks.py @@ -0,0 +1,798 @@ +# Copyright 2024 NVIDIA CORPORATION & AFFILIATES +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# SPDX-License-Identifier: Apache-2.0 + +# This file is modified from https://github.com/PixArt-alpha/PixArt-sigma +import math +import os +from typing import Optional + +import xformers.ops +import torch +import torch.nn as nn +import torch.nn.functional as F +from einops import rearrange +from timm.models.vision_transformer import Attention as Attention_ +from timm.models.vision_transformer import Mlp +from transformers import AutoModelForCausalLM + +from .norms import RMSNorm +from .utils import get_same_padding, to_2tuple + +sdpa_32b = None +Q_4GB_LIMIT = 32000000 +"""If q is greater than this, the operation will likely require >4GB VRAM, which will fail on Intel Arc Alchemist GPUs without a workaround.""" +# 2k = 37 748 736 +# 1024 = 9 437 184 +# 2k model goes very slightly over 4GB + +from comfy import model_management +if model_management.xformers_enabled(): + import xformers + import xformers.ops +else: + if model_management.xpu_available: + import intel_extension_for_pytorch as ipex + import os + if not torch.xpu.has_fp64_dtype() and not os.environ.get('IPEX_FORCE_ATTENTION_SLICE', None): + from ...utils.IPEX.attention import scaled_dot_product_attention_32_bit + sdpa_32b = scaled_dot_product_attention_32_bit + print("Using IPEX 4GB SDPA workaround") + else: + print("No IPEX 4GB workaround") + + +def modulate(x, shift, scale): + return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1) + + +def t2i_modulate(x, shift, scale): + return x * (1 + scale) + shift + + +class MultiHeadCrossAttention(nn.Module): + def __init__(self, d_model, num_heads, attn_drop=0.0, proj_drop=0.0, qk_norm=False, **block_kwargs): + super().__init__() + assert d_model % num_heads == 0, "d_model must be divisible by num_heads" + + self.d_model = d_model + self.num_heads = num_heads + self.head_dim = d_model // num_heads + + self.q_linear = nn.Linear(d_model, d_model) + self.kv_linear = nn.Linear(d_model, d_model * 2) + self.attn_drop = nn.Dropout(attn_drop) + self.proj = nn.Linear(d_model, d_model) + self.proj_drop = nn.Dropout(proj_drop) + if qk_norm: + # not used for now + self.q_norm = RMSNorm(d_model, scale_factor=1.0, eps=1e-6) + self.k_norm = RMSNorm(d_model, scale_factor=1.0, eps=1e-6) + else: + self.q_norm = nn.Identity() + self.k_norm = nn.Identity() + + def forward(self, x, cond, mask=None): + # query/value: img tokens; key: condition; mask: if padding tokens + B, N, C = x.shape + + q = self.q_linear(x).view(1, -1, self.num_heads, self.head_dim) + kv = self.kv_linear(cond).view(1, -1, 2, self.num_heads, self.head_dim) + k, v = kv.unbind(2) + + if model_management.xformers_enabled(): + attn_bias = None + if mask is not None: + attn_bias = xformers.ops.fmha.BlockDiagonalMask.from_seqlens([N] * B, mask) + x = xformers.ops.memory_efficient_attention( + q, k, v, + p=self.attn_drop.p, + attn_bias=attn_bias + ) + else: + q, k, v = map(lambda t: t.permute(0, 2, 1, 3),(q, k, v),) + attn_mask = None + if mask is not None and len(mask) > 1: + + # Create equivalent of xformer diagonal block mask, still only correct for square masks + # But depth doesn't matter as tensors can expand in that dimension + attn_mask_template = torch.ones( + [q.shape[2] // B, mask[0]], + dtype=torch.bool, + device=q.device + ) + attn_mask = torch.block_diag(attn_mask_template) + + # create a mask on the diagonal for each mask in the batch + for n in range(B - 1): + attn_mask = torch.block_diag(attn_mask, attn_mask_template) + + p = getattr(self.attn_drop, "p", 0) # IPEX.optimize() will turn attn_drop into an Identity() + + if sdpa_32b is not None and (q.element_size() * q.nelement()) > Q_4GB_LIMIT: + sdpa = sdpa_32b + else: + sdpa = torch.nn.functional.scaled_dot_product_attention + + x = sdpa( + q, k, v, + attn_mask=attn_mask, + dropout_p=p + ).permute(0, 2, 1, 3).contiguous() + x = x.view(B, -1, C) + x = self.proj(x) + x = self.proj_drop(x) + return x + + +class LiteLA(Attention_): + r"""Lightweight linear attention""" + + PAD_VAL = 1 + + def __init__( + self, + in_dim: int, + out_dim: int, + heads: Optional[int] = None, + heads_ratio: float = 1.0, + dim=32, + eps=1e-15, + use_bias=False, + qk_norm=False, + norm_eps=1e-5, + ): + heads = heads or int(out_dim // dim * heads_ratio) + super().__init__(in_dim, num_heads=heads, qkv_bias=use_bias) + + self.in_dim = in_dim + self.out_dim = out_dim + self.heads = heads + self.dim = out_dim // heads # TODO: need some change + self.eps = eps + + self.kernel_func = nn.ReLU(inplace=False) + if qk_norm: + self.q_norm = RMSNorm(in_dim, scale_factor=1.0, eps=norm_eps) + self.k_norm = RMSNorm(in_dim, scale_factor=1.0, eps=norm_eps) + else: + self.q_norm = nn.Identity() + self.k_norm = nn.Identity() + + def attn_matmul(self, q, k, v: torch.Tensor) -> torch.Tensor: + # lightweight linear attention + q = self.kernel_func(q) # B, h, h_d, N + k = self.kernel_func(k) + + q, k, v = q.float(), k.float(), v.float() + + v = F.pad(v, (0, 0, 0, 1), mode="constant", value=LiteLA.PAD_VAL) + vk = torch.matmul(v, k) + out = torch.matmul(vk, q) + + if out.dtype in [torch.float16, torch.bfloat16]: + out = out.float() + out = out[:, :, :-1] / (out[:, :, -1:] + self.eps) + + return out + + def forward(self, x: torch.Tensor, mask=None, HW=None, block_id=None) -> torch.Tensor: + B, N, C = x.shape + + qkv = self.qkv(x).reshape(B, N, 3, C) + q, k, v = qkv.unbind(2) # B, N, 3, C --> B, N, C + dtype = q.dtype + + q = self.q_norm(q).transpose(-1, -2) # (B, N, C) -> (B, C, N) + k = self.k_norm(k).transpose(-1, -2) # (B, N, C) -> (B, C, N) + v = v.transpose(-1, -2) + + q = q.reshape(B, C // self.dim, self.dim, N) # (B, h, h_d, N) + k = k.reshape(B, C // self.dim, self.dim, N).transpose(-1, -2) # (B, h, N, h_d) + v = v.reshape(B, C // self.dim, self.dim, N) # (B, h, h_d, N) + + out = self.attn_matmul(q, k, v).to(dtype) + + out = out.view(B, C, N).permute(0, 2, 1) # B, N, C + out = self.proj(out) + + if torch.get_autocast_gpu_dtype() == torch.float16: + out = out.clip(-65504, 65504) + + return out + + @property + def module_str(self) -> str: + _str = type(self).__name__ + "(" + eps = f"{self.eps:.1E}" + _str += f"i={self.in_dim},o={self.out_dim},h={self.heads},d={self.dim},eps={eps}" + return _str + + def __repr__(self): + return f"EPS{self.eps}-" + super().__repr__() + + +class PAGCFGIdentitySelfAttnProcessorLiteLA: + r"""Self Attention with Perturbed Attention & CFG Guidance""" + + def __init__(self, attn): + self.attn = attn + + def __call__(self, x: torch.Tensor, mask=None, HW=None, block_id=None) -> torch.Tensor: + x_uncond, x_org, x_ptb = x.chunk(3) + x_org = torch.cat([x_uncond, x_org]) + B, N, C = x_org.shape + + qkv = self.attn.qkv(x_org).reshape(B, N, 3, C) + # B, N, 3, C --> B, N, C + q, k, v = qkv.unbind(2) + dtype = q.dtype + q = self.attn.q_norm(q).transpose(-1, -2) # (B, N, C) -> (B, C, N) + k = self.attn.k_norm(k).transpose(-1, -2) # (B, N, C) -> (B, C, N) + v = v.transpose(-1, -2) + + q = q.reshape(B, C // self.attn.dim, self.attn.dim, N) # (B, h, h_d, N) + k = k.reshape(B, C // self.attn.dim, self.attn.dim, N).transpose(-1, -2) # (B, h, N, h_d) + v = v.reshape(B, C // self.attn.dim, self.attn.dim, N) # (B, h, h_d, N) + + out = self.attn.attn_matmul(q, k, v).to(dtype) + + out = out.view(B, C, N).permute(0, 2, 1) # B, N, C + out = self.attn.proj(out) + + # perturbed path (identity attention) + v_weight = self.attn.qkv.weight[C * 2 : C * 3, :] # Shape: (dim, dim) + if self.attn.qkv.bias: + v_bias = self.attn.qkv.bias[C * 2 : C * 3] # Shape: (dim,) + x_ptb = (torch.matmul(x_ptb, v_weight.t()) + v_bias).to(dtype) + else: + x_ptb = torch.matmul(x_ptb, v_weight.t()).to(dtype) + x_ptb = self.attn.proj(x_ptb) + + out = torch.cat([out, x_ptb]) + + if torch.get_autocast_gpu_dtype() == torch.float16: + out = out.clip(-65504, 65504) + + return out + + +class PAGIdentitySelfAttnProcessorLiteLA: + r"""Self Attention with Perturbed Attention Guidance""" + + def __init__(self, attn): + self.attn = attn + + def __call__(self, x: torch.Tensor, mask=None, HW=None, block_id=None) -> torch.Tensor: + x_org, x_ptb = x.chunk(2) + B, N, C = x_org.shape + + qkv = self.attn.qkv(x_org).reshape(B, N, 3, C) + # B, N, 3, C --> B, N, C + q, k, v = qkv.unbind(2) + dtype = q.dtype + q = self.attn.q_norm(q).transpose(-1, -2) # (B, N, C) -> (B, C, N) + k = self.attn.k_norm(k).transpose(-1, -2) # (B, N, C) -> (B, C, N) + v = v.transpose(-1, -2) + + q = q.reshape(B, C // self.attn.dim, self.attn.dim, N) # (B, h, h_d, N) + k = k.reshape(B, C // self.attn.dim, self.attn.dim, N).transpose(-1, -2) # (B, h, N, h_d) + v = v.reshape(B, C // self.attn.dim, self.attn.dim, N) # (B, h, h_d, N) + + out = self.attn.attn_matmul(q, k, v).to(dtype) + + out = out.view(B, C, N).permute(0, 2, 1) # B, N, C + out = self.attn.proj(out) + + # perturbed path (identity attention) + v_weight = self.attn.qkv.weight[C * 2 : C * 3, :] # Shape: (dim, dim) + if self.attn.qkv.bias: + v_bias = self.attn.qkv.bias[C * 2 : C * 3] # Shape: (dim,) + x_ptb = (torch.matmul(x_ptb, v_weight.t()) + v_bias).to(dtype) + else: + x_ptb = torch.matmul(x_ptb, v_weight.t()).to(dtype) + x_ptb = self.attn.proj(x_ptb) + + out = torch.cat([out, x_ptb]) + + if torch.get_autocast_gpu_dtype() == torch.float16: + out = out.clip(-65504, 65504) + + return out + + +class SelfAttnProcessorLiteLA: + r"""Self Attention with Lite Linear Attention""" + + def __init__(self, attn): + self.attn = attn + + def __call__(self, x: torch.Tensor, mask=None, HW=None, block_id=None) -> torch.Tensor: + B, N, C = x.shape + if HW is None: + H = W = int(N**0.5) + else: + H, W = HW + qkv = self.attn.qkv(x).reshape(B, N, 3, C) + # B, N, 3, C --> B, N, C + q, k, v = qkv.unbind(2) + dtype = q.dtype + q = self.attn.q_norm(q).transpose(-1, -2) # (B, N, C) -> (B, C, N) + k = self.attn.k_norm(k).transpose(-1, -2) # (B, N, C) -> (B, C, N) + v = v.transpose(-1, -2) + + q = q.reshape(B, C // self.attn.dim, self.attn.dim, N) # (B, h, h_d, N) + k = k.reshape(B, C // self.attn.dim, self.attn.dim, N).transpose(-1, -2) # (B, h, N, h_d) + v = v.reshape(B, C // self.attn.dim, self.attn.dim, N) # (B, h, h_d, N) + + out = self.attn.attn_matmul(q, k, v).to(dtype) + + out = out.view(B, C, N).permute(0, 2, 1) # B, N, C + out = self.attn.proj(out) + + if torch.get_autocast_gpu_dtype() == torch.float16: + out = out.clip(-65504, 65504) + + return out + + +class FlashAttention(Attention_): + """Multi-head Flash Attention block with qk norm.""" + + def __init__( + self, + dim, + num_heads=8, + qkv_bias=True, + qk_norm=False, + **block_kwargs, + ): + """ + Args: + dim (int): Number of input channels. + num_heads (int): Number of attention heads. + qkv_bias (bool: If True, add a learnable bias to query, key, value. + """ + super().__init__(dim, num_heads=num_heads, qkv_bias=qkv_bias, **block_kwargs) + + if qk_norm: + self.q_norm = nn.LayerNorm(dim) + self.k_norm = nn.LayerNorm(dim) + else: + self.q_norm = nn.Identity() + self.k_norm = nn.Identity() + + def forward(self, x, mask=None, HW=None, block_id=None): + B, N, C = x.shape + + qkv = self.qkv(x).reshape(B, N, 3, C) + q, k, v = qkv.unbind(2) + dtype = q.dtype + + q = self.q_norm(q) + k = self.k_norm(k) + + q = q.reshape(B, N, self.num_heads, C // self.num_heads).to(dtype) + k = k.reshape(B, N, self.num_heads, C // self.num_heads).to(dtype) + v = v.reshape(B, N, self.num_heads, C // self.num_heads).to(dtype) + + use_fp32_attention = getattr(self, "fp32_attention", False) # necessary for NAN loss + if use_fp32_attention: + q, k, v = q.float(), k.float(), v.float() + + attn_bias = None + if mask is not None: + attn_bias = torch.zeros([B * self.num_heads, q.shape[1], k.shape[1]], dtype=q.dtype, device=q.device) + attn_bias.masked_fill_(mask.squeeze(1).repeat(self.num_heads, 1, 1) == 0, float("-inf")) + + if _xformers_available: + x = xformers.ops.memory_efficient_attention(q, k, v, p=self.attn_drop.p, attn_bias=attn_bias) + else: + q, k, v = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2) + if mask is not None and mask.ndim == 2: + mask = (1 - mask.to(x.dtype)) * -10000.0 + mask = mask[:, None, None].repeat(1, self.num_heads, 1, 1) + x = F.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=0.0, is_causal=False) + x = x.transpose(1, 2) + + x = x.view(B, N, C) + x = self.proj(x) + x = self.proj_drop(x) + + if torch.get_autocast_gpu_dtype() == torch.float16: + x = x.clip(-65504, 65504) + + return x + + +################################################################################# +# AMP attention with fp32 softmax to fix loss NaN problem during training # +################################################################################# +class Attention(Attention_): + def forward(self, x, HW=None): + B, N, C = x.shape + qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4) + # B,N,3,H,C -> B,H,N,C + q, k, v = qkv.unbind(0) # make torchscript happy (cannot use tensor as tuple) + use_fp32_attention = getattr(self, "fp32_attention", False) + if use_fp32_attention: + q, k = q.float(), k.float() + + with torch.cuda.amp.autocast(enabled=not use_fp32_attention): + attn = (q @ k.transpose(-2, -1)) * self.scale + attn = attn.softmax(dim=-1) + + attn = self.attn_drop(attn) + + x = (attn @ v).transpose(1, 2).reshape(B, N, C) + x = self.proj(x) + x = self.proj_drop(x) + return x + + +class FinalLayer(nn.Module): + """ + The final layer of Sana. + """ + + def __init__(self, hidden_size, patch_size, out_channels): + super().__init__() + self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True) + self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True)) + + def forward(self, x, c): + shift, scale = self.adaLN_modulation(c).chunk(2, dim=1) + x = modulate(self.norm_final(x), shift, scale) + x = self.linear(x) + return x + + +class T2IFinalLayer(nn.Module): + """ + The final layer of Sana. + """ + + def __init__(self, hidden_size, patch_size, out_channels): + super().__init__() + self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True) + self.scale_shift_table = nn.Parameter(torch.randn(2, hidden_size) / hidden_size**0.5) + self.out_channels = out_channels + + def forward(self, x, t): + shift, scale = (self.scale_shift_table[None] + t[:, None]).chunk(2, dim=1) + x = t2i_modulate(self.norm_final(x), shift, scale) + x = self.linear(x) + return x + + +class MaskFinalLayer(nn.Module): + """ + The final layer of Sana. + """ + + def __init__(self, final_hidden_size, c_emb_size, patch_size, out_channels): + super().__init__() + self.norm_final = nn.LayerNorm(final_hidden_size, elementwise_affine=False, eps=1e-6) + self.linear = nn.Linear(final_hidden_size, patch_size * patch_size * out_channels, bias=True) + self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(c_emb_size, 2 * final_hidden_size, bias=True)) + + def forward(self, x, t): + shift, scale = self.adaLN_modulation(t).chunk(2, dim=1) + x = modulate(self.norm_final(x), shift, scale) + x = self.linear(x) + return x + + +class DecoderLayer(nn.Module): + """ + The final layer of Sana. + """ + + def __init__(self, hidden_size, decoder_hidden_size): + super().__init__() + self.norm_decoder = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + self.linear = nn.Linear(hidden_size, decoder_hidden_size, bias=True) + self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True)) + + def forward(self, x, t): + shift, scale = self.adaLN_modulation(t).chunk(2, dim=1) + x = modulate(self.norm_decoder(x), shift, scale) + x = self.linear(x) + return x + + +################################################################################# +# Embedding Layers for Timesteps and Class Labels # +################################################################################# +class TimestepEmbedder(nn.Module): + """ + Embeds scalar timesteps into vector representations. + """ + + def __init__(self, hidden_size, frequency_embedding_size=256): + super().__init__() + self.mlp = nn.Sequential( + nn.Linear(frequency_embedding_size, hidden_size, bias=True), + nn.SiLU(), + nn.Linear(hidden_size, hidden_size, bias=True), + ) + self.frequency_embedding_size = frequency_embedding_size + + @staticmethod + def timestep_embedding(t, dim, max_period=10000): + """ + Create sinusoidal timestep embeddings. + :param t: a 1-D Tensor of N indices, one per batch element. + These may be fractional. + :param dim: the dimension of the output. + :param max_period: controls the minimum frequency of the embeddings. + :return: an (N, D) Tensor of positional embeddings. + """ + # https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py + half = dim // 2 + freqs = torch.exp( + -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32, device=t.device) / half + ) + args = t[:, None].float() * freqs[None] + embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) + if dim % 2: + embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) + return embedding + + def forward(self, t): + t_freq = self.timestep_embedding(t, self.frequency_embedding_size).to(self.dtype) + t_emb = self.mlp(t_freq) + return t_emb + + @property + def dtype(self): + try: + return next(self.parameters()).dtype + except StopIteration: + return torch.float32 + + +class SizeEmbedder(TimestepEmbedder): + """ + Embeds scalar timesteps into vector representations. + """ + + def __init__(self, hidden_size, frequency_embedding_size=256): + super().__init__(hidden_size=hidden_size, frequency_embedding_size=frequency_embedding_size) + self.mlp = nn.Sequential( + nn.Linear(frequency_embedding_size, hidden_size, bias=True), + nn.SiLU(), + nn.Linear(hidden_size, hidden_size, bias=True), + ) + self.frequency_embedding_size = frequency_embedding_size + self.outdim = hidden_size + + def forward(self, s, bs): + if s.ndim == 1: + s = s[:, None] + assert s.ndim == 2 + if s.shape[0] != bs: + s = s.repeat(bs // s.shape[0], 1) + assert s.shape[0] == bs + b, dims = s.shape[0], s.shape[1] + s = rearrange(s, "b d -> (b d)") + s_freq = self.timestep_embedding(s, self.frequency_embedding_size).to(self.dtype) + s_emb = self.mlp(s_freq) + s_emb = rearrange(s_emb, "(b d) d2 -> b (d d2)", b=b, d=dims, d2=self.outdim) + return s_emb + + @property + def dtype(self): + try: + return next(self.parameters()).dtype + except StopIteration: + return torch.float32 + + +class LabelEmbedder(nn.Module): + """ + Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance. + """ + + def __init__(self, num_classes, hidden_size, dropout_prob): + super().__init__() + use_cfg_embedding = dropout_prob > 0 + self.embedding_table = nn.Embedding(num_classes + use_cfg_embedding, hidden_size) + self.num_classes = num_classes + self.dropout_prob = dropout_prob + + def token_drop(self, labels, force_drop_ids=None): + """ + Drops labels to enable classifier-free guidance. + """ + if force_drop_ids is None: + drop_ids = torch.rand(labels.shape[0]).cuda() < self.dropout_prob + else: + drop_ids = force_drop_ids == 1 + labels = torch.where(drop_ids, self.num_classes, labels) + return labels + + def forward(self, labels, train, force_drop_ids=None): + use_dropout = self.dropout_prob > 0 + if (train and use_dropout) or (force_drop_ids is not None): + labels = self.token_drop(labels, force_drop_ids) + embeddings = self.embedding_table(labels) + return embeddings + + +class CaptionEmbedder(nn.Module): + """ + Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance. + """ + + def __init__( + self, + in_channels, + hidden_size, + uncond_prob, + act_layer=nn.GELU(approximate="tanh"), + token_num=120, + ): + super().__init__() + self.y_proj = Mlp( + in_features=in_channels, hidden_features=hidden_size, out_features=hidden_size, act_layer=act_layer, drop=0 + ) + self.register_buffer("y_embedding", nn.Parameter(torch.randn(token_num, in_channels) / in_channels**0.5)) + self.uncond_prob = uncond_prob + + def initialize_gemma_params(self, model_name="google/gemma-2b-it"): + num_layers = len(self.custom_gemma_layers) + text_encoder = AutoModelForCausalLM.from_pretrained(model_name).get_decoder() + pretrained_layers = text_encoder.layers[-num_layers:] + for custom_layer, pretrained_layer in zip(self.custom_gemma_layers, pretrained_layers): + info = custom_layer.load_state_dict(pretrained_layer.state_dict(), strict=False) + print(f"**** {info} ****") + print(f"**** Initialized {num_layers} Gemma layers from pretrained model: {model_name} ****") + + def token_drop(self, caption, force_drop_ids=None): + """ + Drops labels to enable classifier-free guidance. + """ + if force_drop_ids is None: + drop_ids = torch.rand(caption.shape[0]).cuda() < self.uncond_prob + else: + drop_ids = force_drop_ids == 1 + caption = torch.where(drop_ids[:, None, None, None], self.y_embedding, caption) + return caption + + def forward(self, caption, train, force_drop_ids=None, mask=None): + if train: + assert caption.shape[2:] == self.y_embedding.shape + use_dropout = self.uncond_prob > 0 + if (train and use_dropout) or (force_drop_ids is not None): + caption = self.token_drop(caption, force_drop_ids) + + caption = self.y_proj(caption) + + return caption + + +class CaptionEmbedderDoubleBr(nn.Module): + """ + Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance. + """ + + def __init__(self, in_channels, hidden_size, uncond_prob, act_layer=nn.GELU(approximate="tanh"), token_num=120): + super().__init__() + self.proj = Mlp( + in_features=in_channels, hidden_features=hidden_size, out_features=hidden_size, act_layer=act_layer, drop=0 + ) + self.embedding = nn.Parameter(torch.randn(1, in_channels) / 10**0.5) + self.y_embedding = nn.Parameter(torch.randn(token_num, in_channels) / 10**0.5) + self.uncond_prob = uncond_prob + + def token_drop(self, global_caption, caption, force_drop_ids=None): + """ + Drops labels to enable classifier-free guidance. + """ + if force_drop_ids is None: + drop_ids = torch.rand(global_caption.shape[0]).cuda() < self.uncond_prob + else: + drop_ids = force_drop_ids == 1 + global_caption = torch.where(drop_ids[:, None], self.embedding, global_caption) + caption = torch.where(drop_ids[:, None, None, None], self.y_embedding, caption) + return global_caption, caption + + def forward(self, caption, train, force_drop_ids=None): + assert caption.shape[2:] == self.y_embedding.shape + global_caption = caption.mean(dim=2).squeeze() + use_dropout = self.uncond_prob > 0 + if (train and use_dropout) or (force_drop_ids is not None): + global_caption, caption = self.token_drop(global_caption, caption, force_drop_ids) + y_embed = self.proj(global_caption) + return y_embed, caption + + +class PatchEmbed(nn.Module): + """2D Image to Patch Embedding""" + + def __init__( + self, + img_size=224, + patch_size=16, + in_chans=3, + embed_dim=768, + kernel_size=None, + padding=0, + norm_layer=None, + flatten=True, + bias=True, + ): + super().__init__() + kernel_size = kernel_size or patch_size + img_size = to_2tuple(img_size) + patch_size = to_2tuple(patch_size) + self.img_size = img_size + self.patch_size = patch_size + self.grid_size = (img_size[0] // patch_size[0], img_size[1] // patch_size[1]) + self.num_patches = self.grid_size[0] * self.grid_size[1] + self.flatten = flatten + if not padding and kernel_size % 2 > 0: + padding = get_same_padding(kernel_size) + self.proj = nn.Conv2d( + in_chans, embed_dim, kernel_size=kernel_size, stride=patch_size, padding=padding, bias=bias + ) + self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity() + + def forward(self, x): + B, C, H, W = x.shape + assert (H == self.img_size[0], f"Input image height ({H}) doesn't match model ({self.img_size[0]}).") + assert (W == self.img_size[1], f"Input image width ({W}) doesn't match model ({self.img_size[1]}).") + x = self.proj(x) + if self.flatten: + x = x.flatten(2).transpose(1, 2) # BCHW -> BNC + x = self.norm(x) + return x + + +class PatchEmbedMS(nn.Module): + """2D Image to Patch Embedding""" + + def __init__( + self, + patch_size=16, + in_chans=3, + embed_dim=768, + kernel_size=None, + padding=0, + norm_layer=None, + flatten=True, + bias=True, + ): + super().__init__() + kernel_size = kernel_size or patch_size + patch_size = to_2tuple(patch_size) + self.patch_size = patch_size + self.flatten = flatten + if not padding and kernel_size % 2 > 0: + padding = get_same_padding(kernel_size) + self.proj = nn.Conv2d( + in_chans, embed_dim, kernel_size=kernel_size, stride=patch_size, padding=padding, bias=bias + ) + self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity() + + def forward(self, x): + x = self.proj(x) + if self.flatten: + x = x.flatten(2).transpose(1, 2) # BCHW -> BNC + x = self.norm(x) + return x diff --git a/Sana/models/sana_multi_scale.py b/Sana/models/sana_multi_scale.py new file mode 100644 index 0000000..7cc3745 --- /dev/null +++ b/Sana/models/sana_multi_scale.py @@ -0,0 +1,374 @@ +# Copyright 2024 NVIDIA CORPORATION & AFFILIATES +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# SPDX-License-Identifier: Apache-2.0 + +# This file is modified from https://github.com/PixArt-alpha/PixArt-sigma +import torch +import torch.nn as nn +from timm.models.layers import DropPath + +from .basic_modules import DWMlp, GLUMBConv, MBConvPreGLU, Mlp +from .sana import Sana, get_2d_sincos_pos_embed +from .sana_blocks import ( + Attention, + CaptionEmbedder, + FlashAttention, + LiteLA, + MultiHeadCrossAttention, + PatchEmbedMS, + T2IFinalLayer, + t2i_modulate, +) +from .utils import auto_grad_checkpoint + + +class SanaMSBlock(nn.Module): + """ + A Sana block with global shared adaptive layer norm zero (adaLN-Zero) conditioning. + """ + + def __init__( + self, + hidden_size, + num_heads, + mlp_ratio=4.0, + drop_path=0.0, + input_size=None, + qk_norm=False, + attn_type="flash", + ffn_type="mlp", + mlp_acts=("silu", "silu", None), + linear_head_dim=32, + cross_norm=False, + **block_kwargs, + ): + super().__init__() + self.hidden_size = hidden_size + self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + if attn_type == "flash": + # flash self attention + self.attn = FlashAttention( + hidden_size, + num_heads=num_heads, + qkv_bias=True, + qk_norm=qk_norm, + **block_kwargs, + ) + elif attn_type == "linear": + # linear self attention + # TODO: Here the num_heads set to 36 for tmp used + self_num_heads = hidden_size // linear_head_dim + self.attn = LiteLA(hidden_size, hidden_size, heads=self_num_heads, eps=1e-8, qk_norm=qk_norm) + elif attn_type == "vanilla": + # vanilla self attention + self.attn = Attention(hidden_size, num_heads=num_heads, qkv_bias=True) + else: + raise ValueError(f"{attn_type} type is not defined.") + + self.cross_attn = MultiHeadCrossAttention(hidden_size, num_heads, qk_norm=cross_norm, **block_kwargs) + self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + if ffn_type == "dwmlp": + approx_gelu = lambda: nn.GELU(approximate="tanh") + self.mlp = DWMlp( + in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu, drop=0 + ) + elif ffn_type == "glumbconv": + self.mlp = GLUMBConv( + in_features=hidden_size, + hidden_features=int(hidden_size * mlp_ratio), + use_bias=(True, True, False), + norm=(None, None, None), + act=mlp_acts, + ) + elif ffn_type == "glumbconv_dilate": + self.mlp = GLUMBConv( + in_features=hidden_size, + hidden_features=int(hidden_size * mlp_ratio), + use_bias=(True, True, False), + norm=(None, None, None), + act=mlp_acts, + dilation=2, + ) + elif ffn_type == "mlp": + approx_gelu = lambda: nn.GELU(approximate="tanh") + self.mlp = Mlp( + in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu, drop=0 + ) + elif ffn_type == "mbconvpreglu": + self.mlp = MBConvPreGLU( + in_dim=hidden_size, + out_dim=hidden_size, + mid_dim=int(hidden_size * mlp_ratio), + use_bias=(True, True, False), + norm=None, + act=mlp_acts, + ) + else: + raise ValueError(f"{ffn_type} type is not defined.") + self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() + self.scale_shift_table = nn.Parameter(torch.randn(6, hidden_size) / hidden_size**0.5) + + def forward(self, x, y, t, mask=None, HW=None, **kwargs): + B, N, C = x.shape + + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( + self.scale_shift_table[None] + t.reshape(B, 6, -1) + ).chunk(6, dim=1) + x = x + self.drop_path(gate_msa * self.attn(t2i_modulate(self.norm1(x), shift_msa, scale_msa), HW=HW)) + x = x + self.cross_attn(x, y, mask) + x = x + self.drop_path(gate_mlp * self.mlp(t2i_modulate(self.norm2(x), shift_mlp, scale_mlp), HW=HW)) + + return x + + +############################################################################# +# Core Sana Model # +################################################################################# +class SanaMS(Sana): + """ + Diffusion model with a Transformer backbone. + """ + + def __init__( + self, + input_size=32, + patch_size=2, + in_channels=32, + hidden_size=1152, + depth=28, + num_heads=16, + mlp_ratio=4.0, + class_dropout_prob=0.1, + learn_sigma=False, + pred_sigma=False, + drop_path: float = 0.0, + caption_channels=2304, + pe_interpolation=1.0, + config=None, + model_max_length=300, + qk_norm=False, + y_norm=False, + norm_eps=1e-5, + attn_type="linear", + ffn_type="glumbconv", + use_pe=False, + y_norm_scale_factor=1.0, + patch_embed_kernel=None, + mlp_acts=("silu", "silu", None), + linear_head_dim=32, + cross_norm=False, + **kwargs, + ): + super().__init__( + input_size=input_size, + patch_size=patch_size, + in_channels=in_channels, + hidden_size=hidden_size, + depth=depth, + num_heads=num_heads, + mlp_ratio=mlp_ratio, + class_dropout_prob=class_dropout_prob, + learn_sigma=learn_sigma, + pred_sigma=pred_sigma, + drop_path=drop_path, + caption_channels=caption_channels, + pe_interpolation=pe_interpolation, + config=config, + model_max_length=model_max_length, + qk_norm=qk_norm, + y_norm=y_norm, + norm_eps=norm_eps, + attn_type=attn_type, + ffn_type=ffn_type, + use_pe=use_pe, + y_norm_scale_factor=y_norm_scale_factor, + patch_embed_kernel=patch_embed_kernel, + mlp_acts=mlp_acts, + linear_head_dim=linear_head_dim, + **kwargs, + ) + self.dtype = torch.get_default_dtype() + self.h = self.w = 0 + approx_gelu = lambda: nn.GELU(approximate="tanh") + self.t_block = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size, bias=True)) + self.pos_embed_ms = None + + kernel_size = patch_embed_kernel or patch_size + self.x_embedder = PatchEmbedMS(patch_size, in_channels, hidden_size, kernel_size=kernel_size, bias=True) + self.y_embedder = CaptionEmbedder( + in_channels=caption_channels, + hidden_size=hidden_size, + uncond_prob=class_dropout_prob, + act_layer=approx_gelu, + token_num=model_max_length, + ) + drop_path = [x.item() for x in torch.linspace(0, drop_path, depth)] # stochastic depth decay rule + self.blocks = nn.ModuleList( + [ + SanaMSBlock( + hidden_size, + num_heads, + mlp_ratio=mlp_ratio, + drop_path=drop_path[i], + input_size=(input_size // patch_size, input_size // patch_size), + qk_norm=qk_norm, + attn_type=attn_type, + ffn_type=ffn_type, + mlp_acts=mlp_acts, + linear_head_dim=linear_head_dim, + cross_norm=cross_norm, + ) + for i in range(depth) + ] + ) + self.final_layer = T2IFinalLayer(hidden_size, patch_size, self.out_channels) + + self.initialize() + + def forward(self, x, timesteps, context, **kwargs): + """ + Forward pass that adapts comfy input to original forward function + x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images) + timesteps: (N,) tensor of diffusion timesteps + context: (N, 1, 120, C) conditioning + """ + ## size/ar from cond with fallback based on the latent image shape. + bs = x.shape[0] + ## Still accepts the input w/o that dim but returns garbage + if len(context.shape) == 3: + context = context.unsqueeze(1) + + ## run original forward pass + out = self.forward_raw( + x = x.to(self.dtype), + timestep = timesteps.to(self.dtype), + y = context.to(self.dtype), + ) + + ## only return EPS + out = out.to(torch.float) + + return out + + def forward_raw(self, x, timestep, y, mask=None, data_info=None, **kwargs): + """ + Forward pass of Sana. + x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images) + t: (N,) tensor of diffusion timesteps + y: (N, 1, 120, C) tensor of class labels + """ + bs = x.shape[0] + x = x.to(self.dtype) + timestep = timestep.to(self.dtype) + y = y.to(self.dtype) + self.h, self.w = x.shape[-2] // self.patch_size, x.shape[-1] // self.patch_size + if self.use_pe: + x = self.x_embedder(x) + if self.pos_embed_ms is None or self.pos_embed_ms.shape[1:] != x.shape[1:]: + self.pos_embed_ms = ( + torch.from_numpy( + get_2d_sincos_pos_embed( + self.pos_embed.shape[-1], + (self.h, self.w), + pe_interpolation=self.pe_interpolation, + base_size=self.base_size, + ) + ) + .unsqueeze(0) + .to(x.device) + .to(self.dtype) + ) + x += self.pos_embed_ms # (N, T, D), where T = H * W / patch_size ** 2 + else: + x = self.x_embedder(x) + + t = self.t_embedder(timestep) # (N, D) + + t0 = self.t_block(t) + y = self.y_embedder(y, self.training, mask=mask) # (N, D) + if self.y_norm: + y = self.attention_y_norm(y) + + if mask is not None: + if mask.shape[0] != y.shape[0]: + mask = mask.repeat(y.shape[0] // mask.shape[0], 1) + mask = mask.squeeze(1).squeeze(1) + y = y.squeeze(1).masked_select(mask.unsqueeze(-1) != 0).view(1, -1, x.shape[-1]) + y_lens = mask.sum(dim=1).tolist() + else: + y_lens = [y.shape[2]] * y.shape[0] + y = y.squeeze(1).view(1, -1, x.shape[-1]) + + for block in self.blocks: + x = auto_grad_checkpoint( + block, x, y, t0, y_lens, (self.h, self.w), **kwargs + ) # (N, T, D) #support grad checkpoint + + x = self.final_layer(x, t) # (N, T, patch_size ** 2 * out_channels) + x = self.unpatchify(x) # (N, out_channels, H, W) + + return x + + def __call__(self, *args, **kwargs): + """ + This method allows the object to be called like a function. + It simply calls the forward method. + """ + return self.forward(*args, **kwargs) + + def forward_with_dpmsolver(self, x, timestep, y, data_info, **kwargs): + """ + dpm solver donnot need variance prediction + """ + # https://github.com/openai/glide-text2im/blob/main/notebooks/text2im.ipynb + model_out = self.forward(x, timestep, y, data_info=data_info, **kwargs) + return model_out.chunk(2, dim=1)[0] if self.pred_sigma else model_out + + def unpatchify(self, x): + """ + x: (N, T, patch_size**2 * C) + imgs: (N, H, W, C) + """ + c = self.out_channels + p = self.x_embedder.patch_size[0] + assert self.h * self.w == x.shape[1] + + x = x.reshape(shape=(x.shape[0], self.h, self.w, p, p, c)) + x = torch.einsum("nhwpqc->nchpwq", x) + imgs = x.reshape(shape=(x.shape[0], c, self.h * p, self.w * p)) + return imgs + + def initialize(self): + # Initialize transformer layers: + def _basic_init(module): + if isinstance(module, nn.Linear): + torch.nn.init.xavier_uniform_(module.weight) + if module.bias is not None: + nn.init.constant_(module.bias, 0) + + self.apply(_basic_init) + + # Initialize patch_embed like nn.Linear (instead of nn.Conv2d): + w = self.x_embedder.proj.weight.data + nn.init.xavier_uniform_(w.view([w.shape[0], -1])) + + # Initialize timestep embedding MLP: + nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02) + nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02) + nn.init.normal_(self.t_block[1].weight, std=0.02) + + # Initialize caption embedding MLP: + nn.init.normal_(self.y_embedder.y_proj.fc1.weight, std=0.02) + nn.init.normal_(self.y_embedder.y_proj.fc2.weight, std=0.02) diff --git a/Sana/models/utils.py b/Sana/models/utils.py new file mode 100644 index 0000000..d74db3b --- /dev/null +++ b/Sana/models/utils.py @@ -0,0 +1,591 @@ +# Copyright 2024 NVIDIA CORPORATION & AFFILIATES +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# SPDX-License-Identifier: Apache-2.0 + +import math +import os +import random +import re +import sys +from collections.abc import Iterable +from itertools import repeat + +import torch +import torch.distributed as dist +import torch.nn as nn +import torch.nn.functional as F +from PIL import Image +from torch.utils.checkpoint import checkpoint, checkpoint_sequential +from torchvision import transforms as T + + +def _ntuple(n): + def parse(x): + if isinstance(x, Iterable) and not isinstance(x, str): + return x + return tuple(repeat(x, n)) + + return parse + + +to_1tuple = _ntuple(1) +to_2tuple = _ntuple(2) + + +def set_grad_checkpoint(model, gc_step=1): + assert isinstance(model, nn.Module) + + def set_attr(module): + module.grad_checkpointing = True + module.grad_checkpointing_step = gc_step + + model.apply(set_attr) + + +def set_fp32_attention(model): + assert isinstance(model, nn.Module) + + def set_attr(module): + module.fp32_attention = True + + model.apply(set_attr) + + +def auto_grad_checkpoint(module, *args, **kwargs): + if getattr(module, "grad_checkpointing", False): + if isinstance(module, Iterable): + gc_step = module[0].grad_checkpointing_step + return checkpoint_sequential(module, gc_step, *args, **kwargs) + else: + return checkpoint(module, *args, **kwargs) + return module(*args, **kwargs) + + +def checkpoint_sequential(functions, step, input, *args, **kwargs): + + # Hack for keyword-only parameter in a python 2.7-compliant way + preserve = kwargs.pop("preserve_rng_state", True) + if kwargs: + raise ValueError("Unexpected keyword arguments: " + ",".join(arg for arg in kwargs)) + + def run_function(start, end, functions): + def forward(input): + for j in range(start, end + 1): + input = functions[j](input, *args) + return input + + return forward + + if isinstance(functions, torch.nn.Sequential): + functions = list(functions.children()) + + # the last chunk has to be non-volatile + end = -1 + segment = len(functions) // step + for start in range(0, step * (segment - 1), step): + end = start + step - 1 + input = checkpoint(run_function(start, end, functions), input, preserve_rng_state=preserve) + return run_function(end + 1, len(functions) - 1, functions)(input) + + +def window_partition(x, window_size): + """ + Partition into non-overlapping windows with padding if needed. + Args: + x (tensor): input tokens with [B, H, W, C]. + window_size (int): window size. + + Returns: + windows: windows after partition with [B * num_windows, window_size, window_size, C]. + (Hp, Wp): padded height and width before partition + """ + B, H, W, C = x.shape + + pad_h = (window_size - H % window_size) % window_size + pad_w = (window_size - W % window_size) % window_size + if pad_h > 0 or pad_w > 0: + x = F.pad(x, (0, 0, 0, pad_w, 0, pad_h)) + Hp, Wp = H + pad_h, W + pad_w + + x = x.view(B, Hp // window_size, window_size, Wp // window_size, window_size, C) + windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C) + return windows, (Hp, Wp) + + +def window_unpartition(windows, window_size, pad_hw, hw): + """ + Window unpartition into original sequences and removing padding. + Args: + x (tensor): input tokens with [B * num_windows, window_size, window_size, C]. + window_size (int): window size. + pad_hw (Tuple): padded height and width (Hp, Wp). + hw (Tuple): original height and width (H, W) before padding. + + Returns: + x: unpartitioned sequences with [B, H, W, C]. + """ + Hp, Wp = pad_hw + H, W = hw + B = windows.shape[0] // (Hp * Wp // window_size // window_size) + x = windows.view(B, Hp // window_size, Wp // window_size, window_size, window_size, -1) + x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, Hp, Wp, -1) + + if Hp > H or Wp > W: + x = x[:, :H, :W, :].contiguous() + return x + + +def get_rel_pos(q_size, k_size, rel_pos): + """ + Get relative positional embeddings according to the relative positions of + query and key sizes. + Args: + q_size (int): size of query q. + k_size (int): size of key k. + rel_pos (Tensor): relative position embeddings (L, C). + + Returns: + Extracted positional embeddings according to relative positions. + """ + max_rel_dist = int(2 * max(q_size, k_size) - 1) + # Interpolate rel pos if needed. + if rel_pos.shape[0] != max_rel_dist: + # Interpolate rel pos. + rel_pos_resized = F.interpolate( + rel_pos.reshape(1, rel_pos.shape[0], -1).permute(0, 2, 1), + size=max_rel_dist, + mode="linear", + ) + rel_pos_resized = rel_pos_resized.reshape(-1, max_rel_dist).permute(1, 0) + else: + rel_pos_resized = rel_pos + + # Scale the coords with short length if shapes for q and k are different. + q_coords = torch.arange(q_size)[:, None] * max(k_size / q_size, 1.0) + k_coords = torch.arange(k_size)[None, :] * max(q_size / k_size, 1.0) + relative_coords = (q_coords - k_coords) + (k_size - 1) * max(q_size / k_size, 1.0) + + return rel_pos_resized[relative_coords.long()] + + +def add_decomposed_rel_pos(attn, q, rel_pos_h, rel_pos_w, q_size, k_size): + """ + Calculate decomposed Relative Positional Embeddings from :paper:`mvitv2`. + https://github.com/facebookresearch/mvit/blob/19786631e330df9f3622e5402b4a419a263a2c80/mvit/models/attention.py # noqa B950 + Args: + attn (Tensor): attention map. + q (Tensor): query q in the attention layer with shape (B, q_h * q_w, C). + rel_pos_h (Tensor): relative position embeddings (Lh, C) for height axis. + rel_pos_w (Tensor): relative position embeddings (Lw, C) for width axis. + q_size (Tuple): spatial sequence size of query q with (q_h, q_w). + k_size (Tuple): spatial sequence size of key k with (k_h, k_w). + + Returns: + attn (Tensor): attention map with added relative positional embeddings. + """ + q_h, q_w = q_size + k_h, k_w = k_size + Rh = get_rel_pos(q_h, k_h, rel_pos_h) + Rw = get_rel_pos(q_w, k_w, rel_pos_w) + + B, _, dim = q.shape + r_q = q.reshape(B, q_h, q_w, dim) + rel_h = torch.einsum("bhwc,hkc->bhwk", r_q, Rh) + rel_w = torch.einsum("bhwc,wkc->bhwk", r_q, Rw) + + attn = (attn.view(B, q_h, q_w, k_h, k_w) + rel_h[:, :, :, :, None] + rel_w[:, :, :, None, :]).view( + B, q_h * q_w, k_h * k_w + ) + + return attn + + +def mean_flat(tensor): + return tensor.mean(dim=list(range(1, tensor.ndim))) + + +################################################################################# +# Token Masking and Unmasking # +################################################################################# +def get_mask(batch, length, mask_ratio, device, mask_type=None, data_info=None, extra_len=0): + """ + Get the binary mask for the input sequence. + Args: + - batch: batch size + - length: sequence length + - mask_ratio: ratio of tokens to mask + - data_info: dictionary with info for reconstruction + return: + mask_dict with following keys: + - mask: binary mask, 0 is keep, 1 is remove + - ids_keep: indices of tokens to keep + - ids_restore: indices to restore the original order + """ + assert mask_type in ["random", "fft", "laplacian", "group"] + mask = torch.ones([batch, length], device=device) + len_keep = int(length * (1 - mask_ratio)) - extra_len + + if mask_type == "random" or mask_type == "group": + noise = torch.rand(batch, length, device=device) # noise in [0, 1] + ids_shuffle = torch.argsort(noise, dim=1) # ascend: small is keep, large is remove + ids_restore = torch.argsort(ids_shuffle, dim=1) + # keep the first subset + ids_keep = ids_shuffle[:, :len_keep] + ids_removed = ids_shuffle[:, len_keep:] + + elif mask_type in ["fft", "laplacian"]: + if "strength" in data_info: + strength = data_info["strength"] + + else: + N = data_info["N"][0] + img = data_info["ori_img"] + # 获取原图的尺寸信息 + _, C, H, W = img.shape + if mask_type == "fft": + # 对图片进行reshape,将其变为patch (3, H/N, N, W/N, N) + reshaped_image = img.reshape((batch, -1, H // N, N, W // N, N)) + fft_image = torch.fft.fftn(reshaped_image, dim=(3, 5)) + # 取绝对值并求和获取频率强度 + strength = torch.sum(torch.abs(fft_image), dim=(1, 3, 5)).reshape( + ( + batch, + -1, + ) + ) + elif type == "laplacian": + laplacian_kernel = torch.tensor([[-1, -1, -1], [-1, 8, -1], [-1, -1, -1]], dtype=torch.float32).reshape( + 1, 1, 3, 3 + ) + laplacian_kernel = laplacian_kernel.repeat(C, 1, 1, 1) + # 对图片进行reshape,将其变为patch (3, H/N, N, W/N, N) + reshaped_image = img.reshape(-1, C, H // N, N, W // N, N).permute(0, 2, 4, 1, 3, 5).reshape(-1, C, N, N) + laplacian_response = F.conv2d(reshaped_image, laplacian_kernel, padding=1, groups=C) + strength = laplacian_response.sum(dim=[1, 2, 3]).reshape( + ( + batch, + -1, + ) + ) + + # 对频率强度进行归一化,然后使用torch.multinomial进行采样 + probabilities = strength / (strength.max(dim=1)[0][:, None] + 1e-5) + ids_shuffle = torch.multinomial(probabilities.clip(1e-5, 1), length, replacement=False) + ids_keep = ids_shuffle[:, :len_keep] + ids_restore = torch.argsort(ids_shuffle, dim=1) + ids_removed = ids_shuffle[:, len_keep:] + + mask[:, :len_keep] = 0 + mask = torch.gather(mask, dim=1, index=ids_restore) + + return {"mask": mask, "ids_keep": ids_keep, "ids_restore": ids_restore, "ids_removed": ids_removed} + + +def mask_out_token(x, ids_keep, ids_removed=None): + """ + Mask out the tokens specified by ids_keep. + Args: + - x: input sequence, [N, L, D] + - ids_keep: indices of tokens to keep + return: + - x_masked: masked sequence + """ + N, L, D = x.shape # batch, length, dim + x_remain = torch.gather(x, dim=1, index=ids_keep.unsqueeze(-1).repeat(1, 1, D)) + if ids_removed is not None: + x_masked = torch.gather(x, dim=1, index=ids_removed.unsqueeze(-1).repeat(1, 1, D)) + return x_remain, x_masked + else: + return x_remain + + +def mask_tokens(x, mask_ratio): + """ + Perform per-sample random masking by per-sample shuffling. + Per-sample shuffling is done by argsort random noise. + x: [N, L, D], sequence + """ + N, L, D = x.shape # batch, length, dim + len_keep = int(L * (1 - mask_ratio)) + + noise = torch.rand(N, L, device=x.device) # noise in [0, 1] + + # sort noise for each sample + ids_shuffle = torch.argsort(noise, dim=1) # ascend: small is keep, large is remove + ids_restore = torch.argsort(ids_shuffle, dim=1) + + # keep the first subset + ids_keep = ids_shuffle[:, :len_keep] + x_masked = torch.gather(x, dim=1, index=ids_keep.unsqueeze(-1).repeat(1, 1, D)) + + # generate the binary mask: 0 is keep, 1 is remove + mask = torch.ones([N, L], device=x.device) + mask[:, :len_keep] = 0 + mask = torch.gather(mask, dim=1, index=ids_restore) + + return x_masked, mask, ids_restore + + +def unmask_tokens(x, ids_restore, mask_token): + # x: [N, T, D] if extras == 0 (i.e., no cls token) else x: [N, T+1, D] + mask_tokens = mask_token.repeat(x.shape[0], ids_restore.shape[1] - x.shape[1], 1) + x = torch.cat([x, mask_tokens], dim=1) + x = torch.gather(x, dim=1, index=ids_restore.unsqueeze(-1).repeat(1, 1, x.shape[2])) # unshuffle + return x + + +# Parse 'None' to None and others to float value +def parse_float_none(s): + assert isinstance(s, str) + return None if s == "None" else float(s) + + +# ---------------------------------------------------------------------------- +# Parse a comma separated list of numbers or ranges and return a list of ints. +# Example: '1,2,5-10' returns [1, 2, 5, 6, 7, 8, 9, 10] + + +def parse_int_list(s): + if isinstance(s, list): + return s + ranges = [] + range_re = re.compile(r"^(\d+)-(\d+)$") + for p in s.split(","): + m = range_re.match(p) + if m: + ranges.extend(range(int(m.group(1)), int(m.group(2)) + 1)) + else: + ranges.append(int(p)) + return ranges + + +def init_processes(fn, args): + """Initialize the distributed environment.""" + os.environ["MASTER_ADDR"] = args.master_address + os.environ["MASTER_PORT"] = str(random.randint(2000, 6000)) + print(f'MASTER_ADDR = {os.environ["MASTER_ADDR"]}') + print(f'MASTER_PORT = {os.environ["MASTER_PORT"]}') + torch.cuda.set_device(args.local_rank) + dist.init_process_group(backend="nccl", init_method="env://", rank=args.global_rank, world_size=args.global_size) + fn(args) + if args.global_size > 1: + cleanup() + + +def mprint(*args, **kwargs): + """ + Print only from rank 0. + """ + if dist.get_rank() == 0: + print(*args, **kwargs) + + +def cleanup(): + """ + End DDP training. + """ + dist.barrier() + mprint("Done!") + dist.barrier() + dist.destroy_process_group() + + +# ---------------------------------------------------------------------------- +# logging info. +class Logger: + """ + Redirect stderr to stdout, optionally print stdout to a file, + and optionally force flushing on both stdout and the file. + """ + + def __init__(self, file_name=None, file_mode="w", should_flush=True): + self.file = None + + if file_name is not None: + self.file = open(file_name, file_mode) + + self.should_flush = should_flush + self.stdout = sys.stdout + self.stderr = sys.stderr + + sys.stdout = self + sys.stderr = self + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_value, traceback): + self.close() + + def write(self, text): + """Write text to stdout (and a file) and optionally flush.""" + if len(text) == 0: # workaround for a bug in VSCode debugger: sys.stdout.write(''); sys.stdout.flush() => crash + return + + if self.file is not None: + self.file.write(text) + + self.stdout.write(text) + + if self.should_flush: + self.flush() + + def flush(self): + """Flush written text to both stdout and a file, if open.""" + if self.file is not None: + self.file.flush() + + self.stdout.flush() + + def close(self): + """Flush, close possible files, and remove stdout/stderr mirroring.""" + self.flush() + + # if using multiple loggers, prevent closing in wrong order + if sys.stdout is self: + sys.stdout = self.stdout + if sys.stderr is self: + sys.stderr = self.stderr + + if self.file is not None: + self.file.close() + + +class StackedRandomGenerator: + def __init__(self, device, seeds): + super().__init__() + self.generators = [torch.Generator(device).manual_seed(int(seed) % (1 << 32)) for seed in seeds] + + def randn(self, size, **kwargs): + assert size[0] == len(self.generators) + return torch.stack([torch.randn(size[1:], generator=gen, **kwargs) for gen in self.generators]) + + def randn_like(self, input): + return self.randn(input.shape, dtype=input.dtype, layout=input.layout, device=input.device) + + def randint(self, *args, size, **kwargs): + assert size[0] == len(self.generators) + return torch.stack([torch.randint(*args, size=size[1:], generator=gen, **kwargs) for gen in self.generators]) + + +def prepare_prompt_ar(prompt, ratios, device="cpu", show=True): + # get aspect_ratio or ar + aspect_ratios = re.findall(r"--aspect_ratio\s+(\d+:\d+)", prompt) + ars = re.findall(r"--ar\s+(\d+:\d+)", prompt) + custom_hw = re.findall(r"--hw\s+(\d+:\d+)", prompt) + if show: + print("aspect_ratios:", aspect_ratios, "ars:", ars, "hws:", custom_hw) + prompt_clean = prompt.split("--aspect_ratio")[0].split("--ar")[0].split("--hw")[0] + if len(aspect_ratios) + len(ars) + len(custom_hw) == 0 and show: + print( + "Wrong prompt format. Set to default ar: 1. change your prompt into format '--ar h:w or --hw h:w' for correct generating" + ) + if len(aspect_ratios) != 0: + ar = float(aspect_ratios[0].split(":")[0]) / float(aspect_ratios[0].split(":")[1]) + elif len(ars) != 0: + ar = float(ars[0].split(":")[0]) / float(ars[0].split(":")[1]) + else: + ar = 1.0 + closest_ratio = min(ratios.keys(), key=lambda ratio: abs(float(ratio) - ar)) + if len(custom_hw) != 0: + custom_hw = [float(custom_hw[0].split(":")[0]), float(custom_hw[0].split(":")[1])] + else: + custom_hw = ratios[closest_ratio] + default_hw = ratios[closest_ratio] + prompt_show = f"prompt: {prompt_clean.strip()}\nSize: --ar {closest_ratio}, --bin hw {ratios[closest_ratio]}, --custom hw {custom_hw}" + return ( + prompt_clean, + prompt_show, + torch.tensor(default_hw, device=device)[None], + torch.tensor([float(closest_ratio)], device=device)[None], + torch.tensor(custom_hw, device=device)[None], + ) + + +def resize_and_crop_tensor(samples: torch.Tensor, new_width: int, new_height: int) -> torch.Tensor: + orig_height, orig_width = samples.shape[2], samples.shape[3] + + # Check if resizing is needed + if orig_height != new_height or orig_width != new_width: + ratio = max(new_height / orig_height, new_width / orig_width) + resized_width = int(orig_width * ratio) + resized_height = int(orig_height * ratio) + + # Resize + samples = F.interpolate(samples, size=(resized_height, resized_width), mode="bilinear", align_corners=False) + + # Center Crop + start_x = (resized_width - new_width) // 2 + end_x = start_x + new_width + start_y = (resized_height - new_height) // 2 + end_y = start_y + new_height + samples = samples[:, :, start_y:end_y, start_x:end_x] + + return samples + + +def resize_and_crop_img(img: Image, new_width, new_height): + orig_width, orig_height = img.size + + ratio = max(new_width / orig_width, new_height / orig_height) + resized_width = int(orig_width * ratio) + resized_height = int(orig_height * ratio) + + img = img.resize((resized_width, resized_height), Image.LANCZOS) + + left = (resized_width - new_width) / 2 + top = (resized_height - new_height) / 2 + right = (resized_width + new_width) / 2 + bottom = (resized_height + new_height) / 2 + + img = img.crop((left, top, right, bottom)) + + return img + + +def mask_feature(emb, mask): + if emb.shape[0] == 1: + keep_index = mask.sum().item() + return emb[:, :, :keep_index, :], keep_index + else: + masked_feature = emb * mask[:, None, :, None] + return masked_feature, emb.shape[2] + + +def val2list(x: list or tuple or any, repeat_time=1) -> list: # type: ignore + """Repeat `val` for `repeat_time` times and return the list or val if list/tuple.""" + if isinstance(x, (list, tuple)): + return list(x) + return [x for _ in range(repeat_time)] + + +def val2tuple(x: list or tuple or any, min_len: int = 1, idx_repeat: int = -1) -> tuple: # type: ignore + """Return tuple with min_len by repeating element at idx_repeat.""" + # convert to list first + x = val2list(x) + + # repeat elements if necessary + if len(x) > 0: + x[idx_repeat:idx_repeat] = [x[idx_repeat] for _ in range(min_len - len(x))] + + return tuple(x) + + +def get_same_padding(kernel_size: int or tuple[int, ...]) -> int or tuple[int, ...]: + if isinstance(kernel_size, tuple): + return tuple([get_same_padding(ks) for ks in kernel_size]) + else: + assert kernel_size % 2 > 0, f"kernel size {kernel_size} should be odd number" + return kernel_size // 2 diff --git a/Sana/nodes.py b/Sana/nodes.py new file mode 100644 index 0000000..f400962 --- /dev/null +++ b/Sana/nodes.py @@ -0,0 +1,223 @@ +import os +import json +import torch +import folder_paths + +from comfy.model_management import get_torch_device, soft_empty_cache, text_encoder_offload_device +from comfy import utils +from .conf import sana_conf, sana_res +from .loader import load_sana +from ..utils.dtype import string_to_dtype + +dtypes = [ + "auto", + "FP32", + "FP16", + "BF16" +] + +class SanaCheckpointLoader: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "ckpt_name": (folder_paths.get_filename_list("checkpoints"),), + "model": (list(sana_conf.keys()),), + "dtype": (dtypes,), + } + } + RETURN_TYPES = ("MODEL",) + RETURN_NAMES = ("model",) + FUNCTION = "load_checkpoint" + CATEGORY = "ExtraModels/Sana" + TITLE = "Sana Checkpoint Loader" + + def load_checkpoint(self, ckpt_name, model, dtype): + ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) + model_conf = sana_conf[model] + model = load_sana( + model_path = ckpt_path, + model_conf = model_conf, + dtype = string_to_dtype(dtype, "text_encoder") + ) + return (model,) + + +class SanaResolutionSelect(): + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": (list(sana_res.keys()),), + "ratio": (list(sana_res["1024px"].keys()),{"default":"1.00"}), + } + } + RETURN_TYPES = ("INT","INT") + RETURN_NAMES = ("width","height") + FUNCTION = "get_res" + CATEGORY = "ExtraModels/Sana" + TITLE = "Sana Resolution Select" + + def get_res(self, model, ratio): + width, height = sana_res[model][ratio] + return (width,height) + + +class SanaResolutionCond: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "cond": ("CONDITIONING", ), + "width": ("INT", {"default": 1024.0, "min": 0, "max": 8192}), + "height": ("INT", {"default": 1024.0, "min": 0, "max": 8192}), + } + } + + RETURN_TYPES = ("CONDITIONING",) + RETURN_NAMES = ("cond",) + FUNCTION = "add_cond" + CATEGORY = "ExtraModels/Sana" + TITLE = "Sana Resolution Conditioning" + + def add_cond(self, cond, width, height): + for c in range(len(cond)): + cond[c][1].update({ + "img_hw": [[height, width]], + "aspect_ratio": [[height/width]], + }) + return (cond,) + + +class SanaTextEncode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "text": ("STRING", {"multiline": True}), + "preset_styles": (STYLE_NAMES,), + "GEMMA": ("GEMMA",), + } + } + + RETURN_TYPES = ("CONDITIONING",) + FUNCTION = "encode" + CATEGORY = "ExtraModels/Sana" + TITLE = "Sana Text Encode" + + def encode(self, text, preset_styles, GEMMA=None): + tokenizer = GEMMA["tokenizer"] + text_encoder = GEMMA["text_encoder"] + + # 应用预设样式 - 只使用正面提示词部分 + text, _ = apply_style(preset_styles, text) + + with torch.no_grad(): + # 处理正面提示词 + chi_prompt = "\n".join(preset_te_prompt) + full_prompt = chi_prompt + text + num_chi_tokens = len(tokenizer.encode(chi_prompt)) + max_length = num_chi_tokens + 300 - 2 # 减去[bos]和[_]标记 + + tokens = tokenizer( + [full_prompt], + max_length=max_length, + padding="max_length", + truncation=True, + return_tensors="pt" + ).to(text_encoder.device) + + select_idx = [0] + list(range(-300 + 1, 0)) + embs = text_encoder(tokens.input_ids, tokens.attention_mask)[0][:, None][:, :, select_idx] + emb_masks = tokens.attention_mask[:, select_idx] + # 利用emb_masks将有效的embs选出来,其他置零 + embs = embs * emb_masks.unsqueeze(-1) + # import IPython + # IPython.embed() + + return ([[embs, {}]], ) + +# 需要添加style相关的辅助函数 +style_list = [ + { + "name": "(No style)", + "prompt": "{prompt}", + "negative_prompt": "", + }, + { + "name": "Cinematic", + "prompt": "cinematic still {prompt} . emotional, harmonious, vignette, highly detailed, high budget, bokeh, " + "cinemascope, moody, epic, gorgeous, film grain, grainy", + "negative_prompt": "anime, cartoon, graphic, text, painting, crayon, graphite, abstract, glitch, deformed, mutated, ugly, disfigured", + }, + { + "name": "Photographic", + "prompt": "cinematic photo {prompt} . 35mm photograph, film, bokeh, professional, 4k, highly detailed", + "negative_prompt": "drawing, painting, crayon, sketch, graphite, impressionist, noisy, blurry, soft, deformed, ugly", + }, + { + "name": "Anime", + "prompt": "anime artwork {prompt} . anime style, key visual, vibrant, studio anime, highly detailed", + "negative_prompt": "photo, deformed, black and white, realism, disfigured, low contrast", + }, + { + "name": "Manga", + "prompt": "manga style {prompt} . vibrant, high-energy, detailed, iconic, Japanese comic style", + "negative_prompt": "ugly, deformed, noisy, blurry, low contrast, realism, photorealistic, Western comic style", + }, + { + "name": "Digital Art", + "prompt": "concept art {prompt} . digital artwork, illustrative, painterly, matte painting, highly detailed", + "negative_prompt": "photo, photorealistic, realism, ugly", + }, + { + "name": "Pixel art", + "prompt": "pixel-art {prompt} . low-res, blocky, pixel art style, 8-bit graphics", + "negative_prompt": "sloppy, messy, blurry, noisy, highly detailed, ultra textured, photo, realistic", + }, + { + "name": "Fantasy art", + "prompt": "ethereal fantasy concept art of {prompt} . magnificent, celestial, ethereal, painterly, epic, " + "majestic, magical, fantasy art, cover art, dreamy", + "negative_prompt": "photographic, realistic, realism, 35mm film, dslr, cropped, frame, text, deformed, " + "glitch, noise, noisy, off-center, deformed, cross-eyed, closed eyes, bad anatomy, ugly, " + "disfigured, sloppy, duplicate, mutated, black and white", + }, + { + "name": "Neonpunk", + "prompt": "neonpunk style {prompt} . cyberpunk, vaporwave, neon, vibes, vibrant, stunningly beautiful, crisp, " + "detailed, sleek, ultramodern, magenta highlights, dark purple shadows, high contrast, cinematic, " + "ultra detailed, intricate, professional", + "negative_prompt": "painting, drawing, illustration, glitch, deformed, mutated, cross-eyed, ugly, disfigured", + }, + { + "name": "3D Model", + "prompt": "professional 3d model {prompt} . octane render, highly detailed, volumetric, dramatic lighting", + "negative_prompt": "ugly, deformed, noisy, low poly, blurry, painting", + }, +] + +styles = {k["name"]: (k["prompt"], k["negative_prompt"]) for k in style_list} +STYLE_NAMES = list(styles.keys()) + +def apply_style(style_name: str, positive: str, negative: str = "") -> tuple[str, str]: + p, n = styles.get(style_name, styles[style_name]) + if not negative: + negative = "" + return p.replace("{prompt}", positive), n + negative + +preset_te_prompt = ['Given a user prompt, generate an "Enhanced prompt" that provides detailed visual descriptions suitable for image generation. Evaluate the level of detail in the user prompt:', '- If the prompt is simple, focus on adding specifics about colors, shapes, sizes, textures, and spatial relationships to create vivid and concrete scenes.', '- If the prompt is already detailed, refine and enhance the existing details slightly without overcomplicating.', 'Here are examples of how to transform or refine prompts:', '- User Prompt: A cat sleeping -> Enhanced: A small, fluffy white cat curled up in a round shape, sleeping peacefully on a warm sunny windowsill, surrounded by pots of blooming red flowers.', '- User Prompt: A busy city street -> Enhanced: A bustling city street scene at dusk, featuring glowing street lamps, a diverse crowd of people in colorful clothing, and a double-decker bus passing by towering glass skyscrapers.', 'Please generate only the enhanced description for the prompt below and avoid including any additional commentary or evaluations:', 'User Prompt: '] + +NODE_CLASS_MAPPINGS = { + "SanaCheckpointLoader" : SanaCheckpointLoader, + "SanaResolutionSelect" : SanaResolutionSelect, + "SanaTextEncode" : SanaTextEncode, + "SanaResolutionCond" : SanaResolutionCond, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "Sana Checkpoint Loader": "SanaCheckpointLoader", + "Sana Resolution Select": "SanaResolutionSelect", + "Sana Text Encoder": "SanaTextEncode", + "Sana Resolution Cond": "SanaResolutionCond", +} diff --git a/VAE/nodes.py b/VAE/nodes.py index b3639ae..00be226 100644 --- a/VAE/nodes.py +++ b/VAE/nodes.py @@ -1,4 +1,6 @@ import folder_paths +import torch +import comfy from .conf import vae_conf from .loader import EXVAE @@ -12,6 +14,8 @@ dtypes = [ "BF16" ] +MAX_RESOLUTION=16384 + class ExtraVAELoader: @classmethod def INPUT_TYPES(s): @@ -33,6 +37,34 @@ class ExtraVAELoader: vae = EXVAE(model_path, model_conf, string_to_dtype(dtype, "vae")) return (vae,) + +class EmptyDCAELatentImage: + def __init__(self): + self.device = comfy.model_management.intermediate_device() + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "width": ("INT", {"default": 512, "min": 16, "max": MAX_RESOLUTION, "step": 8, "tooltip": "The width of the latent images in pixels."}), + "height": ("INT", {"default": 512, "min": 16, "max": MAX_RESOLUTION, "step": 8, "tooltip": "The height of the latent images in pixels."}), + "batch_size": ("INT", {"default": 1, "min": 1, "max": 4096, "tooltip": "The number of latent images in the batch."}) + } + } + RETURN_TYPES = ("LATENT",) + OUTPUT_TOOLTIPS = ("The empty latent image batch.",) + FUNCTION = "generate" + TITLE = "Empty DCAE Latent Image" + + CATEGORY = "latent" + DESCRIPTION = "Create a new batch of empty latent images to be denoised via sampling." + + def generate(self, width, height, batch_size=1): + latent = torch.zeros([batch_size, 32, height // 32, width // 32], device=self.device) + return ({"samples":latent}, ) + + NODE_CLASS_MAPPINGS = { "ExtraVAELoader" : ExtraVAELoader, + "EmptyDCAELatentImage" : EmptyDCAELatentImage, } diff --git a/__init__.py b/__init__.py index b4260bb..1fff84c 100644 --- a/__init__.py +++ b/__init__.py @@ -38,5 +38,14 @@ else: from .utils.nodes import NODE_CLASS_MAPPINGS as Extra_Nodes NODE_CLASS_MAPPINGS.update(Extra_Nodes) + # Sana + from .Sana.nodes import NODE_CLASS_MAPPINGS as Sana_Nodes + NODE_CLASS_MAPPINGS.update(Sana_Nodes) + + # Gemma + from .Gemma.nodes import NODE_CLASS_MAPPINGS as Gemma_Nodes + NODE_CLASS_MAPPINGS.update(Gemma_Nodes) + NODE_DISPLAY_NAME_MAPPINGS = {k:v.TITLE for k,v in NODE_CLASS_MAPPINGS.items()} __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] + From 81f58d87cdc69a24ad425ac44da7bc81ee4682d7 Mon Sep 17 00:00:00 2001 From: junsong Date: Wed, 4 Dec 2024 03:25:35 -0800 Subject: [PATCH 3/4] fix the cross-attention num_head bug. --- Gemma/nodes.py | 5 ++--- Sana/conf.py | 4 ++-- Sana/loader.py | 4 ++-- Sana/models/sana.py | 3 ++- Sana/models/sana_multi_scale.py | 17 ++++++++--------- Sana/nodes.py | 6 +----- 6 files changed, 17 insertions(+), 22 deletions(-) diff --git a/Gemma/nodes.py b/Gemma/nodes.py index 55a8a03..55e0883 100644 --- a/Gemma/nodes.py +++ b/Gemma/nodes.py @@ -109,11 +109,10 @@ class GemmaTextEncode: return_tensors="pt" ).to(text_encoder.device) - cond = text_encoder(tokens.input_ids, tokens.attention_mask)[0][:, None] + cond = text_encoder(tokens.input_ids, tokens.attention_mask)[0] emb_masks = tokens.attention_mask - # 利用emb_masks将有效的cond选出来,其他置零 - # cond = cond * emb_masks.unsqueeze(-1) + cond = cond * emb_masks.unsqueeze(-1) return ([[cond, {}]], ) diff --git a/Sana/conf.py b/Sana/conf.py index 7f5a046..3719c0f 100644 --- a/Sana/conf.py +++ b/Sana/conf.py @@ -14,7 +14,7 @@ sana_conf = { "depth": 28, "hidden_size": 1152, "patch_size": 1, - "num_heads": 36, + "num_heads": 16, "linear_head_dim": 32, "model_max_length": 300, "y_norm": True, @@ -36,7 +36,7 @@ sana_conf = { "depth": 20, "hidden_size": 2240, "patch_size": 1, - "num_heads": 70, + "num_heads": 20, "linear_head_dim": 32, "model_max_length": 300, "y_norm": True, diff --git a/Sana/loader.py b/Sana/loader.py index 34ba85f..806ca64 100644 --- a/Sana/loader.py +++ b/Sana/loader.py @@ -48,7 +48,7 @@ class EXM_Sana_Model(comfy.model_base.BaseModel): return out -def load_sana(model_path, model_conf, dtype): +def load_sana(model_path, model_conf): state_dict = comfy.utils.load_torch_file(model_path) state_dict = state_dict.get("model", state_dict) @@ -62,7 +62,7 @@ def load_sana(model_path, model_conf, dtype): state_dict = convert_state_dict(state_dict) # Diffusers parameters = comfy.utils.calculate_parameters(state_dict) - unet_dtype = dtype + unet_dtype = comfy.model_management.unet_dtype() load_device = comfy.model_management.get_torch_device() offload_device = comfy.model_management.unet_offload_device() diff --git a/Sana/models/sana.py b/Sana/models/sana.py index 0dd6551..9da2d71 100644 --- a/Sana/models/sana.py +++ b/Sana/models/sana.py @@ -159,7 +159,7 @@ class Sana(nn.Module): caption_channels=2304, pe_interpolation=1.0, config=None, - model_max_length=120, + model_max_length=300, qk_norm=False, y_norm=False, norm_eps=1e-5, @@ -182,6 +182,7 @@ class Sana(nn.Module): self.depth = depth self.use_pe = use_pe self.y_norm = y_norm + self.model_max_length = model_max_length self.fp32_attention = kwargs.get("use_fp32_attention", False) kernel_size = patch_embed_kernel or patch_size diff --git a/Sana/models/sana_multi_scale.py b/Sana/models/sana_multi_scale.py index 7cc3745..2d5452c 100644 --- a/Sana/models/sana_multi_scale.py +++ b/Sana/models/sana_multi_scale.py @@ -296,20 +296,19 @@ class SanaMS(Sana): t = self.t_embedder(timestep) # (N, D) + y_lens = ((y != 0).sum(dim=3) > 0).sum(dim=2).squeeze().tolist() + y_lens = [y_lens[1]] * bs + + mask = torch.zeros((len(y_lens), self.model_max_length), dtype=torch.int).to(x.device) + for i, count in enumerate(y_lens): + mask[i, :count] = 1 + t0 = self.t_block(t) y = self.y_embedder(y, self.training, mask=mask) # (N, D) if self.y_norm: y = self.attention_y_norm(y) - if mask is not None: - if mask.shape[0] != y.shape[0]: - mask = mask.repeat(y.shape[0] // mask.shape[0], 1) - mask = mask.squeeze(1).squeeze(1) - y = y.squeeze(1).masked_select(mask.unsqueeze(-1) != 0).view(1, -1, x.shape[-1]) - y_lens = mask.sum(dim=1).tolist() - else: - y_lens = [y.shape[2]] * y.shape[0] - y = y.squeeze(1).view(1, -1, x.shape[-1]) + y = y.squeeze(1).masked_select(mask.unsqueeze(-1).bool()).view(1, -1, y.shape[-1]) for block in self.blocks: x = auto_grad_checkpoint( diff --git a/Sana/nodes.py b/Sana/nodes.py index f400962..57b5928 100644 --- a/Sana/nodes.py +++ b/Sana/nodes.py @@ -23,7 +23,6 @@ class SanaCheckpointLoader: "required": { "ckpt_name": (folder_paths.get_filename_list("checkpoints"),), "model": (list(sana_conf.keys()),), - "dtype": (dtypes,), } } RETURN_TYPES = ("MODEL",) @@ -32,13 +31,12 @@ class SanaCheckpointLoader: CATEGORY = "ExtraModels/Sana" TITLE = "Sana Checkpoint Loader" - def load_checkpoint(self, ckpt_name, model, dtype): + def load_checkpoint(self, ckpt_name, model): ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) model_conf = sana_conf[model] model = load_sana( model_path = ckpt_path, model_conf = model_conf, - dtype = string_to_dtype(dtype, "text_encoder") ) return (model,) @@ -132,8 +130,6 @@ class SanaTextEncode: emb_masks = tokens.attention_mask[:, select_idx] # 利用emb_masks将有效的embs选出来,其他置零 embs = embs * emb_masks.unsqueeze(-1) - # import IPython - # IPython.embed() return ([[embs, {}]], ) From dde1273804b5d624a43ca3eb36e1e4e0db5c89bf Mon Sep 17 00:00:00 2001 From: junsong Date: Thu, 5 Dec 2024 00:18:56 -0800 Subject: [PATCH 4/4] fix the conversation and run success; --- Gemma/nodes.py | 16 +- Sana/diffusers_convert.py | 44 ++-- Sana/lora.py | 146 ----------- Sana/models/sana.py | 2 - Sana/models/sana_blocks.py | 2 - Sana/models/utils.py | 499 +------------------------------------ Sana/nodes.py | 93 +------ 7 files changed, 33 insertions(+), 769 deletions(-) delete mode 100644 Sana/lora.py diff --git a/Gemma/nodes.py b/Gemma/nodes.py index 55e0883..9781ea1 100644 --- a/Gemma/nodes.py +++ b/Gemma/nodes.py @@ -6,17 +6,11 @@ from ..utils.dtype import string_to_dtype from huggingface_hub import snapshot_download -# 初始化自定义文件夹路径 -os.makedirs( - os.path.join(folder_paths.models_dir, "text_encoders"), - exist_ok=True -) -folder_paths.folder_names_and_paths["text_encoders"] = ( - [ - os.path.join(folder_paths.models_dir, "text_encoders"), - *folder_paths.folder_names_and_paths.get("text_encoders", [[],set()])[0] - ], - folder_paths.supported_pt_extensions +tenc_root = ( + folder_paths.folder_names_and_paths.get( + "text_encoders", + folder_paths.folder_names_and_paths.get("clip", [[], set()]) + ) ) dtypes = [ diff --git a/Sana/diffusers_convert.py b/Sana/diffusers_convert.py index 312ea9d..dd15b95 100644 --- a/Sana/diffusers_convert.py +++ b/Sana/diffusers_convert.py @@ -1,20 +1,8 @@ # For using the diffusers format weights # Based on the original ComfyUI function + -# https://github.com/PixArt-alpha/PixArt-alpha/blob/master/tools/convert_pixart_alpha_to_diffusers.py +# https://github.com/NVlabs/Sana/blob/main/tools/convert_sana_to_diffusers.py import torch -conversion_map_ms = [ # for multi_scale_train (MS) - # Resolution - ("csize_embedder.mlp.0.weight", "adaln_single.emb.resolution_embedder.linear_1.weight"), - ("csize_embedder.mlp.0.bias", "adaln_single.emb.resolution_embedder.linear_1.bias"), - ("csize_embedder.mlp.2.weight", "adaln_single.emb.resolution_embedder.linear_2.weight"), - ("csize_embedder.mlp.2.bias", "adaln_single.emb.resolution_embedder.linear_2.bias"), - # Aspect ratio - ("ar_embedder.mlp.0.weight", "adaln_single.emb.aspect_ratio_embedder.linear_1.weight"), - ("ar_embedder.mlp.0.bias", "adaln_single.emb.aspect_ratio_embedder.linear_1.bias"), - ("ar_embedder.mlp.2.weight", "adaln_single.emb.aspect_ratio_embedder.linear_2.weight"), - ("ar_embedder.mlp.2.bias", "adaln_single.emb.aspect_ratio_embedder.linear_2.bias"), -] def get_depth(state_dict): return sum(key.endswith('.attn1.to_k.bias') for key in state_dict.keys()) @@ -30,7 +18,7 @@ def get_lora_depth(state_dict): return cnt def get_conversion_map(state_dict): - conversion_map = [ # main SD conversion map (PixArt reference, HF Diffusers) + conversion_map = [ # main SD conversion map (Sana reference, HF Diffusers) # Patch embeddings ("x_embedder.proj.weight", "pos_embed.proj.weight"), ("x_embedder.proj.bias", "pos_embed.proj.bias"), @@ -82,10 +70,7 @@ def find_prefix(state_dict, target_key): return prefix def convert_state_dict(state_dict): - if "adaln_single.emb.resolution_embedder.linear_1.weight" in state_dict.keys(): - cmap = get_conversion_map(state_dict) + conversion_map_ms - else: - cmap = get_conversion_map(state_dict) + cmap = get_conversion_map(state_dict) missing = [k for k,v in cmap if v not in state_dict] new_state_dict = {k: state_dict[v] for k,v in cmap if k not in missing} @@ -109,16 +94,17 @@ def convert_state_dict(state_dict): matched += [key('q'), key('k'), key('v')] if len(matched) < len(state_dict): - print(f"PixArt: UNET conversion has leftover keys! ({len(matched)} vs {len(state_dict)})") + print(f"Sana: UNET conversion has leftover keys! ({len(matched)} vs {len(state_dict)})") print(list( set(state_dict.keys()) - set(matched) )) if len(missing) > 0: - print(f"PixArt: UNET conversion has missing keys!") + print(f"Sana: UNET conversion has missing keys!") print(missing) return new_state_dict # Same as above but for LoRA weights: +# TODO: Not used yet, need to support LoRA for Sana def convert_lora_state_dict(state_dict, peft=True): # koyha rep_ak = lambda x: x.replace(".weight", ".lora_down.weight") @@ -137,18 +123,18 @@ def convert_lora_state_dict(state_dict, peft=True): rep_pp = lambda x: x.replace(".", "_")[:-7] + ".alpha" prefix = "lora_transformer_" - t5_marker = "lora_te_encoder" - t5_keys = [] + gemma_marker = "lora_te_encoder" + gemma_keys = [] for key in list(state_dict.keys()): if key.startswith(prefix): state_dict[key[len(prefix):]] = state_dict.pop(key) - elif t5_marker in key: - t5_keys.append(state_dict.pop(key)) - if len(t5_keys) > 0: - print(f"Text Encoder not supported for PixArt LoRA, ignoring {len(t5_keys)} keys") + elif gemma_marker in key: + gemma_keys.append(state_dict.pop(key)) + if len(gemma_keys) > 0: + print(f"Text Encoder not supported for Sana LoRA, ignoring {len(gemma_keys)} keys") cmap = [] - cmap_unet = get_conversion_map(state_dict) + conversion_map_ms # todo: 512 model + cmap_unet = get_conversion_map(state_dict) # todo: 512 model for k, v in cmap_unet: if v.endswith(".weight"): cmap.append((rep_ak(k), rep_ap(v))) @@ -213,11 +199,11 @@ def convert_lora_state_dict(state_dict, peft=True): pass if len(matched) < len(state_dict): - print(f"PixArt: LoRA conversion has leftover keys! ({len(matched)} vs {len(state_dict)})") + print(f"Sana: LoRA conversion has leftover keys! ({len(matched)} vs {len(state_dict)})") print(list( set(state_dict.keys()) - set(matched) )) if len(missing) > 0: - print(f"PixArt: LoRA conversion has missing keys! (probably)") + print(f"Sana: LoRA conversion has missing keys! (probably)") print(missing) return new_state_dict diff --git a/Sana/lora.py b/Sana/lora.py deleted file mode 100644 index fca5931..0000000 --- a/Sana/lora.py +++ /dev/null @@ -1,146 +0,0 @@ -import os -import copy -import json -import torch -import comfy.lora -import comfy.model_management -from comfy.model_patcher import ModelPatcher -from .diffusers_convert import convert_lora_state_dict - -class EXM_PixArt_ModelPatcher(ModelPatcher): - def calculate_weight(self, patches, weight, key): - """ - This is almost the same as the comfy function, but stripped down to just the LoRA patch code. - The problem with the original code is the q/k/v keys being combined into one for the attention. - In the diffusers code, they're treated as separate keys, but in the reference code they're recombined (q+kv|qkv). - This means, for example, that the [1152,1152] weights become [3456,1152] in the state dict. - The issue with this is that the LoRA weights are [128,1152],[1152,128] and become [384,1162],[3456,128] instead. - - This is the best thing I could think of that would fix that, but it's very fragile. - - Check key shape to determine if it needs the fallback logic - - Cut the input into parts based on the shape (undoing the torch.cat) - - Do the matrix multiplication logic - - Recombine them to match the expected shape - """ - for p in patches: - alpha = p[0] - v = p[1] - strength_model = p[2] - if strength_model != 1.0: - weight *= strength_model - - if isinstance(v, list): - v = (self.calculate_weight(v[1:], v[0].clone(), key), ) - - if len(v) == 2: - patch_type = v[0] - v = v[1] - - if patch_type == "lora": - mat1 = comfy.model_management.cast_to_device(v[0], weight.device, torch.float32) - mat2 = comfy.model_management.cast_to_device(v[1], weight.device, torch.float32) - if v[2] is not None: - alpha *= v[2] / mat2.shape[0] - try: - mat1 = mat1.flatten(start_dim=1) - mat2 = mat2.flatten(start_dim=1) - - ch1 = mat1.shape[0] // mat2.shape[1] - ch2 = mat2.shape[0] // mat1.shape[1] - ### Fallback logic for shape mismatch ### - if mat1.shape[0] != mat2.shape[1] and ch1 == ch2 and (mat1.shape[0]/mat2.shape[1])%1 == 0: - mat1 = mat1.chunk(ch1, dim=0) - mat2 = mat2.chunk(ch1, dim=0) - weight += torch.cat( - [alpha * torch.mm(mat1[x], mat2[x]) for x in range(ch1)], - dim=0, - ).reshape(weight.shape).type(weight.dtype) - else: - weight += (alpha * torch.mm(mat1, mat2)).reshape(weight.shape).type(weight.dtype) - except Exception as e: - print("ERROR", key, e) - return weight - - def clone(self): - n = EXM_PixArt_ModelPatcher(self.model, self.load_device, self.offload_device, self.size, self.current_device, weight_inplace_update=self.weight_inplace_update) - n.patches = {} - for k in self.patches: - n.patches[k] = self.patches[k][:] - - n.object_patches = self.object_patches.copy() - n.model_options = copy.deepcopy(self.model_options) - n.model_keys = self.model_keys - return n - -def replace_model_patcher(model): - n = EXM_PixArt_ModelPatcher( - model = model.model, - size = model.size, - load_device = model.load_device, - offload_device = model.offload_device, - weight_inplace_update = model.weight_inplace_update, - ) - n.patches = {} - for k in model.patches: - n.patches[k] = model.patches[k][:] - - n.object_patches = model.object_patches.copy() - n.model_options = copy.deepcopy(model.model_options) - return n - -def find_peft_alpha(path): - def load_json(json_path): - with open(json_path) as f: - data = json.load(f) - alpha = data.get("lora_alpha") - alpha = alpha or data.get("alpha") - if not alpha: - print(" Found config but `lora_alpha` is missing!") - else: - print(f" Found config at {json_path} [alpha:{alpha}]") - return alpha - - # For some weird reason peft doesn't include the alpha in the actual model - print("PixArt: Warning! This is a PEFT LoRA. Trying to find config...") - files = [ - f"{os.path.splitext(path)[0]}.json", - f"{os.path.splitext(path)[0]}.config.json", - os.path.join(os.path.dirname(path),"adapter_config.json"), - ] - for file in files: - if os.path.isfile(file): - return load_json(file) - - print(" Missing config/alpha! assuming alpha of 8. Consider converting it/adding a config json to it.") - return 8.0 - -def load_pixart_lora(model, lora, lora_path, strength): - k_back = lambda x: x.replace(".lora_up.weight", "") - # need to convert the actual weights for this to work. - if any(True for x in lora.keys() if x.endswith("adaln_single.linear.lora_A.weight")): - lora = convert_lora_state_dict(lora, peft=True) - alpha = find_peft_alpha(lora_path) - lora.update({f"{k_back(x)}.alpha":torch.tensor(alpha) for x in lora.keys() if "lora_up" in x}) - else: # OneTrainer - lora = convert_lora_state_dict(lora, peft=False) - - key_map = {k_back(x):f"diffusion_model.{k_back(x)}.weight" for x in lora.keys() if "lora_up" in x} # fake - - loaded = comfy.lora.load_lora(lora, key_map) - if model is not None: - # switch to custom model patcher when using LoRAs - if isinstance(model, EXM_PixArt_ModelPatcher): - new_modelpatcher = model.clone() - else: - new_modelpatcher = replace_model_patcher(model) - k = new_modelpatcher.add_patches(loaded, strength) - else: - k = () - new_modelpatcher = None - - k = set(k) - for x in loaded: - if (x not in k): - print("NOT LOADED", x) - - return new_modelpatcher diff --git a/Sana/models/sana.py b/Sana/models/sana.py index 9da2d71..81c2211 100644 --- a/Sana/models/sana.py +++ b/Sana/models/sana.py @@ -226,8 +226,6 @@ class Sana(nn.Module): ) self.final_layer = T2IFinalLayer(hidden_size, patch_size, self.out_channels) - self.initialize_weights() - def forward(self, x, timestep, y, mask=None, data_info=None, **kwargs): """ Forward pass of Sana. diff --git a/Sana/models/sana_blocks.py b/Sana/models/sana_blocks.py index 31ac821..cdd37e4 100644 --- a/Sana/models/sana_blocks.py +++ b/Sana/models/sana_blocks.py @@ -16,10 +16,8 @@ # This file is modified from https://github.com/PixArt-alpha/PixArt-sigma import math -import os from typing import Optional -import xformers.ops import torch import torch.nn as nn import torch.nn.functional as F diff --git a/Sana/models/utils.py b/Sana/models/utils.py index d74db3b..f6965d2 100644 --- a/Sana/models/utils.py +++ b/Sana/models/utils.py @@ -14,21 +14,12 @@ # # SPDX-License-Identifier: Apache-2.0 -import math -import os -import random -import re -import sys from collections.abc import Iterable from itertools import repeat +from typing import Union, Tuple import torch -import torch.distributed as dist -import torch.nn as nn -import torch.nn.functional as F -from PIL import Image from torch.utils.checkpoint import checkpoint, checkpoint_sequential -from torchvision import transforms as T def _ntuple(n): @@ -44,25 +35,6 @@ to_1tuple = _ntuple(1) to_2tuple = _ntuple(2) -def set_grad_checkpoint(model, gc_step=1): - assert isinstance(model, nn.Module) - - def set_attr(module): - module.grad_checkpointing = True - module.grad_checkpointing_step = gc_step - - model.apply(set_attr) - - -def set_fp32_attention(model): - assert isinstance(model, nn.Module) - - def set_attr(module): - module.fp32_attention = True - - model.apply(set_attr) - - def auto_grad_checkpoint(module, *args, **kwargs): if getattr(module, "grad_checkpointing", False): if isinstance(module, Iterable): @@ -99,478 +71,12 @@ def checkpoint_sequential(functions, step, input, *args, **kwargs): input = checkpoint(run_function(start, end, functions), input, preserve_rng_state=preserve) return run_function(end + 1, len(functions) - 1, functions)(input) - -def window_partition(x, window_size): - """ - Partition into non-overlapping windows with padding if needed. - Args: - x (tensor): input tokens with [B, H, W, C]. - window_size (int): window size. - - Returns: - windows: windows after partition with [B * num_windows, window_size, window_size, C]. - (Hp, Wp): padded height and width before partition - """ - B, H, W, C = x.shape - - pad_h = (window_size - H % window_size) % window_size - pad_w = (window_size - W % window_size) % window_size - if pad_h > 0 or pad_w > 0: - x = F.pad(x, (0, 0, 0, pad_w, 0, pad_h)) - Hp, Wp = H + pad_h, W + pad_w - - x = x.view(B, Hp // window_size, window_size, Wp // window_size, window_size, C) - windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C) - return windows, (Hp, Wp) - - -def window_unpartition(windows, window_size, pad_hw, hw): - """ - Window unpartition into original sequences and removing padding. - Args: - x (tensor): input tokens with [B * num_windows, window_size, window_size, C]. - window_size (int): window size. - pad_hw (Tuple): padded height and width (Hp, Wp). - hw (Tuple): original height and width (H, W) before padding. - - Returns: - x: unpartitioned sequences with [B, H, W, C]. - """ - Hp, Wp = pad_hw - H, W = hw - B = windows.shape[0] // (Hp * Wp // window_size // window_size) - x = windows.view(B, Hp // window_size, Wp // window_size, window_size, window_size, -1) - x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, Hp, Wp, -1) - - if Hp > H or Wp > W: - x = x[:, :H, :W, :].contiguous() - return x - - -def get_rel_pos(q_size, k_size, rel_pos): - """ - Get relative positional embeddings according to the relative positions of - query and key sizes. - Args: - q_size (int): size of query q. - k_size (int): size of key k. - rel_pos (Tensor): relative position embeddings (L, C). - - Returns: - Extracted positional embeddings according to relative positions. - """ - max_rel_dist = int(2 * max(q_size, k_size) - 1) - # Interpolate rel pos if needed. - if rel_pos.shape[0] != max_rel_dist: - # Interpolate rel pos. - rel_pos_resized = F.interpolate( - rel_pos.reshape(1, rel_pos.shape[0], -1).permute(0, 2, 1), - size=max_rel_dist, - mode="linear", - ) - rel_pos_resized = rel_pos_resized.reshape(-1, max_rel_dist).permute(1, 0) - else: - rel_pos_resized = rel_pos - - # Scale the coords with short length if shapes for q and k are different. - q_coords = torch.arange(q_size)[:, None] * max(k_size / q_size, 1.0) - k_coords = torch.arange(k_size)[None, :] * max(q_size / k_size, 1.0) - relative_coords = (q_coords - k_coords) + (k_size - 1) * max(q_size / k_size, 1.0) - - return rel_pos_resized[relative_coords.long()] - - -def add_decomposed_rel_pos(attn, q, rel_pos_h, rel_pos_w, q_size, k_size): - """ - Calculate decomposed Relative Positional Embeddings from :paper:`mvitv2`. - https://github.com/facebookresearch/mvit/blob/19786631e330df9f3622e5402b4a419a263a2c80/mvit/models/attention.py # noqa B950 - Args: - attn (Tensor): attention map. - q (Tensor): query q in the attention layer with shape (B, q_h * q_w, C). - rel_pos_h (Tensor): relative position embeddings (Lh, C) for height axis. - rel_pos_w (Tensor): relative position embeddings (Lw, C) for width axis. - q_size (Tuple): spatial sequence size of query q with (q_h, q_w). - k_size (Tuple): spatial sequence size of key k with (k_h, k_w). - - Returns: - attn (Tensor): attention map with added relative positional embeddings. - """ - q_h, q_w = q_size - k_h, k_w = k_size - Rh = get_rel_pos(q_h, k_h, rel_pos_h) - Rw = get_rel_pos(q_w, k_w, rel_pos_w) - - B, _, dim = q.shape - r_q = q.reshape(B, q_h, q_w, dim) - rel_h = torch.einsum("bhwc,hkc->bhwk", r_q, Rh) - rel_w = torch.einsum("bhwc,wkc->bhwk", r_q, Rw) - - attn = (attn.view(B, q_h, q_w, k_h, k_w) + rel_h[:, :, :, :, None] + rel_w[:, :, :, None, :]).view( - B, q_h * q_w, k_h * k_w - ) - - return attn - - -def mean_flat(tensor): - return tensor.mean(dim=list(range(1, tensor.ndim))) - - -################################################################################# -# Token Masking and Unmasking # -################################################################################# -def get_mask(batch, length, mask_ratio, device, mask_type=None, data_info=None, extra_len=0): - """ - Get the binary mask for the input sequence. - Args: - - batch: batch size - - length: sequence length - - mask_ratio: ratio of tokens to mask - - data_info: dictionary with info for reconstruction - return: - mask_dict with following keys: - - mask: binary mask, 0 is keep, 1 is remove - - ids_keep: indices of tokens to keep - - ids_restore: indices to restore the original order - """ - assert mask_type in ["random", "fft", "laplacian", "group"] - mask = torch.ones([batch, length], device=device) - len_keep = int(length * (1 - mask_ratio)) - extra_len - - if mask_type == "random" or mask_type == "group": - noise = torch.rand(batch, length, device=device) # noise in [0, 1] - ids_shuffle = torch.argsort(noise, dim=1) # ascend: small is keep, large is remove - ids_restore = torch.argsort(ids_shuffle, dim=1) - # keep the first subset - ids_keep = ids_shuffle[:, :len_keep] - ids_removed = ids_shuffle[:, len_keep:] - - elif mask_type in ["fft", "laplacian"]: - if "strength" in data_info: - strength = data_info["strength"] - - else: - N = data_info["N"][0] - img = data_info["ori_img"] - # 获取原图的尺寸信息 - _, C, H, W = img.shape - if mask_type == "fft": - # 对图片进行reshape,将其变为patch (3, H/N, N, W/N, N) - reshaped_image = img.reshape((batch, -1, H // N, N, W // N, N)) - fft_image = torch.fft.fftn(reshaped_image, dim=(3, 5)) - # 取绝对值并求和获取频率强度 - strength = torch.sum(torch.abs(fft_image), dim=(1, 3, 5)).reshape( - ( - batch, - -1, - ) - ) - elif type == "laplacian": - laplacian_kernel = torch.tensor([[-1, -1, -1], [-1, 8, -1], [-1, -1, -1]], dtype=torch.float32).reshape( - 1, 1, 3, 3 - ) - laplacian_kernel = laplacian_kernel.repeat(C, 1, 1, 1) - # 对图片进行reshape,将其变为patch (3, H/N, N, W/N, N) - reshaped_image = img.reshape(-1, C, H // N, N, W // N, N).permute(0, 2, 4, 1, 3, 5).reshape(-1, C, N, N) - laplacian_response = F.conv2d(reshaped_image, laplacian_kernel, padding=1, groups=C) - strength = laplacian_response.sum(dim=[1, 2, 3]).reshape( - ( - batch, - -1, - ) - ) - - # 对频率强度进行归一化,然后使用torch.multinomial进行采样 - probabilities = strength / (strength.max(dim=1)[0][:, None] + 1e-5) - ids_shuffle = torch.multinomial(probabilities.clip(1e-5, 1), length, replacement=False) - ids_keep = ids_shuffle[:, :len_keep] - ids_restore = torch.argsort(ids_shuffle, dim=1) - ids_removed = ids_shuffle[:, len_keep:] - - mask[:, :len_keep] = 0 - mask = torch.gather(mask, dim=1, index=ids_restore) - - return {"mask": mask, "ids_keep": ids_keep, "ids_restore": ids_restore, "ids_removed": ids_removed} - - -def mask_out_token(x, ids_keep, ids_removed=None): - """ - Mask out the tokens specified by ids_keep. - Args: - - x: input sequence, [N, L, D] - - ids_keep: indices of tokens to keep - return: - - x_masked: masked sequence - """ - N, L, D = x.shape # batch, length, dim - x_remain = torch.gather(x, dim=1, index=ids_keep.unsqueeze(-1).repeat(1, 1, D)) - if ids_removed is not None: - x_masked = torch.gather(x, dim=1, index=ids_removed.unsqueeze(-1).repeat(1, 1, D)) - return x_remain, x_masked - else: - return x_remain - - -def mask_tokens(x, mask_ratio): - """ - Perform per-sample random masking by per-sample shuffling. - Per-sample shuffling is done by argsort random noise. - x: [N, L, D], sequence - """ - N, L, D = x.shape # batch, length, dim - len_keep = int(L * (1 - mask_ratio)) - - noise = torch.rand(N, L, device=x.device) # noise in [0, 1] - - # sort noise for each sample - ids_shuffle = torch.argsort(noise, dim=1) # ascend: small is keep, large is remove - ids_restore = torch.argsort(ids_shuffle, dim=1) - - # keep the first subset - ids_keep = ids_shuffle[:, :len_keep] - x_masked = torch.gather(x, dim=1, index=ids_keep.unsqueeze(-1).repeat(1, 1, D)) - - # generate the binary mask: 0 is keep, 1 is remove - mask = torch.ones([N, L], device=x.device) - mask[:, :len_keep] = 0 - mask = torch.gather(mask, dim=1, index=ids_restore) - - return x_masked, mask, ids_restore - - -def unmask_tokens(x, ids_restore, mask_token): - # x: [N, T, D] if extras == 0 (i.e., no cls token) else x: [N, T+1, D] - mask_tokens = mask_token.repeat(x.shape[0], ids_restore.shape[1] - x.shape[1], 1) - x = torch.cat([x, mask_tokens], dim=1) - x = torch.gather(x, dim=1, index=ids_restore.unsqueeze(-1).repeat(1, 1, x.shape[2])) # unshuffle - return x - - -# Parse 'None' to None and others to float value -def parse_float_none(s): - assert isinstance(s, str) - return None if s == "None" else float(s) - - -# ---------------------------------------------------------------------------- -# Parse a comma separated list of numbers or ranges and return a list of ints. -# Example: '1,2,5-10' returns [1, 2, 5, 6, 7, 8, 9, 10] - - -def parse_int_list(s): - if isinstance(s, list): - return s - ranges = [] - range_re = re.compile(r"^(\d+)-(\d+)$") - for p in s.split(","): - m = range_re.match(p) - if m: - ranges.extend(range(int(m.group(1)), int(m.group(2)) + 1)) - else: - ranges.append(int(p)) - return ranges - - -def init_processes(fn, args): - """Initialize the distributed environment.""" - os.environ["MASTER_ADDR"] = args.master_address - os.environ["MASTER_PORT"] = str(random.randint(2000, 6000)) - print(f'MASTER_ADDR = {os.environ["MASTER_ADDR"]}') - print(f'MASTER_PORT = {os.environ["MASTER_PORT"]}') - torch.cuda.set_device(args.local_rank) - dist.init_process_group(backend="nccl", init_method="env://", rank=args.global_rank, world_size=args.global_size) - fn(args) - if args.global_size > 1: - cleanup() - - -def mprint(*args, **kwargs): - """ - Print only from rank 0. - """ - if dist.get_rank() == 0: - print(*args, **kwargs) - - -def cleanup(): - """ - End DDP training. - """ - dist.barrier() - mprint("Done!") - dist.barrier() - dist.destroy_process_group() - - -# ---------------------------------------------------------------------------- -# logging info. -class Logger: - """ - Redirect stderr to stdout, optionally print stdout to a file, - and optionally force flushing on both stdout and the file. - """ - - def __init__(self, file_name=None, file_mode="w", should_flush=True): - self.file = None - - if file_name is not None: - self.file = open(file_name, file_mode) - - self.should_flush = should_flush - self.stdout = sys.stdout - self.stderr = sys.stderr - - sys.stdout = self - sys.stderr = self - - def __enter__(self): - return self - - def __exit__(self, exc_type, exc_value, traceback): - self.close() - - def write(self, text): - """Write text to stdout (and a file) and optionally flush.""" - if len(text) == 0: # workaround for a bug in VSCode debugger: sys.stdout.write(''); sys.stdout.flush() => crash - return - - if self.file is not None: - self.file.write(text) - - self.stdout.write(text) - - if self.should_flush: - self.flush() - - def flush(self): - """Flush written text to both stdout and a file, if open.""" - if self.file is not None: - self.file.flush() - - self.stdout.flush() - - def close(self): - """Flush, close possible files, and remove stdout/stderr mirroring.""" - self.flush() - - # if using multiple loggers, prevent closing in wrong order - if sys.stdout is self: - sys.stdout = self.stdout - if sys.stderr is self: - sys.stderr = self.stderr - - if self.file is not None: - self.file.close() - - -class StackedRandomGenerator: - def __init__(self, device, seeds): - super().__init__() - self.generators = [torch.Generator(device).manual_seed(int(seed) % (1 << 32)) for seed in seeds] - - def randn(self, size, **kwargs): - assert size[0] == len(self.generators) - return torch.stack([torch.randn(size[1:], generator=gen, **kwargs) for gen in self.generators]) - - def randn_like(self, input): - return self.randn(input.shape, dtype=input.dtype, layout=input.layout, device=input.device) - - def randint(self, *args, size, **kwargs): - assert size[0] == len(self.generators) - return torch.stack([torch.randint(*args, size=size[1:], generator=gen, **kwargs) for gen in self.generators]) - - -def prepare_prompt_ar(prompt, ratios, device="cpu", show=True): - # get aspect_ratio or ar - aspect_ratios = re.findall(r"--aspect_ratio\s+(\d+:\d+)", prompt) - ars = re.findall(r"--ar\s+(\d+:\d+)", prompt) - custom_hw = re.findall(r"--hw\s+(\d+:\d+)", prompt) - if show: - print("aspect_ratios:", aspect_ratios, "ars:", ars, "hws:", custom_hw) - prompt_clean = prompt.split("--aspect_ratio")[0].split("--ar")[0].split("--hw")[0] - if len(aspect_ratios) + len(ars) + len(custom_hw) == 0 and show: - print( - "Wrong prompt format. Set to default ar: 1. change your prompt into format '--ar h:w or --hw h:w' for correct generating" - ) - if len(aspect_ratios) != 0: - ar = float(aspect_ratios[0].split(":")[0]) / float(aspect_ratios[0].split(":")[1]) - elif len(ars) != 0: - ar = float(ars[0].split(":")[0]) / float(ars[0].split(":")[1]) - else: - ar = 1.0 - closest_ratio = min(ratios.keys(), key=lambda ratio: abs(float(ratio) - ar)) - if len(custom_hw) != 0: - custom_hw = [float(custom_hw[0].split(":")[0]), float(custom_hw[0].split(":")[1])] - else: - custom_hw = ratios[closest_ratio] - default_hw = ratios[closest_ratio] - prompt_show = f"prompt: {prompt_clean.strip()}\nSize: --ar {closest_ratio}, --bin hw {ratios[closest_ratio]}, --custom hw {custom_hw}" - return ( - prompt_clean, - prompt_show, - torch.tensor(default_hw, device=device)[None], - torch.tensor([float(closest_ratio)], device=device)[None], - torch.tensor(custom_hw, device=device)[None], - ) - - -def resize_and_crop_tensor(samples: torch.Tensor, new_width: int, new_height: int) -> torch.Tensor: - orig_height, orig_width = samples.shape[2], samples.shape[3] - - # Check if resizing is needed - if orig_height != new_height or orig_width != new_width: - ratio = max(new_height / orig_height, new_width / orig_width) - resized_width = int(orig_width * ratio) - resized_height = int(orig_height * ratio) - - # Resize - samples = F.interpolate(samples, size=(resized_height, resized_width), mode="bilinear", align_corners=False) - - # Center Crop - start_x = (resized_width - new_width) // 2 - end_x = start_x + new_width - start_y = (resized_height - new_height) // 2 - end_y = start_y + new_height - samples = samples[:, :, start_y:end_y, start_x:end_x] - - return samples - - -def resize_and_crop_img(img: Image, new_width, new_height): - orig_width, orig_height = img.size - - ratio = max(new_width / orig_width, new_height / orig_height) - resized_width = int(orig_width * ratio) - resized_height = int(orig_height * ratio) - - img = img.resize((resized_width, resized_height), Image.LANCZOS) - - left = (resized_width - new_width) / 2 - top = (resized_height - new_height) / 2 - right = (resized_width + new_width) / 2 - bottom = (resized_height + new_height) / 2 - - img = img.crop((left, top, right, bottom)) - - return img - - -def mask_feature(emb, mask): - if emb.shape[0] == 1: - keep_index = mask.sum().item() - return emb[:, :, :keep_index, :], keep_index - else: - masked_feature = emb * mask[:, None, :, None] - return masked_feature, emb.shape[2] - - def val2list(x: list or tuple or any, repeat_time=1) -> list: # type: ignore """Repeat `val` for `repeat_time` times and return the list or val if list/tuple.""" if isinstance(x, (list, tuple)): return list(x) return [x for _ in range(repeat_time)] - def val2tuple(x: list or tuple or any, min_len: int = 1, idx_repeat: int = -1) -> tuple: # type: ignore """Return tuple with min_len by repeating element at idx_repeat.""" # convert to list first @@ -582,8 +88,7 @@ def val2tuple(x: list or tuple or any, min_len: int = 1, idx_repeat: int = -1) - return tuple(x) - -def get_same_padding(kernel_size: int or tuple[int, ...]) -> int or tuple[int, ...]: +def get_same_padding(kernel_size: Union[int, Tuple[int, ...]]) -> Union[int, Tuple[int, ...]]: if isinstance(kernel_size, tuple): return tuple([get_same_padding(ks) for ks in kernel_size]) else: diff --git a/Sana/nodes.py b/Sana/nodes.py index 57b5928..8291c55 100644 --- a/Sana/nodes.py +++ b/Sana/nodes.py @@ -1,13 +1,8 @@ -import os -import json import torch import folder_paths -from comfy.model_management import get_torch_device, soft_empty_cache, text_encoder_offload_device -from comfy import utils from .conf import sana_conf, sana_res from .loader import load_sana -from ..utils.dtype import string_to_dtype dtypes = [ "auto", @@ -93,7 +88,6 @@ class SanaTextEncode: return { "required": { "text": ("STRING", {"multiline": True}), - "preset_styles": (STYLE_NAMES,), "GEMMA": ("GEMMA",), } } @@ -103,19 +97,15 @@ class SanaTextEncode: CATEGORY = "ExtraModels/Sana" TITLE = "Sana Text Encode" - def encode(self, text, preset_styles, GEMMA=None): + def encode(self, text, GEMMA=None): tokenizer = GEMMA["tokenizer"] text_encoder = GEMMA["text_encoder"] - # 应用预设样式 - 只使用正面提示词部分 - text, _ = apply_style(preset_styles, text) - with torch.no_grad(): - # 处理正面提示词 chi_prompt = "\n".join(preset_te_prompt) full_prompt = chi_prompt + text num_chi_tokens = len(tokenizer.encode(chi_prompt)) - max_length = num_chi_tokens + 300 - 2 # 减去[bos]和[_]标记 + max_length = num_chi_tokens + 300 - 2 tokens = tokenizer( [full_prompt], @@ -128,82 +118,21 @@ class SanaTextEncode: select_idx = [0] + list(range(-300 + 1, 0)) embs = text_encoder(tokens.input_ids, tokens.attention_mask)[0][:, None][:, :, select_idx] emb_masks = tokens.attention_mask[:, select_idx] - # 利用emb_masks将有效的embs选出来,其他置零 embs = embs * emb_masks.unsqueeze(-1) return ([[embs, {}]], ) -# 需要添加style相关的辅助函数 -style_list = [ - { - "name": "(No style)", - "prompt": "{prompt}", - "negative_prompt": "", - }, - { - "name": "Cinematic", - "prompt": "cinematic still {prompt} . emotional, harmonious, vignette, highly detailed, high budget, bokeh, " - "cinemascope, moody, epic, gorgeous, film grain, grainy", - "negative_prompt": "anime, cartoon, graphic, text, painting, crayon, graphite, abstract, glitch, deformed, mutated, ugly, disfigured", - }, - { - "name": "Photographic", - "prompt": "cinematic photo {prompt} . 35mm photograph, film, bokeh, professional, 4k, highly detailed", - "negative_prompt": "drawing, painting, crayon, sketch, graphite, impressionist, noisy, blurry, soft, deformed, ugly", - }, - { - "name": "Anime", - "prompt": "anime artwork {prompt} . anime style, key visual, vibrant, studio anime, highly detailed", - "negative_prompt": "photo, deformed, black and white, realism, disfigured, low contrast", - }, - { - "name": "Manga", - "prompt": "manga style {prompt} . vibrant, high-energy, detailed, iconic, Japanese comic style", - "negative_prompt": "ugly, deformed, noisy, blurry, low contrast, realism, photorealistic, Western comic style", - }, - { - "name": "Digital Art", - "prompt": "concept art {prompt} . digital artwork, illustrative, painterly, matte painting, highly detailed", - "negative_prompt": "photo, photorealistic, realism, ugly", - }, - { - "name": "Pixel art", - "prompt": "pixel-art {prompt} . low-res, blocky, pixel art style, 8-bit graphics", - "negative_prompt": "sloppy, messy, blurry, noisy, highly detailed, ultra textured, photo, realistic", - }, - { - "name": "Fantasy art", - "prompt": "ethereal fantasy concept art of {prompt} . magnificent, celestial, ethereal, painterly, epic, " - "majestic, magical, fantasy art, cover art, dreamy", - "negative_prompt": "photographic, realistic, realism, 35mm film, dslr, cropped, frame, text, deformed, " - "glitch, noise, noisy, off-center, deformed, cross-eyed, closed eyes, bad anatomy, ugly, " - "disfigured, sloppy, duplicate, mutated, black and white", - }, - { - "name": "Neonpunk", - "prompt": "neonpunk style {prompt} . cyberpunk, vaporwave, neon, vibes, vibrant, stunningly beautiful, crisp, " - "detailed, sleek, ultramodern, magenta highlights, dark purple shadows, high contrast, cinematic, " - "ultra detailed, intricate, professional", - "negative_prompt": "painting, drawing, illustration, glitch, deformed, mutated, cross-eyed, ugly, disfigured", - }, - { - "name": "3D Model", - "prompt": "professional 3d model {prompt} . octane render, highly detailed, volumetric, dramatic lighting", - "negative_prompt": "ugly, deformed, noisy, low poly, blurry, painting", - }, +preset_te_prompt = [ + 'Given a user prompt, generate an "Enhanced prompt" that provides detailed visual descriptions suitable for image generation. Evaluate the level of detail in the user prompt:', + '- If the prompt is simple, focus on adding specifics about colors, shapes, sizes, textures, and spatial relationships to create vivid and concrete scenes.', + '- If the prompt is already detailed, refine and enhance the existing details slightly without overcomplicating.', + 'Here are examples of how to transform or refine prompts:', + '- User Prompt: A cat sleeping -> Enhanced: A small, fluffy white cat curled up in a round shape, sleeping peacefully on a warm sunny windowsill, surrounded by pots of blooming red flowers.', + '- User Prompt: A busy city street -> Enhanced: A bustling city street scene at dusk, featuring glowing street lamps, a diverse crowd of people in colorful clothing, and a double-decker bus passing by towering glass skyscrapers.', + 'Please generate only the enhanced description for the prompt below and avoid including any additional commentary or evaluations:', + 'User Prompt: ' ] -styles = {k["name"]: (k["prompt"], k["negative_prompt"]) for k in style_list} -STYLE_NAMES = list(styles.keys()) - -def apply_style(style_name: str, positive: str, negative: str = "") -> tuple[str, str]: - p, n = styles.get(style_name, styles[style_name]) - if not negative: - negative = "" - return p.replace("{prompt}", positive), n + negative - -preset_te_prompt = ['Given a user prompt, generate an "Enhanced prompt" that provides detailed visual descriptions suitable for image generation. Evaluate the level of detail in the user prompt:', '- If the prompt is simple, focus on adding specifics about colors, shapes, sizes, textures, and spatial relationships to create vivid and concrete scenes.', '- If the prompt is already detailed, refine and enhance the existing details slightly without overcomplicating.', 'Here are examples of how to transform or refine prompts:', '- User Prompt: A cat sleeping -> Enhanced: A small, fluffy white cat curled up in a round shape, sleeping peacefully on a warm sunny windowsill, surrounded by pots of blooming red flowers.', '- User Prompt: A busy city street -> Enhanced: A bustling city street scene at dusk, featuring glowing street lamps, a diverse crowd of people in colorful clothing, and a double-decker bus passing by towering glass skyscrapers.', 'Please generate only the enhanced description for the prompt below and avoid including any additional commentary or evaluations:', 'User Prompt: '] - NODE_CLASS_MAPPINGS = { "SanaCheckpointLoader" : SanaCheckpointLoader, "SanaResolutionSelect" : SanaResolutionSelect,