| 38 | return x * self.weight.view(-1, 1, 1) |
| 39 | |
| 40 | class TransformerStage(nn.Module): |
| 41 | |
| 42 | def __init__(self, fmap_size, window_size, ns_per_pt, |
| 43 | dim_in, dim_embed, depths, stage_spec, n_groups, |
| 44 | use_pe, sr_ratio, |
| 45 | heads, heads_q, stride, |
| 46 | offset_range_factor, |
| 47 | dwc_pe, no_off, fixed_pe, |
| 48 | attn_drop, proj_drop, expansion, drop, drop_path_rate, |
| 49 | use_dwc_mlp, ksize, nat_ksize, |
| 50 | k_qna, nq_qna, qna_activation, |
| 51 | layer_scale_value, use_lpu, log_cpb): |
| 52 | |
| 53 | super().__init__() |
| 54 | fmap_size = to_2tuple(fmap_size) |
| 55 | self.depths = depths |
| 56 | hc = dim_embed // heads |
| 57 | assert dim_embed == heads * hc |
| 58 | self.proj = nn.Conv2d(dim_in, dim_embed, 1, 1, 0) if dim_in != dim_embed else nn.Identity() |
| 59 | self.stage_spec = stage_spec |
| 60 | self.use_lpu = use_lpu |
| 61 | |
| 62 | self.ln_cnvnxt = nn.ModuleDict( |
| 63 | {str(d): LayerNormProxy(dim_embed) for d in range(depths) if stage_spec[d] == 'X'} |
| 64 | ) |
| 65 | self.layer_norms = nn.ModuleList( |
| 66 | [LayerNormProxy(dim_embed) if stage_spec[d // 2] != 'X' else nn.Identity() for d in range(2 * depths)] |
| 67 | ) |
| 68 | |
| 69 | mlp_fn = TransformerMLPWithConv if use_dwc_mlp else TransformerMLP |
| 70 | |
| 71 | self.mlps = nn.ModuleList( |
| 72 | [ |
| 73 | mlp_fn(dim_embed, expansion, drop) for _ in range(depths) |
| 74 | ] |
| 75 | ) |
| 76 | self.attns = nn.ModuleList() |
| 77 | self.drop_path = nn.ModuleList() |
| 78 | self.layer_scales = nn.ModuleList( |
| 79 | [ |
| 80 | LayerScale(dim_embed, init_values=layer_scale_value) if layer_scale_value > 0.0 else nn.Identity() |
| 81 | for _ in range(2 * depths) |
| 82 | ] |
| 83 | ) |
| 84 | self.local_perception_units = nn.ModuleList( |
| 85 | [ |
| 86 | nn.Conv2d(dim_embed, dim_embed, kernel_size=3, stride=1, padding=1, groups=dim_embed) if use_lpu else nn.Identity() |
| 87 | for _ in range(depths) |
| 88 | ] |
| 89 | ) |
| 90 | |
| 91 | for i in range(depths): |
| 92 | if stage_spec[i] == 'L': |
| 93 | self.attns.append( |
| 94 | LocalAttention(dim_embed, heads, window_size, attn_drop, proj_drop) |
| 95 | ) |
| 96 | elif stage_spec[i] == 'D': |
| 97 | self.attns.append( |