Transformer class.
| 561 | |
| 562 | |
| 563 | class 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] |