Runs eval for the given step. Args: step: Current step. prng_key: PRNG key. model_params: Model parameters. return_aux: Boolean to determine whether outputs are returned. train_summaries: Summaries from the most recent training ste
(
self,
step: int,
*,
prng_key: Tensor,
model_params: NestedTensor,
return_aux: bool = False,
train_summaries: Optional[NestedTensor] = None,
force_run: bool = False,
)
| 617 | self._eval_policy: EvalPolicy = cfg.eval_policy.instantiate() |
| 618 | |
| 619 | def eval_step( |
| 620 | self, |
| 621 | step: int, |
| 622 | *, |
| 623 | prng_key: Tensor, |
| 624 | model_params: NestedTensor, |
| 625 | return_aux: bool = False, |
| 626 | train_summaries: Optional[NestedTensor] = None, |
| 627 | force_run: bool = False, |
| 628 | ) -> tuple[Tensor, Optional[dict[str, Any]], Optional[list[NestedTensor]]]: |
| 629 | """Runs eval for the given step. |
| 630 | |
| 631 | Args: |
| 632 | step: Current step. |
| 633 | prng_key: PRNG key. |
| 634 | model_params: Model parameters. |
| 635 | return_aux: Boolean to determine whether outputs are returned. |
| 636 | train_summaries: Summaries from the most recent training step. Can be used in the |
| 637 | `evaler_policy`. |
| 638 | force_run: If True, force run the eval for the given step. |
| 639 | |
| 640 | Returns: |
| 641 | A tuple (prng_key, summaries, outputs), where |
| 642 | prng_key can be used for a future training step, |
| 643 | summaries contains replicated eval summaries, or None if eval did not run this step, |
| 644 | and outputs contains an optional list of evaler outputs for all of the input. |
| 645 | |
| 646 | Raises: |
| 647 | RuntimeError: If attempting to nest profilers. |
| 648 | """ |
| 649 | cfg = self.config |
| 650 | |
| 651 | if not force_run and not self._eval_policy( |
| 652 | step=step, train_summaries=(train_summaries or {}) |
| 653 | ): |
| 654 | return prng_key, None, None |
| 655 | |
| 656 | if isinstance(self.input, ElasticInput) and self.input.is_in_elastic_mode: |
| 657 | # TODO(jtian22): When evaluating in elastic mode, the |
| 658 | # `_pad_for_evaluation` is called when iterating the elastic feed. |
| 659 | # However, there's a global sync inside and not all processes are |
| 660 | # participating, which will result in a hang issue. Here we simply |
| 661 | # skip it in the elastic mode. We'll support it as a future work. |
| 662 | # Related code: |
| 663 | # https://github.com/apple/axlearn/blob/dffc5135669ae3548a81de2b8cfa4c9d43f4a5de/axlearn/common/input_tf_data.py#L730-L734 |
| 664 | logging.warning( |
| 665 | "Evaluation in elastic mode is currently not supported and will be skipped." |
| 666 | ) |
| 667 | return prng_key, None, None |
| 668 | |
| 669 | self.vlog( |
| 670 | 2, |
| 671 | "%s: Process % 3d step % 8d: starting", |
| 672 | self.path(), |
| 673 | jax.process_index(), |
| 674 | step, |
| 675 | ) |
| 676 |