r"""forward pass of the model
(self, input_ids,
input_position,
encoder_masks,
init_reset=True,
batch_valid_length=None)
| 386 | # self.load_embedding_from_ckpt(config.load_ckpt_path) |
| 387 | |
| 388 | def construct(self, input_ids, |
| 389 | input_position, |
| 390 | encoder_masks, |
| 391 | init_reset=True, |
| 392 | batch_valid_length=None): |
| 393 | r"""forward pass of the model""" |
| 394 | embed, word_table = self.embedding(input_ids, input_position, init_reset, batch_valid_length) |
| 395 | hidden_state = P.Cast()(embed, self.dtype) |
| 396 | if init_reset is False: |
| 397 | hidden_state = self.reshape_to_2d(hidden_state) |
| 398 | # encoder_mask = self.create_encoder_mask(encoder_masks) |
| 399 | if self.blocks is not None: |
| 400 | for i in range(self.num_layers - 1): |
| 401 | hidden_state, _ = self.blocks[i](hidden_state, encoder_masks, init_reset, batch_valid_length) |
| 402 | if self.is_pipeline: |
| 403 | top_query_hidden_states, _ = self.top_query_embedding(input_position) |
| 404 | top_query_hidden_states = self.reshape_to_2d(top_query_hidden_states) |
| 405 | encoder_output, _ = self.top_query_layer(hidden_state, top_query_hidden_states, |
| 406 | encoder_masks, init_reset, batch_valid_length) |
| 407 | encoder_output = self.layernorm(encoder_output) |
| 408 | else: |
| 409 | hidden_state = self.reshape_to_2d(hidden_state) |
| 410 | encoder_output = self.layernorm(hidden_state) |
| 411 | encoder_output = P.Cast()(encoder_output, self.dtype) |
| 412 | top_query_hidden_states, _ = self.top_query_embedding(input_position) |
| 413 | top_query_hidden_states = self.reshape_to_2d(top_query_hidden_states) |
| 414 | encoder_output, _ = self.top_query_layer(encoder_output, top_query_hidden_states, |
| 415 | encoder_masks, init_reset, batch_valid_length) |
| 416 | |
| 417 | return encoder_output, word_table |
| 418 | |
| 419 | def reshape_to_2d(self, x): |
| 420 | r"""reshape nd tensor to 2d, if n <= 2, keep original shape.""" |
nothing calls this directly
no test coverage detected