| 454 | |
| 455 | |
| 456 | class BertIntermediate(nn.Module): |
| 457 | def __init__(self, config): |
| 458 | super(BertIntermediate, self).__init__() |
| 459 | self.dense = nn.Linear(config.hidden_size, config.intermediate_size, bias=True) |
| 460 | # self.dense = mpu.ColumnParallelLinear( |
| 461 | # input_size=config.hidden_size, |
| 462 | # output_size=config.intermediate_size, |
| 463 | # bias=True, |
| 464 | # gather_output=False, |
| 465 | # stride=1, |
| 466 | # init_method=normal_init_method(mean=0.0, |
| 467 | # std=config.initializer_range)) |
| 468 | self.intermediate_act_fn = ACT2FN[config.hidden_act] \ |
| 469 | if isinstance(config.hidden_act, str) else config.hidden_act |
| 470 | |
| 471 | def forward(self, hidden_states): |
| 472 | hidden_states = self.dense(hidden_states) |
| 473 | hidden_states = self.intermediate_act_fn(hidden_states) |
| 474 | return hidden_states |
| 475 | |
| 476 | |
| 477 | class BertOutput(nn.Module): |