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

Class Transformer

codegeex/paddle/codegeex_model.py:563–675  ·  view source on GitHub ↗

Transformer class.

Source from the content-addressed store, hash-verified

561
562
563class Transformer(paddle.nn.Layer):
564 """Transformer class."""
565
566 def __init__(
567 self,
568 hidden_size,
569 num_attention_heads,
570 num_layers,
571 layernorm_epsilon=1e-5,
572 ):
573 super(Transformer, self).__init__()
574 self.hidden_size = hidden_size
575 self.num_attention_heads = num_attention_heads
576 self.layernorm_epsilon = layernorm_epsilon
577 # Number of layers:
578 self.num_layers = num_layers
579 self.num_unique_layers = None
580
581 #################
582 assert self.num_unique_layers is None
583 #################
584
585 if self.num_unique_layers is None:
586 self.num_unique_layers = self.num_layers
587 assert self.num_layers % self.num_unique_layers == 0, \
588 'number of layers should be divisible by number of unique layers'
589
590 # Transformer layers.
591 def build_layer(layer_number):
592 return TransformerLayer(self.hidden_size, self.num_attention_heads, layer_number)
593
594 self.layers = paddle.nn.LayerList(
595 [build_layer(i + 1) for i in range(self.num_unique_layers)])
596
597 self.topQueryLayer = TopQueryLayer(self.hidden_size,
598 self.num_attention_heads,
599 self.num_unique_layers)
600
601 self.final_layernorm = paddle.nn.LayerNorm(self.hidden_size,
602 epsilon=self.layernorm_epsilon)
603
604 def _get_layer_index(self, layer_number):
605 return layer_number % self.num_unique_layers
606
607 def _get_layer(self, layer_number):
608 return self.layers[self._get_layer_index(layer_number)]
609
610 def forward(
611 self,
612 hidden_states,
613 query_hidden_state,
614 attention_mask,
615 layer_past=None,
616 get_key_value=False,
617 prompt_length=None,
618 context_length=None,
619 ):
620 # data format change to avoid explicit tranposes : [b s h] --> [s b h]

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected