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

Class TopQueryLayer

codegeex/paddle/codegeex_model.py:478–560  ·  view source on GitHub ↗

A single top query layer. Top query layer takes input with size [b, s, h] and returns an output of the same size.

Source from the content-addressed store, hash-verified

476
477
478class TopQueryLayer(paddle.nn.Layer):
479 """A single top query layer.
480
481 Top query layer takes input with size [b, s, h] and returns an
482 output of the same size.
483 """
484
485 def __init__(
486 self,
487 hidden_size,
488 num_attention_heads,
489 layer_number,
490 layernorm_epsilon=1e-5,
491 ):
492 super(TopQueryLayer, self).__init__()
493 self.hidden_size = hidden_size
494 self.num_attention_heads = num_attention_heads
495 self.layernorm_epsilon = layernorm_epsilon
496 self.layer_number = layer_number
497
498 # Use FP32 for Layernorm
499 self.input_layernorm = paddle.nn.LayerNorm(self.hidden_size,
500 epsilon=self.layernorm_epsilon)
501
502 # Self attention.
503 self.attention = TopQuerySelfAttention(self.hidden_size,
504 self.num_attention_heads,
505 self.layer_number)
506 # Layernorm on the input data.
507 self.post_attention_layernorm = paddle.nn.LayerNorm(self.hidden_size,
508 epsilon=self.layernorm_epsilon)
509
510 # MLP
511 self.mlp = MLP(self.hidden_size)
512
513 def forward(
514 self,
515 hidden_states,
516 query_hidden_state,
517 attention_mask,
518 layer_past=None,
519 get_key_value=False,
520 prompt_length=None,
521 context_length=None,
522 ):
523 # hidden_states: [b, s, h]
524 # assert query_hidden_state != None
525
526 # Use FP32 for Layernorm
527 # layernorm_output = self.input_layernorm(hidden_states.cast("float32")).cast("float16")
528 layernorm_output = self.input_layernorm(hidden_states)
529
530 # Self attention.
531 attention_output = self.attention(layernorm_output,
532 query_hidden_state,
533 attention_mask,
534 layer_past=layer_past,
535 get_key_value=get_key_value,

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected