jaxpi.evaluator
Metric logging.
Metric names use "section/name" keys (e.g. "loss/res", "error/u", "weights/u_ic") so that W&B groups them into separate chart sections: losses, errors, adaptive weights, gradient norms, etc.
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, state, batch)BaseEvaluator.__call__() [source]
python
BaseEvaluator.__call__(self, model, state, loss_dict, batch, *args)Call self as a function.