(self, config: ChatGLMConfig, layer_number, device=None)
| 227 | """ |
| 228 | |
| 229 | def __init__(self, config: ChatGLMConfig, layer_number, device=None): |
| 230 | super(SelfAttention, self).__init__() |
| 231 | self.layer_number = max(1, layer_number) |
| 232 | |
| 233 | self.projection_size = config.kv_channels * config.num_attention_heads |
| 234 | |
| 235 | # Per attention head and per partition values. |
| 236 | self.hidden_size_per_attention_head = self.projection_size // config.num_attention_heads |
| 237 | self.num_attention_heads_per_partition = config.num_attention_heads |
| 238 | |
| 239 | self.multi_query_attention = config.multi_query_attention |
| 240 | self.qkv_hidden_size = 3 * self.projection_size |
| 241 | if self.multi_query_attention: |
| 242 | self.num_multi_query_groups_per_partition = config.multi_query_group_num |
| 243 | self.qkv_hidden_size = ( |
| 244 | self.projection_size + 2 * self.hidden_size_per_attention_head * config.multi_query_group_num |
| 245 | ) |
| 246 | self.query_key_value = nn.Linear(config.hidden_size, self.qkv_hidden_size, |
| 247 | bias=config.add_bias_linear or config.add_qkv_bias, |
| 248 | device=device, **_config_to_kwargs(config) |
| 249 | ) |
| 250 | |
| 251 | self.core_attention = CoreAttention(config, self.layer_number) |
| 252 | |
| 253 | # Output. |
| 254 | self.dense = nn.Linear(self.projection_size, config.hidden_size, bias=config.add_bias_linear, |
| 255 | device=device, **_config_to_kwargs(config) |
| 256 | ) |
| 257 | |
| 258 | def _allocate_memory(self, inference_max_sequence_len, batch_size, device=None, dtype=None): |
| 259 | if self.multi_query_attention: |
nothing calls this directly
no test coverage detected