(self, init_method, output_layer_init_method)
| 896 | """Transformer class.""" |
| 897 | |
| 898 | def __init__(self, init_method, output_layer_init_method): |
| 899 | super(ParallelTransformer, self).__init__() |
| 900 | args = get_args() |
| 901 | |
| 902 | # Store activation checkpoiting flag. |
| 903 | self.checkpoint_activations = args.checkpoint_activations |
| 904 | self.checkpoint_num_layers = args.checkpoint_num_layers |
| 905 | |
| 906 | # Number of layers: |
| 907 | self.num_layers = args.num_layers |
| 908 | self.num_unique_layers = None |
| 909 | |
| 910 | ################# |
| 911 | assert self.num_unique_layers is None |
| 912 | ################# |
| 913 | |
| 914 | if self.num_unique_layers is None: |
| 915 | self.num_unique_layers = self.num_layers |
| 916 | assert self.num_layers % self.num_unique_layers == 0, \ |
| 917 | 'number of layers should be divisible by number of unique layers' |
| 918 | self.param_sharing_style = 'grouped' |
| 919 | |
| 920 | # Transformer layers. |
| 921 | def build_layer(layer_number): |
| 922 | return ParallelTransformerLayer( |
| 923 | init_method, |
| 924 | output_layer_init_method, layer_number) |
| 925 | |
| 926 | self.layers = torch.nn.ModuleList( |
| 927 | [build_layer(i + 1) for i in range(self.num_unique_layers)]) |
| 928 | |
| 929 | self.topQueryLayer = ParallelTopQueryLayer( |
| 930 | init_method, |
| 931 | output_layer_init_method, self.num_unique_layers) |
| 932 | |
| 933 | # Final layer norm before output. |
| 934 | if hasattr(args, 'ln_fp16'): |
| 935 | self.ln_fp16 = args.ln_fp16 |
| 936 | else: |
| 937 | self.ln_fp16 = False |
| 938 | |
| 939 | self.final_layernorm = LayerNorm( |
| 940 | args.hidden_size, |
| 941 | eps=args.layernorm_epsilon) |
| 942 | |
| 943 | def _get_layer_index(self, layer_number): |
| 944 | if self.param_sharing_style == 'grouped': |
nothing calls this directly
no test coverage detected