| 147 | |
| 148 | |
| 149 | class AutoCorrelationLayer(nn.Module): |
| 150 | def __init__(self, correlation, d_model, n_heads, d_keys=None, |
| 151 | d_values=None): |
| 152 | super(AutoCorrelationLayer, self).__init__() |
| 153 | |
| 154 | d_keys = d_keys or (d_model // n_heads) |
| 155 | d_values = d_values or (d_model // n_heads) |
| 156 | |
| 157 | self.inner_correlation = correlation |
| 158 | self.query_projection = nn.Linear(d_model, d_keys * n_heads) |
| 159 | self.key_projection = nn.Linear(d_model, d_keys * n_heads) |
| 160 | self.value_projection = nn.Linear(d_model, d_values * n_heads) |
| 161 | self.out_projection = nn.Linear(d_values * n_heads, d_model) |
| 162 | self.n_heads = n_heads |
| 163 | |
| 164 | def forward(self, queries, keys, values, attn_mask): |
| 165 | B, L, _ = queries.shape |
| 166 | _, S, _ = keys.shape |
| 167 | H = self.n_heads |
| 168 | |
| 169 | queries = self.query_projection(queries).view(B, L, H, -1) |
| 170 | keys = self.key_projection(keys).view(B, S, H, -1) |
| 171 | values = self.value_projection(values).view(B, S, H, -1) |
| 172 | |
| 173 | out, attn = self.inner_correlation( |
| 174 | queries, |
| 175 | keys, |
| 176 | values, |
| 177 | attn_mask |
| 178 | ) |
| 179 | out = out.view(B, L, -1) |
| 180 | |
| 181 | return self.out_projection(out), attn |