MCPcopy Create free account
hub / github.com/MotrixLab/AiOS / __init__

Method __init__

models/aios/backbones/swin_transformer.py:82–127  ·  view source on GitHub ↗
(self,
                 dim,
                 window_size,
                 num_heads,
                 qkv_bias=True,
                 qk_scale=None,
                 attn_drop=0.,
                 proj_drop=0.)

Source from the content-addressed store, hash-verified

80 proj_drop (float, optional): Dropout ratio of output. Default: 0.0
81 """
82 def __init__(self,
83 dim,
84 window_size,
85 num_heads,
86 qkv_bias=True,
87 qk_scale=None,
88 attn_drop=0.,
89 proj_drop=0.):
90
91 super().__init__()
92 self.dim = dim
93 self.window_size = window_size # Wh, Ww
94 self.num_heads = num_heads
95 head_dim = dim // num_heads
96 self.scale = qk_scale or head_dim**-0.5
97
98 # define a parameter table of relative position bias
99 self.relative_position_bias_table = nn.Parameter(
100 torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1),
101 num_heads)) # 2*Wh-1 * 2*Ww-1, nH
102
103 # get pair-wise relative position index for each token inside the window
104 coords_h = torch.arange(self.window_size[0])
105 coords_w = torch.arange(self.window_size[1])
106 coords = torch.stack(torch.meshgrid([coords_h, coords_w])) # 2, Wh, Ww
107 coords_flatten = torch.flatten(coords, 1) # 2, Wh*Ww
108 relative_coords = coords_flatten[:, :,
109 None] - coords_flatten[:,
110 None, :] # 2, Wh*Ww, Wh*Ww
111 relative_coords = relative_coords.permute(
112 1, 2, 0).contiguous() # Wh*Ww, Wh*Ww, 2
113 relative_coords[:, :,
114 0] += self.window_size[0] - 1 # shift to start from 0
115 relative_coords[:, :, 1] += self.window_size[1] - 1
116 relative_coords[:, :, 0] *= 2 * self.window_size[1] - 1
117 relative_position_index = relative_coords.sum(-1) # Wh*Ww, Wh*Ww
118 self.register_buffer('relative_position_index',
119 relative_position_index)
120
121 self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
122 self.attn_drop = nn.Dropout(attn_drop)
123 self.proj = nn.Linear(dim, dim)
124 self.proj_drop = nn.Dropout(proj_drop)
125
126 trunc_normal_(self.relative_position_bias_table, std=.02)
127 self.softmax = nn.Softmax(dim=-1)
128
129 def forward(self, x, mask=None):
130 """Forward function.

Callers

nothing calls this directly

Calls 1

__init__Method · 0.45

Tested by

no test coverage detected