(
self, dim, dim_out=None, feat_size=None, stride=1, num_heads=4, dim_head=16, r=9,
qk_ratio=1.0, qkv_bias=False)
| 65 | qkv_bias (bool): add bias to q, k, and v projections |
| 66 | """ |
| 67 | def __init__( |
| 68 | self, dim, dim_out=None, feat_size=None, stride=1, num_heads=4, dim_head=16, r=9, |
| 69 | qk_ratio=1.0, qkv_bias=False): |
| 70 | super().__init__() |
| 71 | dim_out = dim_out or dim |
| 72 | assert dim_out % num_heads == 0, ' should be divided by num_heads' |
| 73 | self.dim_qk = dim_head or make_divisible(dim_out * qk_ratio, divisor=8) // num_heads |
| 74 | self.num_heads = num_heads |
| 75 | self.dim_v = dim_out // num_heads |
| 76 | |
| 77 | self.qkv = nn.Conv2d( |
| 78 | dim, |
| 79 | num_heads * self.dim_qk + self.dim_qk + self.dim_v, |
| 80 | kernel_size=1, bias=qkv_bias) |
| 81 | self.norm_q = nn.BatchNorm2d(num_heads * self.dim_qk) |
| 82 | self.norm_v = nn.BatchNorm2d(self.dim_v) |
| 83 | |
| 84 | if r is not None: |
| 85 | # local lambda convolution for pos |
| 86 | self.conv_lambda = nn.Conv3d(1, self.dim_qk, (r, r, 1), padding=(r // 2, r // 2, 0)) |
| 87 | self.pos_emb = None |
| 88 | self.rel_pos_indices = None |
| 89 | else: |
| 90 | # relative pos embedding |
| 91 | assert feat_size is not None |
| 92 | feat_size = to_2tuple(feat_size) |
| 93 | rel_size = [2 * s - 1 for s in feat_size] |
| 94 | self.conv_lambda = None |
| 95 | self.pos_emb = nn.Parameter(torch.zeros(rel_size[0], rel_size[1], self.dim_qk)) |
| 96 | self.register_buffer('rel_pos_indices', rel_pos_indices(feat_size), persistent=False) |
| 97 | |
| 98 | self.pool = nn.AvgPool2d(2, 2) if stride == 2 else nn.Identity() |
| 99 | |
| 100 | self.reset_parameters() |
| 101 | |
| 102 | def reset_parameters(self): |
| 103 | trunc_normal_(self.qkv.weight, std=self.qkv.weight.shape[1] ** -0.5) # fan-in |
nothing calls this directly
no test coverage detected