Skip to content

jaxpi.evaluator

Metric logging during training.

BaseEvaluator [source]

BaseEvaluator.log_lr() [source]

python
BaseEvaluator.log_lr(self, model, state)

BaseEvaluator.log_losses() [source]

python
BaseEvaluator.log_losses(self, loss_dict)

BaseEvaluator.log_raw_losses() [source]

python
BaseEvaluator.log_raw_losses(self, model, params, state, batch)

BaseEvaluator.log_loss_weights() [source]

python
BaseEvaluator.log_loss_weights(self, state)

BaseEvaluator.log_pts_weights() [source]

python
BaseEvaluator.log_pts_weights(self, state)

BaseEvaluator.log_grads() [source]

python
BaseEvaluator.log_grads(self, model, params, batch)

BaseEvaluator.__call__() [source]

python
BaseEvaluator.__call__(self, model, state, loss_dict, batch, *args)

Call self as a function.

Released under the Apache 2.0 License.