(
self, q_size, kv_size, n_heads, n_head_channels, n_groups,
attn_drop, proj_drop, stride,
offset_range_factor, use_pe, dwc_pe,
no_off, fixed_pe, ksize, log_cpb
)
| 130 | class DAttentionBaseline(nn.Module): |
| 131 | |
| 132 | def __init__( |
| 133 | self, q_size, kv_size, n_heads, n_head_channels, n_groups, |
| 134 | attn_drop, proj_drop, stride, |
| 135 | offset_range_factor, use_pe, dwc_pe, |
| 136 | no_off, fixed_pe, ksize, log_cpb |
| 137 | ): |
| 138 | |
| 139 | super().__init__() |
| 140 | self.dwc_pe = dwc_pe |
| 141 | self.n_head_channels = n_head_channels |
| 142 | self.scale = self.n_head_channels ** -0.5 |
| 143 | self.n_heads = n_heads |
| 144 | self.q_h, self.q_w = q_size |
| 145 | # self.kv_h, self.kv_w = kv_size |
| 146 | self.kv_h, self.kv_w = self.q_h // stride, self.q_w // stride |
| 147 | self.nc = n_head_channels * n_heads |
| 148 | self.n_groups = n_groups |
| 149 | self.n_group_channels = self.nc // self.n_groups |
| 150 | self.n_group_heads = self.n_heads // self.n_groups |
| 151 | self.use_pe = use_pe |
| 152 | self.fixed_pe = fixed_pe |
| 153 | self.no_off = no_off |
| 154 | self.offset_range_factor = offset_range_factor |
| 155 | self.ksize = ksize |
| 156 | self.log_cpb = log_cpb |
| 157 | self.stride = stride |
| 158 | kk = self.ksize |
| 159 | pad_size = kk // 2 if kk != stride else 0 |
| 160 | |
| 161 | self.conv_offset = nn.Sequential( |
| 162 | nn.Conv2d(self.n_group_channels, self.n_group_channels, kk, stride, pad_size, groups=self.n_group_channels), |
| 163 | LayerNormProxy(self.n_group_channels), |
| 164 | nn.GELU(), |
| 165 | nn.Conv2d(self.n_group_channels, 2, 1, 1, 0, bias=False) |
| 166 | ) |
| 167 | if self.no_off: |
| 168 | for m in self.conv_offset.parameters(): |
| 169 | m.requires_grad_(False) |
| 170 | |
| 171 | self.proj_q = nn.Conv2d( |
| 172 | self.nc, self.nc, |
| 173 | kernel_size=1, stride=1, padding=0 |
| 174 | ) |
| 175 | |
| 176 | self.proj_k = nn.Conv2d( |
| 177 | self.nc, self.nc, |
| 178 | kernel_size=1, stride=1, padding=0 |
| 179 | ) |
| 180 | |
| 181 | self.proj_v = nn.Conv2d( |
| 182 | self.nc, self.nc, |
| 183 | kernel_size=1, stride=1, padding=0 |
| 184 | ) |
| 185 | |
| 186 | self.proj_out = nn.Conv2d( |
| 187 | self.nc, self.nc, |
| 188 | kernel_size=1, stride=1, padding=0 |
| 189 | ) |
nothing calls this directly
no test coverage detected