(
self,
hidden_states,
attention_mask,
layer_past=None,
get_key_value=False,
prompt_length=None,
context_length=None,
layer_id=0,
)
| 528 | self.mlp = MLP(self.hidden_size) |
| 529 | |
| 530 | def forward( |
| 531 | self, |
| 532 | hidden_states, |
| 533 | attention_mask, |
| 534 | layer_past=None, |
| 535 | get_key_value=False, |
| 536 | prompt_length=None, |
| 537 | context_length=None, |
| 538 | layer_id=0, |
| 539 | ): |
| 540 | # hidden_states: [b, s, h] |
| 541 | # Use FP32 for Layernorm |
| 542 | # layernorm_output = self.input_layernorm(hidden_states.float()).half() |
| 543 | layernorm_output = self.input_layernorm(hidden_states) |
| 544 | |
| 545 | # Self attention. |
| 546 | attention_output, attention_mask = self.attention(layernorm_output, |
| 547 | attention_mask, |
| 548 | layer_past=layer_past, |
| 549 | get_key_value=get_key_value, |
| 550 | prompt_length=prompt_length, |
| 551 | context_length=context_length, |
| 552 | layer_id=layer_id) |
| 553 | |
| 554 | if get_key_value: |
| 555 | attention_output, presents = attention_output |
| 556 | |
| 557 | # Residual connection. |
| 558 | residual = hidden_states |
| 559 | layernorm_input = attention_output + residual |
| 560 | |
| 561 | # Use FP32 for Layernorm |
| 562 | # layernorm_output = self.post_attention_layernorm(layernorm_input.float()).half() |
| 563 | layernorm_output = self.post_attention_layernorm(layernorm_input) |
| 564 | mlp_output = self.mlp(layernorm_output) |
| 565 | output = mlp_output + layernorm_input |
| 566 | |
| 567 | if get_key_value: |
| 568 | output = [output, presents] |
| 569 | |
| 570 | return output, attention_mask |
| 571 | |
| 572 | |
| 573 | class TopQueryLayer(torch.nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected