(
self,
hidden_states,
attention_mask,
layer_past=None,
get_key_value=False,
prompt_length=None,
context_length=None,
)
| 436 | self.mlp = MLP(self.hidden_size) |
| 437 | |
| 438 | def forward( |
| 439 | self, |
| 440 | hidden_states, |
| 441 | attention_mask, |
| 442 | layer_past=None, |
| 443 | get_key_value=False, |
| 444 | prompt_length=None, |
| 445 | context_length=None, |
| 446 | ): |
| 447 | # hidden_states: [b, s, h] |
| 448 | # Use FP32 for Layernorm |
| 449 | # layernorm_output = self.input_layernorm(hidden_states.float()).half() |
| 450 | layernorm_output = self.input_layernorm(hidden_states) |
| 451 | |
| 452 | # Self attention. |
| 453 | attention_output = self.attention(layernorm_output, |
| 454 | attention_mask, |
| 455 | layer_past=layer_past, |
| 456 | get_key_value=get_key_value, |
| 457 | prompt_length=prompt_length, |
| 458 | context_length=context_length) |
| 459 | |
| 460 | if get_key_value: |
| 461 | attention_output, presents = attention_output |
| 462 | |
| 463 | # Residual connection. |
| 464 | residual = hidden_states |
| 465 | layernorm_input = attention_output + residual |
| 466 | |
| 467 | # Use FP32 for Layernorm |
| 468 | # layernorm_output = self.post_attention_layernorm(layernorm_input.float()).half() |
| 469 | layernorm_output = self.post_attention_layernorm(layernorm_input) |
| 470 | mlp_output = self.mlp(layernorm_output) |
| 471 | output = mlp_output + layernorm_input |
| 472 | |
| 473 | if get_key_value: |
| 474 | output = [output, presents] |
| 475 | |
| 476 | return output |
| 477 | |
| 478 | |
| 479 | class TopQueryLayer(torch.nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected