(
self,
hidden_size,
num_attention_heads,
num_layers,
layernorm_epsilon=1e-5,
)
| 565 | """Transformer class.""" |
| 566 | |
| 567 | def __init__( |
| 568 | self, |
| 569 | hidden_size, |
| 570 | num_attention_heads, |
| 571 | num_layers, |
| 572 | layernorm_epsilon=1e-5, |
| 573 | ): |
| 574 | super(Transformer, self).__init__() |
| 575 | self.hidden_size = hidden_size |
| 576 | self.num_attention_heads = num_attention_heads |
| 577 | self.layernorm_epsilon = layernorm_epsilon |
| 578 | # Number of layers: |
| 579 | self.num_layers = num_layers |
| 580 | self.num_unique_layers = None |
| 581 | |
| 582 | ################# |
| 583 | assert self.num_unique_layers is None |
| 584 | ################# |
| 585 | |
| 586 | if self.num_unique_layers is None: |
| 587 | self.num_unique_layers = self.num_layers |
| 588 | assert self.num_layers % self.num_unique_layers == 0, \ |
| 589 | 'number of layers should be divisible by number of unique layers' |
| 590 | |
| 591 | # Transformer layers. |
| 592 | def build_layer(layer_number): |
| 593 | return TransformerLayer(self.hidden_size, self.num_attention_heads, layer_number) |
| 594 | |
| 595 | self.layers = torch.nn.ModuleList( |
| 596 | [build_layer(i + 1) for i in range(self.num_unique_layers)]) |
| 597 | |
| 598 | self.topQueryLayer = TopQueryLayer(self.hidden_size, |
| 599 | self.num_attention_heads, |
| 600 | self.num_unique_layers) |
| 601 | |
| 602 | self.final_layernorm = torch.nn.LayerNorm(self.hidden_size, |
| 603 | eps=self.layernorm_epsilon) |
| 604 | |
| 605 | def _get_layer_index(self, layer_number): |
| 606 | return layer_number % self.num_unique_layers |
nothing calls this directly
no test coverage detected