(self, hidden_size, output_dropout_prob, init_method, inner_hidden_size=None,
output_layer_init_method=None, layer_id=None, row_parallel_linear_final_bias=True, hooks={}, bias=True, activation_func=gelu, transformer_pointer=None, is_gated_mlp=False, num_experts=1,
params_dtype=torch.float, skip_init=False, device=torch.device('cpu'))
| 200 | |
| 201 | class MLP(torch.nn.Module): |
| 202 | def __init__(self, hidden_size, output_dropout_prob, init_method, inner_hidden_size=None, |
| 203 | output_layer_init_method=None, layer_id=None, row_parallel_linear_final_bias=True, hooks={}, bias=True, activation_func=gelu, transformer_pointer=None, is_gated_mlp=False, num_experts=1, |
| 204 | params_dtype=torch.float, skip_init=False, device=torch.device('cpu')): |
| 205 | super(MLP, self).__init__() |
| 206 | self.layer_id = layer_id |
| 207 | self.activation_func = activation_func |
| 208 | # Set output layer initialization if not provided. |
| 209 | if output_layer_init_method is None: |
| 210 | output_layer_init_method = init_method |
| 211 | self.hooks = hooks |
| 212 | # Project to 4h. |
| 213 | self.hidden_size = hidden_size |
| 214 | if inner_hidden_size is None: |
| 215 | inner_hidden_size = 4 * hidden_size |
| 216 | self.inner_hidden_size = inner_hidden_size |
| 217 | self.dense_h_to_4h = ColumnParallelLinear( |
| 218 | self.hidden_size, |
| 219 | self.inner_hidden_size, |
| 220 | gather_output=False, |
| 221 | init_method=init_method, |
| 222 | bias=bias, |
| 223 | params_dtype=params_dtype, |
| 224 | module=self, |
| 225 | name="dense_h_to_4h", |
| 226 | skip_init=skip_init, |
| 227 | device=device |
| 228 | ) |
| 229 | # Project back to h. |
| 230 | self.dense_4h_to_h = RowParallelLinear( |
| 231 | self.inner_hidden_size, |
| 232 | self.hidden_size, |
| 233 | input_is_parallel=True, |
| 234 | init_method=output_layer_init_method, |
| 235 | bias=bias, |
| 236 | params_dtype=params_dtype, |
| 237 | module=self, |
| 238 | name="dense_4h_to_h", |
| 239 | skip_init=skip_init, |
| 240 | device=device, |
| 241 | final_bias=row_parallel_linear_final_bias |
| 242 | ) |
| 243 | self.is_gated_mlp = is_gated_mlp |
| 244 | if is_gated_mlp: |
| 245 | self.dense_h_to_4h_gate = ColumnParallelLinear( |
| 246 | self.hidden_size, |
| 247 | self.inner_hidden_size, |
| 248 | gather_output=False, |
| 249 | init_method=init_method, |
| 250 | bias=False, |
| 251 | params_dtype=params_dtype, |
| 252 | module=self, |
| 253 | name="dense_h_to_4h_gate", |
| 254 | skip_init=skip_init, |
| 255 | device=device |
| 256 | ) |
| 257 | self.num_experts = num_experts |
| 258 | for i in range(1, num_experts): |
| 259 | self.register_module(f"dense_h_to_4h_{i}", ColumnParallelLinear( |
no test coverage detected