(self, hidden_size, num_attention_heads, attention_dropout_prob, output_dropout_prob, init_method,
layer_id, hidden_size_per_attention_head=None, output_layer_init_method=None, bias=True, cross_num_multi_query_heads=0, row_parallel_linear_final_bias=True, hooks={},
cross_attn_hidden_size=None, transformer_pointer=None, params_dtype=torch.float, skip_init=False, device=torch.device('cpu'))
| 121 | """Parallel cross-attention layer for Transformer""" |
| 122 | |
| 123 | def __init__(self, hidden_size, num_attention_heads, attention_dropout_prob, output_dropout_prob, init_method, |
| 124 | layer_id, hidden_size_per_attention_head=None, output_layer_init_method=None, bias=True, cross_num_multi_query_heads=0, row_parallel_linear_final_bias=True, hooks={}, |
| 125 | cross_attn_hidden_size=None, transformer_pointer=None, params_dtype=torch.float, skip_init=False, device=torch.device('cpu')): |
| 126 | super().__init__() |
| 127 | # Set output layer initialization if not provided. |
| 128 | if output_layer_init_method is None: |
| 129 | output_layer_init_method = init_method |
| 130 | self.hooks = hooks |
| 131 | self.layer_id = layer_id |
| 132 | self.num_attention_heads = num_attention_heads |
| 133 | self.hidden_size = hidden_size |
| 134 | # Per attention head and per partition values. |
| 135 | world_size = get_model_parallel_world_size() |
| 136 | if hidden_size_per_attention_head is None: |
| 137 | self.hidden_size_per_attention_head = divide(hidden_size, num_attention_heads) |
| 138 | else: |
| 139 | self.hidden_size_per_attention_head = hidden_size_per_attention_head |
| 140 | self.num_attention_heads_per_partition = divide(num_attention_heads, world_size) |
| 141 | self.inner_hidden_size = num_attention_heads * self.hidden_size_per_attention_head |
| 142 | self.hidden_size_per_partition = self.hidden_size_per_attention_head * self.num_attention_heads_per_partition |
| 143 | self.cross_num_multi_query_heads = cross_num_multi_query_heads |
| 144 | # Strided linear layer. |
| 145 | if cross_num_multi_query_heads == 0: |
| 146 | kv_size = 2 * self.inner_hidden_size |
| 147 | else: # multi-query |
| 148 | kv_size = self.hidden_size_per_attention_head * self.cross_num_multi_query_heads * 2 |
| 149 | |
| 150 | self.query = ColumnParallelLinear(hidden_size, self.inner_hidden_size, |
| 151 | gather_output=False, |
| 152 | init_method=init_method, bias=bias, params_dtype=params_dtype, module=self, name="query", skip_init=skip_init, device=device) |
| 153 | if cross_attn_hidden_size is None: |
| 154 | cross_attn_hidden_size = hidden_size |
| 155 | self.cross_attn_hidden_size = cross_attn_hidden_size |
| 156 | self.key_value = ColumnParallelLinear(cross_attn_hidden_size, kv_size, |
| 157 | stride=2, |
| 158 | gather_output=False, |
| 159 | init_method=init_method, bias=bias, params_dtype=params_dtype, module=self, name="key_value", |
| 160 | skip_init=skip_init, device=device) |
| 161 | # Dropout. Note that for a single iteration, this layer will generate |
| 162 | # different outputs on different number of parallel partitions but |
| 163 | # on average it should not be partition dependent. |
| 164 | self.attention_dropout = torch.nn.Dropout(attention_dropout_prob) |
| 165 | |
| 166 | # Output. |
| 167 | self.dense = RowParallelLinear( |
| 168 | self.inner_hidden_size, |
| 169 | hidden_size, |
| 170 | input_is_parallel=True, |
| 171 | init_method=output_layer_init_method, bias=bias, params_dtype=params_dtype, module=self, name="dense",skip_init=skip_init, |
| 172 | device=device, final_bias=row_parallel_linear_final_bias) |
| 173 | self.output_dropout = torch.nn.Dropout(output_dropout_prob) |
| 174 | |
| 175 | object.__setattr__(self, 'transformer', transformer_pointer) |
| 176 | assert transformer_pointer is not None |
| 177 | |
| 178 | def _transpose_for_scores(self, tensor): |
| 179 | """Transpose a 3D tensor [b, s, np*hn] into a 4D tensor with |
nothing calls this directly
no test coverage detected