regularization_fn

probly.losses.torch.regularization_fn(alpha: Tensor, y: Tensor) Tensor[source]

Information Robust Dirichlet regularizer from [Tsi19].

Penalizes high Dirichlet concentration values for incorrect classes to encourage confident but well-calibrated predictions.

Parameters:
  • alpha – Dirichlet concentration parameters, shape (B, K), must be > 0.

  • y – One-hot encoded class labels, shape (B, K).

Returns:

Scalar regularization loss summed over classes and batch.

Raises:

ValueError – If alpha and y shapes do not match.