| 126 | |
| 127 | |
| 128 | class DecoderLayer(nn.Module): |
| 129 | |
| 130 | def __init__( |
| 131 | self, |
| 132 | dim, |
| 133 | num_heads, |
| 134 | mlp_ratio=4.0, |
| 135 | qkv_bias=False, |
| 136 | qk_scale=None, |
| 137 | drop=0.0, |
| 138 | attn_drop=0.0, |
| 139 | drop_path=0.0, |
| 140 | act_layer=nn.GELU, |
| 141 | norm_layer='nn.LayerNorm', |
| 142 | epsilon=1e-6, |
| 143 | ): |
| 144 | super().__init__() |
| 145 | self.norm1 = eval(norm_layer)(dim, eps=epsilon) |
| 146 | self.normkv = eval(norm_layer)(dim, eps=epsilon) |
| 147 | |
| 148 | self.mixer = CrossAttention( |
| 149 | dim, |
| 150 | num_heads=num_heads, |
| 151 | qkv_bias=qkv_bias, |
| 152 | qk_scale=qk_scale, |
| 153 | attn_drop=attn_drop, |
| 154 | proj_drop=drop, |
| 155 | ) |
| 156 | |
| 157 | # NOTE: drop path for stochastic depth, we shall see if this is better than dropout here |
| 158 | self.drop_path = DropPath(drop_path) if drop_path > 0.0 else Identity() |
| 159 | |
| 160 | self.norm2 = eval(norm_layer)(dim, eps=epsilon) |
| 161 | |
| 162 | mlp_hidden_dim = int(dim * mlp_ratio) |
| 163 | self.mlp_ratio = mlp_ratio |
| 164 | self.mlp = Mlp( |
| 165 | in_features=dim, |
| 166 | hidden_features=mlp_hidden_dim, |
| 167 | act_layer=act_layer, |
| 168 | drop=drop, |
| 169 | ) |
| 170 | |
| 171 | def forward(self, q, kv, key_mask=None): |
| 172 | x1 = q + self.drop_path( |
| 173 | self.mixer(self.norm1(q), self.normkv(kv), key_mask)) |
| 174 | x = x1 + self.drop_path(self.mlp(self.norm2(x1))) |
| 175 | return x |
| 176 | |
| 177 | |
| 178 | class CMFFLayer(nn.Module): |