MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / TopQuerySelfAttention

Class TopQuerySelfAttention

codegeex/oneflow/codegeex_model.py:279–491  ·  view source on GitHub ↗

Top query self-attention layer abstract class. Self-attention layer takes input with size [b, s, h] and returns output of the same size.

Source from the content-addressed store, hash-verified

277
278
279class TopQuerySelfAttention(torch.nn.Module):
280 """Top query self-attention layer abstract class.
281 Self-attention layer takes input with size [b, s, h]
282 and returns output of the same size.
283 """
284
285 def __init__(
286 self,
287 hidden_size,
288 num_attention_heads,
289 layer_number,
290 fp16=True,
291 attention_softmax_in_fp32=True,
292 ):
293 super(TopQuerySelfAttention, self).__init__()
294 self.hidden_size = hidden_size
295 self.num_attention_heads = num_attention_heads
296 self.fp16 = fp16
297 self.attention_softmax_in_fp32 = attention_softmax_in_fp32
298 self.layer_number = max(1, layer_number)
299
300 assert self.hidden_size % self.num_attention_heads == 0
301 self.hidden_size_per_attention_head = int(self.hidden_size // self.num_attention_heads)
302
303 self.query = torch.nn.Linear(self.hidden_size, self.hidden_size)
304 self.key = torch.nn.Linear(self.hidden_size, self.hidden_size)
305 self.value = torch.nn.Linear(self.hidden_size, self.hidden_size)
306
307 self.norm_factor = math.sqrt(self.hidden_size_per_attention_head)
308 self.softmax = torch.nn.Softmax(dim=-1)
309
310 self.dense = torch.nn.Linear(self.hidden_size, self.hidden_size)
311
312 def forward(
313 self,
314 hidden_states,
315 query_hidden_state,
316 attention_mask,
317 layer_past=None,
318 get_key_value=False,
319 prompt_length=None,
320 context_length=None,
321 ):
322
323 # hidden_states: [sq, b, h]
324 if hasattr(torch._C, 'grouped_matmul_bias') and not isinstance(self.query, QuantizedLinear):
325 query_layer, key_layer, value_layer = torch._C.grouped_matmul_bias([query_hidden_state, hidden_states, hidden_states],
326 [self.query.weight, self.key.weight, self.value.weight],
327 [self.query.bias, self.key.bias, self.value.bias])
328 else:
329 query_layer = self.query(query_hidden_state)
330 key_layer = self.key(hidden_states)
331 value_layer = self.value(hidden_states)
332
333 fallback = not hasattr(torch._C, 'fused_multi_head_attention_inference_v2')
334
335 if fallback:
336 if hasattr(torch._C, 'fused_codegeex_qkv_reshape'):

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected