first run sucessfull with text encoder mask bug not fix;
This commit is contained in:
@@ -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")
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user