| """ |
| ART layer in "ICLR 2023: ACCURATE IMAGE RESTORATION WITH ATTENTION RETRACTABLE TRANSFORMER". |
| """ |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from timm.models.layers import DropPath, to_2tuple, trunc_normal_ |
| from pdb import set_trace as stx |
| import numbers |
|
|
| from einops import rearrange |
| import math |
|
|
| NEG_INF = -1000000 |
|
|
|
|
| |
| class Mlp(nn.Module): |
| def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.): |
| super().__init__() |
| out_features = out_features or in_features |
| hidden_features = hidden_features or in_features |
| self.fc1 = nn.Linear(in_features, hidden_features) |
| self.act = act_layer() |
| self.fc2 = nn.Linear(hidden_features, out_features) |
| self.drop = nn.Dropout(drop) |
|
|
| def forward(self, x): |
| x = self.fc1(x) |
| x = self.act(x) |
| x = self.drop(x) |
| x = self.fc2(x) |
| x = self.drop(x) |
| return x |
|
|
|
|
| class DynamicPosBias(nn.Module): |
| def __init__(self, dim, num_heads): |
| super().__init__() |
| self.num_heads = num_heads |
| self.pos_dim = dim // 4 |
| self.pos_proj = nn.Linear(2, self.pos_dim) |
| self.pos1 = nn.Sequential( |
| nn.LayerNorm(self.pos_dim), |
| nn.ReLU(inplace=True), |
| nn.Linear(self.pos_dim, self.pos_dim), |
| ) |
| self.pos2 = nn.Sequential( |
| nn.LayerNorm(self.pos_dim), |
| nn.ReLU(inplace=True), |
| nn.Linear(self.pos_dim, self.pos_dim) |
| ) |
| self.pos3 = nn.Sequential( |
| nn.LayerNorm(self.pos_dim), |
| nn.ReLU(inplace=True), |
| nn.Linear(self.pos_dim, self.num_heads) |
| ) |
|
|
| def forward(self, biases): |
| pos = self.pos3(self.pos2(self.pos1(self.pos_proj(biases)))) |
| return pos |
|
|
| def flops(self, N): |
| flops = N * 2 * self.pos_dim |
| flops += N * self.pos_dim * self.pos_dim |
| flops += N * self.pos_dim * self.pos_dim |
| flops += N * self.pos_dim * self.num_heads |
| return flops |
|
|
|
|
| |
| class Attention(nn.Module): |
| r""" Multi-head self attention module with dynamic position bias. |
| |
| Args: |
| dim (int): Number of input channels. |
| num_heads (int): Number of attention heads. |
| qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True |
| qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set |
| attn_drop (float, optional): Dropout ratio of attention weight. Default: 0.0 |
| proj_drop (float, optional): Dropout ratio of output. Default: 0.0 |
| """ |
|
|
| def __init__(self, dim, num_heads, qkv_bias=True, qk_scale=None, attn_drop=0., proj_drop=0., |
| position_bias=True): |
|
|
| super().__init__() |
| self.dim = dim |
| self.num_heads = num_heads |
| head_dim = dim // num_heads |
| self.scale = qk_scale or head_dim ** -0.5 |
| self.position_bias = position_bias |
| if self.position_bias: |
| self.pos = DynamicPosBias(self.dim // 4, self.num_heads) |
|
|
| self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) |
| self.attn_drop = nn.Dropout(attn_drop) |
| self.proj = nn.Linear(dim, dim) |
| self.proj_drop = nn.Dropout(proj_drop) |
|
|
| self.softmax = nn.Softmax(dim=-1) |
|
|
| def forward(self, x, H, W, mask=None): |
| """ |
| Args: |
| x: input features with shape of (num_groups*B, N, C) |
| mask: (0/-inf) mask with shape of (num_groups, Gh*Gw, Gh*Gw) or None |
| H: height of each group |
| W: width of each group |
| """ |
| |
| group_size = (H, W) |
| B_, N, C = x.shape |
| assert H * W == N |
| qkv = self.qkv(x).reshape(B_, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4).contiguous() |
| q, k, v = qkv[0], qkv[1], qkv[2] |
|
|
| |
| |
|
|
| q = q * self.scale |
| attn = q @ k.transpose(-2, -1).contiguous() |
| |
|
|
| if self.position_bias: |
| |
| position_bias_h = torch.arange(1 - group_size[0], group_size[0], device=attn.device) |
| position_bias_w = torch.arange(1 - group_size[1], group_size[1], device=attn.device) |
| biases = torch.stack(torch.meshgrid([position_bias_h, position_bias_w])) |
| biases = biases.flatten(1).transpose(0, 1).contiguous().float() |
|
|
| |
| coords_h = torch.arange(group_size[0], device=attn.device) |
| coords_w = torch.arange(group_size[1], device=attn.device) |
| coords = torch.stack(torch.meshgrid([coords_h, coords_w])) |
| coords_flatten = torch.flatten(coords, 1) |
| relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :] |
| relative_coords = relative_coords.permute(1, 2, 0).contiguous() |
| relative_coords[:, :, 0] += group_size[0] - 1 |
| relative_coords[:, :, 1] += group_size[1] - 1 |
| relative_coords[:, :, 0] *= 2 * group_size[1] - 1 |
| relative_position_index = relative_coords.sum(-1) |
|
|
| pos = self.pos(biases) |
| |
| relative_position_bias = pos[relative_position_index.view(-1)].view( |
| group_size[0] * group_size[1], group_size[0] * group_size[1], -1) |
| relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous() |
| attn = attn + relative_position_bias.unsqueeze(0) |
|
|
|
|
| if mask is not None: |
| nP = mask.shape[0] |
| attn = attn.view(B_ // nP, nP, self.num_heads, N, N) + mask.unsqueeze(1).unsqueeze(0) |
| attn = attn.view(-1, self.num_heads, N, N) |
| attn = self.softmax(attn) |
| else: |
| attn = self.softmax(attn) |
|
|
| 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 TransformerBlock(nn.Module): |
| r""" ART Transformer Block. |
| |
| Args: |
| dim (int): Number of input channels. |
| num_heads (int): Number of attention heads. |
| window_size: window size of dense attention |
| interval: interval size of sparse attention |
| ds_flag (int): use Dense Attention or Sparse Attention, 0 for DAB and 1 for SAB. |
| mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. |
| qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True |
| qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set. |
| drop (float, optional): Dropout rate. Default: 0.0 |
| attn_drop (float, optional): Attention dropout rate. Default: 0.0 |
| drop_path (float, optional): Stochastic depth rate. Default: 0.0 |
| # act_layer (nn.Module, optional): Activation layer. Default: nn.GELU |
| norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm |
| """ |
|
|
| def __init__(self, |
| dim, |
| num_heads, |
| window_size=7, |
| interval=8, |
| ds_flag=0, |
| mlp_ratio=4., |
| qkv_bias=True, |
| qk_scale=None, |
| drop=0., |
| attn_drop=0., |
| drop_path=0., |
| act_layer=nn.GELU, |
| norm_layer=nn.LayerNorm, |
| pre_norm=True): |
| super().__init__() |
| self.dim = dim |
| self.num_heads = num_heads |
| self.window_size = window_size |
| self.interval = interval |
| self.ds_flag = ds_flag |
| self.mlp_ratio = mlp_ratio |
|
|
| |
|
|
| self.norm1 = norm_layer(dim) |
|
|
| self.attn = Attention( |
| dim, num_heads=num_heads, |
| qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop, |
| position_bias=True) |
| |
| |
|
|
| self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity() |
| self.norm2 = norm_layer(dim) |
| mlp_hidden_dim = int(dim * mlp_ratio) |
| self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop) |
| |
| self.pre_norm = pre_norm |
|
|
| def forward(self, x, x_size): |
| H, W = x_size |
| B, L, C = x.shape |
| assert L == H * W, "input feature has wrong size %d, %d, %d" % (L, H, W) |
|
|
| if min(H, W) <= self.window_size: |
| |
| self.ds_flag = 0 |
| self.window_size = min(H, W) |
|
|
| |
| shortcut = x |
| if self.pre_norm: |
| |
| x = self.norm1(x) |
| x = x.view(B, H, W, C) |
| |
|
|
| |
| size_par = self.interval if self.ds_flag == 1 else self.window_size |
| pad_l = pad_t = 0 |
| pad_r = (size_par - W % size_par) % size_par |
| pad_b = (size_par - H % size_par) % size_par |
| x = F.pad(x, (0, 0, pad_l, pad_r, pad_t, pad_b)) |
| _, Hd, Wd, _ = x.shape |
|
|
| mask = torch.zeros((1, Hd, Wd, 1), device=x.device) |
| if pad_b > 0: |
| mask[:, -pad_b:, :, :] = -1 |
| if pad_r > 0: |
| mask[:, :, -pad_r:, :] = -1 |
|
|
| |
|
|
| |
| if self.ds_flag == 0: |
| G = Gh = Gw = self.window_size |
| x = x.reshape(B, Hd // G, G, Wd // G, G, C).permute(0, 1, 3, 2, 4, 5).contiguous() |
| x = x.reshape(B * Hd * Wd // G ** 2, G ** 2, C) |
| nP = Hd * Wd // G ** 2 |
| |
| if pad_r > 0 or pad_b > 0: |
| mask = mask.reshape(1, Hd // G, G, Wd // G, G, 1).permute(0, 1, 3, 2, 4, 5).contiguous() |
| mask = mask.reshape(nP, 1, G * G) |
| attn_mask = torch.zeros((nP, G * G, G * G), device=x.device) |
| attn_mask = attn_mask.masked_fill(mask < 0, NEG_INF) |
| else: |
| attn_mask = None |
| if self.ds_flag == 1: |
| I, Gh, Gw = self.interval, Hd // self.interval, Wd // self.interval |
| x = x.reshape(B, Gh, I, Gw, I, C).permute(0, 2, 4, 1, 3, 5).contiguous() |
| x = x.reshape(B * I * I, Gh * Gw, C) |
| nP = I ** 2 |
| |
| if pad_r > 0 or pad_b > 0: |
| mask = mask.reshape(1, Gh, I, Gw, I, 1).permute(0, 2, 4, 1, 3, 5).contiguous() |
| mask = mask.reshape(nP, 1, Gh * Gw) |
| attn_mask = torch.zeros((nP, Gh * Gw, Gh * Gw), device=x.device) |
| attn_mask = attn_mask.masked_fill(mask < 0, NEG_INF) |
| else: |
| attn_mask = None |
|
|
| |
| |
| x = self.attn(x, Gh, Gw, mask=attn_mask) |
|
|
| |
| if self.ds_flag == 0: |
| x = x.reshape(B, Hd // G, Wd // G, G, G, C).permute(0, 1, 3, 2, 4, |
| 5).contiguous() |
| else: |
| x = x.reshape(B, I, I, Gh, Gw, C).permute(0, 3, 1, 4, 2, 5).contiguous() |
| x = x.reshape(B, Hd, Wd, C) |
| |
| |
|
|
| |
|
|
| |
| if pad_r > 0 or pad_b > 0: |
| x = x[:, :H, :W, :].contiguous() |
| x = x.view(B, H * W, C) |
|
|
| |
| |
| |
| |
| if self.pre_norm: |
| x = shortcut + self.drop_path(x) |
| x = x + self.drop_path(self.mlp(self.norm2(x))) |
| else: |
| |
| x = self.norm1(shortcut + self.drop_path(x)) |
| x = self.norm2(x + self.drop_path(self.mlp(x))) |
| |
| |
|
|
| |
| |
|
|
| return x |
|
|
|
|
| def extra_repr(self) -> str: |
| return f"dim={self.dim}, num_heads={self.num_heads}, " \ |
| f"window_size={self.window_size}, ds_flag={self.ds_flag}, mlp_ratio={self.mlp_ratio}" |
|
|
|
|
|
|
| class ConvBlock(nn.Module): |
| """ |
| A Convolutional Block that consists of two convolution layers each followed by |
| instance normalization, LeakyReLU activation and dropout. |
| """ |
|
|
| def __init__(self, in_chans: int, out_chans: int, drop_prob: float): |
| """ |
| Args: |
| in_chans: Number of channels in the input. |
| out_chans: Number of channels in the output. |
| drop_prob: Dropout probability. |
| """ |
| super().__init__() |
|
|
| self.in_chans = in_chans |
| self.out_chans = out_chans |
| self.drop_prob = drop_prob |
|
|
| self.layers = nn.Sequential( |
| nn.Conv2d(in_chans, out_chans, kernel_size=3, padding=1, bias=False), |
| nn.InstanceNorm2d(out_chans), |
| nn.LeakyReLU(negative_slope=0.2, inplace=True), |
| nn.Dropout2d(drop_prob), |
| nn.Conv2d(out_chans, out_chans, kernel_size=3, padding=1, bias=False), |
| nn.InstanceNorm2d(out_chans), |
| nn.LeakyReLU(negative_slope=0.2, inplace=True), |
| nn.Dropout2d(drop_prob), |
| ) |
|
|
| def forward(self, image: torch.Tensor) -> torch.Tensor: |
| return self.layers(image) |
|
|
|
|
|
|
| |
| class Cross_TransformerBlock(nn.Module): |
| r""" ART Transformer Block. |
| |
| Args: |
| dim (int): Number of input channels. |
| num_heads (int): Number of attention heads. |
| window_size: window size of dense attention |
| interval: interval size of sparse attention |
| ds_flag (int): use Dense Attention or Sparse Attention, 0 for DAB and 1 for SAB. |
| mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. |
| qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True |
| qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set. |
| drop (float, optional): Dropout rate. Default: 0.0 |
| attn_drop (float, optional): Attention dropout rate. Default: 0.0 |
| drop_path (float, optional): Stochastic depth rate. Default: 0.0 |
| # act_layer (nn.Module, optional): Activation layer. Default: nn.GELU |
| norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm |
| """ |
|
|
| def __init__(self, |
| dim, |
| num_heads, |
| window_size=7, |
| interval=8, |
| ds_flag=0, |
| mlp_ratio=4., |
| qkv_bias=True, |
| qk_scale=None, |
| drop=0., |
| attn_drop=0., |
| drop_path=0., |
| act_layer=nn.GELU, |
| norm_layer=nn.LayerNorm): |
| super().__init__() |
| self.dim = dim |
| self.num_heads = num_heads |
| self.window_size = window_size |
| self.interval = interval |
| self.ds_flag = ds_flag |
| self.mlp_ratio = mlp_ratio |
|
|
| self.norm1 = norm_layer(dim) |
|
|
| self.attn = Attention( |
| dim, num_heads=num_heads, |
| qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop, |
| position_bias=True) |
|
|
| self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity() |
| self.norm2 = norm_layer(dim) |
| mlp_hidden_dim = int(dim * mlp_ratio) |
| self.mlp_x = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop) |
| self.mlp_y = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop) |
|
|
| def forward(self, x, y, x_size): |
| """ |
| x, y: feature maps of two modalities. They are from the same level with same feature size. |
| """ |
| H, W = x_size |
| B, L, C = x.shape |
| assert L == H * W, "input feature has wrong size %d, %d, %d" % (L, H, W) |
|
|
| if min(H, W) <= self.window_size: |
| |
| self.ds_flag = 0 |
| self.window_size = min(H, W) |
|
|
| shortcut_x = x |
| x = self.norm1(x) |
| x = x.view(B, H, W, C) |
|
|
| shortcut_y = y |
| y = self.norm1(y) |
| y = y.view(B, H, W, C) |
|
|
| |
| size_par = self.interval if self.ds_flag == 1 else self.window_size |
| pad_l = pad_t = 0 |
| pad_r = (size_par - W % size_par) % size_par |
| pad_b = (size_par - H % size_par) % size_par |
| x = F.pad(x, (0, 0, pad_l, pad_r, pad_t, pad_b)) |
| y = F.pad(y, (0, 0, pad_l, pad_r, pad_t, pad_b)) |
| _, Hd, Wd, _ = x.shape |
|
|
| mask = torch.zeros((1, Hd, Wd, 1), device=x.device) |
| if pad_b > 0: |
| mask[:, -pad_b:, :, :] = -1 |
| if pad_r > 0: |
| mask[:, :, -pad_r:, :] = -1 |
|
|
| |
|
|
| |
| if self.ds_flag == 0: |
| G = Gh = Gw = self.window_size |
| x = x.reshape(B, Hd // G, G, Wd // G, G, C).permute(0, 1, 3, 2, 4, 5).contiguous() |
| x = x.reshape(B * Hd * Wd // G ** 2, G ** 2, C) |
| y = y.reshape(B, Hd // G, G, Wd // G, G, C).permute(0, 1, 3, 2, 4, 5).contiguous() |
| y = y.reshape(B * Hd * Wd // G ** 2, G ** 2, C) |
| nP = Hd * Wd // G ** 2 |
| |
| if pad_r > 0 or pad_b > 0: |
| mask = mask.reshape(1, Hd // G, G, Wd // G, G, 1).permute(0, 1, 3, 2, 4, 5).contiguous() |
| mask = mask.reshape(nP, 1, G * G) |
| attn_mask = torch.zeros((nP, G * G, G * G), device=x.device) |
| attn_mask = attn_mask.masked_fill(mask < 0, NEG_INF) |
| else: |
| attn_mask = None |
| if self.ds_flag == 1: |
| I, Gh, Gw = self.interval, Hd // self.interval, Wd // self.interval |
| x = x.reshape(B, Gh, I, Gw, I, C).permute(0, 2, 4, 1, 3, 5).contiguous() |
| x = x.reshape(B * I * I, Gh * Gw, C) |
| y = y.reshape(B, Gh, I, Gw, I, C).permute(0, 2, 4, 1, 3, 5).contiguous() |
| y = y.reshape(B * I * I, Gh * Gw, C) |
| nP = I ** 2 |
| |
| if pad_r > 0 or pad_b > 0: |
| mask = mask.reshape(1, Gh, I, Gw, I, 1).permute(0, 2, 4, 1, 3, 5).contiguous() |
| mask = mask.reshape(nP, 1, Gh * Gw) |
| attn_mask = torch.zeros((nP, Gh * Gw, Gh * Gw), device=x.device) |
| attn_mask = attn_mask.masked_fill(mask < 0, NEG_INF) |
| else: |
| attn_mask = None |
|
|
| |
| xy = torch.cat((x, y), 1) |
| xy = self.attn(xy, Gh, 2*Gw, mask=attn_mask) |
| |
|
|
| |
| if self.ds_flag == 0: |
| x = xy[:, :xy.shape[1]//2, :].reshape(B, Hd // G, Wd // G, G, G, C).permute(0, 1, 3, 2, 4, 5).contiguous() |
| y = xy[:, xy.shape[1]//2:, :].reshape(B, Hd // G, Wd // G, G, G, C).permute(0, 1, 3, 2, 4, 5).contiguous() |
| else: |
| x = xy[:, :xy.shape[1]//2, :].reshape(B, I, I, Gh, Gw, C).permute(0, 3, 1, 4, 2, 5).contiguous() |
| y = xy[:, xy.shape[1]//2:, :].reshape(B, I, I, Gh, Gw, C).permute(0, 3, 1, 4, 2, 5).contiguous() |
| |
| x = x.reshape(B, Hd, Wd, C) |
| y = y.reshape(B, Hd, Wd, C) |
|
|
| |
| if pad_r > 0 or pad_b > 0: |
| x = x[:, :H, :W, :].contiguous() |
| y = y[:, :H, :W, :].contiguous() |
| x = x.view(B, H * W, C) |
| y = y.view(B, H * W, C) |
|
|
| |
| x = shortcut_x + self.drop_path(x) |
| x = x + self.drop_path(self.mlp_x(self.norm2(x))) |
| |
| y = shortcut_y + self.drop_path(y) |
| y = y + self.drop_path(self.mlp_y(self.norm2(y))) |
|
|
| return x, y |
|
|
|
|
| def extra_repr(self) -> str: |
| return f"dim={self.dim}, num_heads={self.num_heads}, " \ |
| f"window_size={self.window_size}, ds_flag={self.ds_flag}, mlp_ratio={self.mlp_ratio}" |
|
|
|
|
|
|
|
|
|
|
|
|
| |
| class Cross_TransformerBlock_v2(nn.Module): |
| r""" ART Transformer Block. |
| 将两个模态的特征沿着channel方向concat, 一起取window. 之后把取出来的两个模态中同一个window的所有patches合并,送到transformer处理。 |
| Args: |
| dim (int): Number of input channels. |
| num_heads (int): Number of attention heads. |
| window_size: window size of dense attention |
| interval: interval size of sparse attention |
| ds_flag (int): use Dense Attention or Sparse Attention, 0 for DAB and 1 for SAB. |
| mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. |
| qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True |
| qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set. |
| drop (float, optional): Dropout rate. Default: 0.0 |
| attn_drop (float, optional): Attention dropout rate. Default: 0.0 |
| drop_path (float, optional): Stochastic depth rate. Default: 0.0 |
| # act_layer (nn.Module, optional): Activation layer. Default: nn.GELU |
| norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm |
| """ |
|
|
| def __init__(self, |
| dim, |
| num_heads, |
| window_size=7, |
| interval=8, |
| ds_flag=0, |
| mlp_ratio=4., |
| qkv_bias=True, |
| qk_scale=None, |
| drop=0., |
| attn_drop=0., |
| drop_path=0., |
| act_layer=nn.GELU, |
| norm_layer=nn.LayerNorm): |
| super().__init__() |
| self.dim = dim |
| self.num_heads = num_heads |
| self.window_size = window_size |
| self.interval = interval |
| self.ds_flag = ds_flag |
| self.mlp_ratio = mlp_ratio |
|
|
| self.norm1 = norm_layer(dim) |
|
|
| self.attn = Attention( |
| dim, num_heads=num_heads, |
| qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop, |
| position_bias=True) |
|
|
| self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity() |
| self.norm2 = norm_layer(dim) |
| mlp_hidden_dim = int(dim * mlp_ratio) |
| self.mlp_x = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop) |
| self.mlp_y = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop) |
|
|
| def forward(self, x, y, x_size): |
| """ |
| x, y: feature maps of two modalities. They are from the same level with same feature size. |
| """ |
| H, W = x_size |
| B, L, C = x.shape |
| assert L == H * W, "input feature has wrong size %d, %d, %d" % (L, H, W) |
|
|
| if min(H, W) <= self.window_size: |
| |
| self.ds_flag = 0 |
| self.window_size = min(H, W) |
|
|
| shortcut_x = x |
| x = self.norm1(x) |
| x = x.view(B, H, W, C) |
|
|
| shortcut_y = y |
| y = self.norm1(y) |
| y = y.view(B, H, W, C) |
|
|
| |
| xy = torch.cat((x, y), -1) |
| size_par = self.interval if self.ds_flag == 1 else self.window_size |
| pad_l = pad_t = 0 |
| pad_r = (size_par - W % size_par) % size_par |
| pad_b = (size_par - H % size_par) % size_par |
| xy = F.pad(xy, (0, 0, pad_l, pad_r, pad_t, pad_b)) |
| _, Hd, Wd, _ = x.shape |
|
|
| mask = torch.zeros((1, Hd, Wd, 1), device=x.device) |
| if pad_b > 0: |
| mask[:, -pad_b:, :, :] = -1 |
| if pad_r > 0: |
| mask[:, :, -pad_r:, :] = -1 |
|
|
| |
|
|
| |
| if self.ds_flag == 0: |
| G = Gh = Gw = self.window_size |
| xy = xy.reshape(B, Hd // G, G, Wd // G, G, 2*C).permute(0, 1, 3, 2, 4, 5).contiguous() |
| xy = xy.reshape(B * Hd * Wd // G ** 2, G ** 2, 2*C) |
| nP = Hd * Wd // G ** 2 |
| |
| if pad_r > 0 or pad_b > 0: |
| mask = mask.reshape(1, Hd // G, G, Wd // G, G, 1).permute(0, 1, 3, 2, 4, 5).contiguous() |
| mask = mask.reshape(nP, 1, G * G) |
| attn_mask = torch.zeros((nP, G * G, G * G), device=x.device) |
| attn_mask = attn_mask.masked_fill(mask < 0, NEG_INF) |
| else: |
| attn_mask = None |
| if self.ds_flag == 1: |
| I, Gh, Gw = self.interval, Hd // self.interval, Wd // self.interval |
| xy = xy.reshape(B, Gh, I, Gw, I, 2*C).permute(0, 2, 4, 1, 3, 5).contiguous() |
| xy = xy.reshape(B * I * I, Gh * Gw, 2*C) |
| nP = I ** 2 |
| |
| if pad_r > 0 or pad_b > 0: |
| mask = mask.reshape(1, Gh, I, Gw, I, 1).permute(0, 2, 4, 1, 3, 5).contiguous() |
| mask = mask.reshape(nP, 1, Gh * Gw) |
| attn_mask = torch.zeros((nP, Gh * Gw, Gh * Gw), device=x.device) |
| attn_mask = attn_mask.masked_fill(mask < 0, NEG_INF) |
| else: |
| attn_mask = None |
|
|
| |
| x = xy[:, :, :C] |
| y = xy[:, :, C:] |
| xy = torch.cat((x, y), 1) |
| xy = self.attn(xy, Gh, 2*Gw, mask=attn_mask) |
| |
|
|
| |
| if self.ds_flag == 0: |
| x = xy[:, :xy.shape[1]//2, :].reshape(B, Hd // G, Wd // G, G, G, C).permute(0, 1, 3, 2, 4, 5).contiguous() |
| y = xy[:, xy.shape[1]//2:, :].reshape(B, Hd // G, Wd // G, G, G, C).permute(0, 1, 3, 2, 4, 5).contiguous() |
| else: |
| x = xy[:, :xy.shape[1]//2, :].reshape(B, I, I, Gh, Gw, C).permute(0, 3, 1, 4, 2, 5).contiguous() |
| y = xy[:, xy.shape[1]//2:, :].reshape(B, I, I, Gh, Gw, C).permute(0, 3, 1, 4, 2, 5).contiguous() |
| |
| x = x.reshape(B, Hd, Wd, C) |
| y = y.reshape(B, Hd, Wd, C) |
|
|
| |
| if pad_r > 0 or pad_b > 0: |
| x = x[:, :H, :W, :].contiguous() |
| y = y[:, :H, :W, :].contiguous() |
| x = x.view(B, H * W, C) |
| y = y.view(B, H * W, C) |
|
|
| |
| x = shortcut_x + self.drop_path(x) |
| x = x + self.drop_path(self.mlp_x(self.norm2(x))) |
| |
| y = shortcut_y + self.drop_path(y) |
| y = y + self.drop_path(self.mlp_y(self.norm2(y))) |
|
|
| return x, y |
|
|
|
|
| def extra_repr(self) -> str: |
| return f"dim={self.dim}, num_heads={self.num_heads}, " \ |
| f"window_size={self.window_size}, ds_flag={self.ds_flag}, mlp_ratio={self.mlp_ratio}" |
|
|
|
|
|
|
|
|
|
|