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
alphaandyshapes do not match.