| 130 | |
| 131 | |
| 132 | class AutoCorrelationLayer(nn.Module): |
| 133 | def __init__(self, correlation, d_model, n_heads, d_keys=None, |
| 134 | d_values=None): |
| 135 | super(AutoCorrelationLayer, self).__init__() |
| 136 | |
| 137 | d_keys = d_keys or (d_model // n_heads) |
| 138 | d_values = d_values or (d_model // n_heads) |
| 139 | |
| 140 | self.inner_correlation = correlation |
| 141 | self.query_projection = nn.Linear(d_model, d_keys * n_heads) |
| 142 | self.key_projection = nn.Linear(d_model, d_keys * n_heads) |
| 143 | self.value_projection = nn.Linear(d_model, d_values * n_heads) |
| 144 | self.out_projection = nn.Linear(d_values * n_heads, d_model) |
| 145 | self.n_heads = n_heads |
| 146 | |
| 147 | def forward(self, queries, keys, values, attn_mask): |
| 148 | B, L, _ = queries.shape |
| 149 | _, S, _ = keys.shape |
| 150 | H = self.n_heads |
| 151 | |
| 152 | queries = self.query_projection(queries).view(B, L, H, -1) |
| 153 | keys = self.key_projection(keys).view(B, S, H, -1) |
| 154 | values = self.value_projection(values).view(B, S, H, -1) |
| 155 | |
| 156 | out, attn = self.inner_correlation( |
| 157 | queries, |
| 158 | keys, |
| 159 | values, |
| 160 | attn_mask |
| 161 | ) |
| 162 | out = out.view(B, L, -1) |
| 163 | |
| 164 | return self.out_projection(out), attn |