evaluation net
(self, input_ids, current_index, init_reset=True, batch_valid_length=None)
| 573 | self.not_equal = P.NotEqual().shard(((1, 1), ())) |
| 574 | |
| 575 | def construct(self, input_ids, current_index, init_reset=True, batch_valid_length=None): |
| 576 | """evaluation net""" |
| 577 | # input_mask = F.cast(F.not_equal(input_ids, self.pad_token), mstype.float32) |
| 578 | input_mask = F.cast(self.not_equal(input_ids, self.pad_token), mstype.float32) |
| 579 | bs, seq_length = F.shape(input_ids) |
| 580 | if self.is_first_iteration is False: |
| 581 | attention_mask = P.Tile()(self.all_ones_attention_mask, (bs, 1, 1)) |
| 582 | else: |
| 583 | attention_mask = self.get_attention_mask(input_mask) |
| 584 | input_position = F.tuple_to_array(F.make_range(seq_length)) |
| 585 | input_position = P.Tile()(input_position, (bs, 1)) |
| 586 | logits = self.backbone(input_ids, input_position, attention_mask, |
| 587 | init_reset, batch_valid_length) |
| 588 | index = current_index.view(-1, ) |
| 589 | # P.Print()("==logits_is:", logits, ",shape is:", logits.shape) |
| 590 | # P.Print()("==index_is:", index, ",shape is:", index.shape) |
| 591 | logits = self.gather(logits, index, 0) |
| 592 | logits = logits.view(bs, 1, -1) |
| 593 | log_probs = self.log_softmax(logits) |
| 594 | return log_probs |
| 595 | |
| 596 | |
| 597 | class LogitsNet(nn.Cell): |
nothing calls this directly
no outgoing calls
no test coverage detected