(
self,
hidden_size,
num_attention_heads,
num_layers,
layernorm_epsilon=1e-5,
)
| 658 | """Transformer class.""" |
| 659 | |
| 660 | def __init__( |
| 661 | self, |
| 662 | hidden_size, |
| 663 | num_attention_heads, |
| 664 | num_layers, |
| 665 | layernorm_epsilon=1e-5, |
| 666 | ): |
| 667 | super(Transformer, self).__init__() |
| 668 | self.hidden_size = hidden_size |
| 669 | self.num_attention_heads = num_attention_heads |
| 670 | self.layernorm_epsilon = layernorm_epsilon |
| 671 | # Number of layers: |
| 672 | self.num_layers = num_layers |
| 673 | self.num_unique_layers = None |
| 674 | |
| 675 | ################# |
| 676 | assert self.num_unique_layers is None |
| 677 | ################# |
| 678 | |
| 679 | if self.num_unique_layers is None: |
| 680 | self.num_unique_layers = self.num_layers |
| 681 | assert self.num_layers % self.num_unique_layers == 0, \ |
| 682 | 'number of layers should be divisible by number of unique layers' |
| 683 | |
| 684 | # Transformer layers. |
| 685 | def build_layer(layer_number): |
| 686 | return TransformerLayer(self.hidden_size, self.num_attention_heads, layer_number) |
| 687 | |
| 688 | self.layers = torch.nn.ModuleList( |
| 689 | [build_layer(i + 1) for i in range(self.num_unique_layers)]) |
| 690 | |
| 691 | self.topQueryLayer = TopQueryLayer(self.hidden_size, |
| 692 | self.num_attention_heads, |
| 693 | self.num_unique_layers) |
| 694 | |
| 695 | self.final_layernorm = torch.nn.LayerNorm(self.hidden_size, |
| 696 | eps=self.layernorm_epsilon) |
| 697 | |
| 698 | def _get_layer_index(self, layer_number): |
| 699 | return layer_number % self.num_unique_layers |
nothing calls this directly
no test coverage detected