elbo_loss¶
- probly.losses.torch.elbo_loss(inputs: Tensor, targets: Tensor, kl: Tensor, *, kl_penalty: float = 1e-05) Tensor[source]¶
Evidence lower bound loss based on [BCKW15].
- Parameters:
inputs – Logits of size (n_instances, n_classes).
targets – Class labels of size (n_instances,).
kl – KL divergence of the model.
kl_penalty – Weight for KL divergence term.
- Returns:
The mean loss value.