MCPcopy Create free account
hub / github.com/apple/axlearn / eval_step

Method eval_step

axlearn/common/evaler.py:619–759  ·  view source on GitHub ↗

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,
    )

Source from the content-addressed store, hash-verified

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

Callers 5

test_spmd_evalerMethod · 0.80
test_min_stepMethod · 0.80
test_output_writerMethod · 0.80
test_eval_policyMethod · 0.80
_run_evalMethod · 0.80

Calls 10

vlogMethod · 0.80
mapMethod · 0.80
addMethod · 0.80
pathMethod · 0.45
init_stateMethod · 0.45
datasetMethod · 0.45
batchesMethod · 0.45
forwardMethod · 0.45
writeMethod · 0.45
get_summariesMethod · 0.45

Tested by 4

test_spmd_evalerMethod · 0.64
test_min_stepMethod · 0.64
test_output_writerMethod · 0.64
test_eval_policyMethod · 0.64