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