Top query self-attention layer abstract class. Self-attention layer takes input with size [b, s, h] and returns output of the same size.
| 225 | |
| 226 | |
| 227 | class TopQuerySelfAttention(torch.nn.Module): |
| 228 | """Top query self-attention layer abstract class. |
| 229 | |
| 230 | Self-attention layer takes input with size [b, s, h] |
| 231 | and returns output of the same size. |
| 232 | """ |
| 233 | |
| 234 | def __init__( |
| 235 | self, |
| 236 | hidden_size, |
| 237 | num_attention_heads, |
| 238 | layer_number, |
| 239 | fp16=True, |
| 240 | attention_softmax_in_fp32=True, |
| 241 | ): |
| 242 | super(TopQuerySelfAttention, self).__init__() |
| 243 | self.hidden_size = hidden_size |
| 244 | self.num_attention_heads = num_attention_heads |
| 245 | self.fp16 = fp16 |
| 246 | self.attention_softmax_in_fp32 = attention_softmax_in_fp32 |
| 247 | self.layer_number = max(1, layer_number) |
| 248 | |
| 249 | assert self.hidden_size % self.num_attention_heads == 0 |
| 250 | self.hidden_size_per_attention_head = int(self.hidden_size // self.num_attention_heads) |
| 251 | |
| 252 | self.query = torch.nn.Linear(self.hidden_size, self.hidden_size) |
| 253 | self.key = torch.nn.Linear(self.hidden_size, self.hidden_size) |
| 254 | self.value = torch.nn.Linear(self.hidden_size, self.hidden_size) |
| 255 | |
| 256 | self.norm_factor = math.sqrt(self.hidden_size_per_attention_head) |
| 257 | self.softmax = torch.nn.Softmax(dim=-1) |
| 258 | |
| 259 | self.dense = torch.nn.Linear(self.hidden_size, self.hidden_size) |
| 260 | |
| 261 | def forward( |
| 262 | self, |
| 263 | hidden_states, |
| 264 | query_hidden_state, |
| 265 | attention_mask, |
| 266 | layer_past=None, |
| 267 | get_key_value=False, |
| 268 | prompt_length=None, |
| 269 | context_length=None, |
| 270 | ): |
| 271 | |
| 272 | # hidden_states: [sq, b, h] |
| 273 | query_layer = self.query(query_hidden_state) |
| 274 | key_layer = self.key(hidden_states) |
| 275 | value_layer = self.value(hidden_states) |
| 276 | |
| 277 | new_query_layer_shape = query_layer.size()[:-1] + \ |
| 278 | (self.num_attention_heads, |
| 279 | self.hidden_size_per_attention_head) |
| 280 | query_layer = query_layer.view(*new_query_layer_shape) |
| 281 | |
| 282 | new_query_layer_shape = key_layer.size()[:-1] + \ |
| 283 | (self.num_attention_heads, |
| 284 | self.hidden_size_per_attention_head) |