NegLogProbLoss#
- class liesel.optim.NegLogProbLoss(model, split, validation_strategy='log_lik', scale=False)[source]#
Bases:
LossMixinNegative log-probability loss for Liesel models.
The training objective is the negative sum of the model log-likelihood and log-prior. During mini-batch optimization, likelihood terms are scaled through
carry.batches.scaled_log_lik(...)soBatchManagercan apply branch-specific scaling for multi-size observed data. Validation loss usessplit.scaled_log_lik(...)for the same reason.The model must use the standard sums of observed distribution factors and parameter priors. Weak observed variables and weak parameters contribute their likelihoods and priors like strong ones. Custom aggregate nodes or additional unclassified distribution factors require a custom
Loss; supplying a manual split is not sufficient.- Parameters:
model (
Model) – Liesel model evaluated by the loss.split (
PositionSplit|PositionSplitManager) – Train/validation/test split. UsePositionSplitManagerfor models with observed branches of different sample sizes.validation_strategy (
Literal['log_lik','log_prob'], default:'log_lik') – Validation objective."log_lik"uses the scaled log-likelihood only."log_prob"also includes the model log-prior.scale (
bool, default:False) – IfTrue, divide losses by the training sample size. ForPositionSplitManager, the scalar is the sum of all branch-specific training sizes.
Examples
Construct a default loss for a simple observed model:
>>> import jax.numpy as jnp >>> import liesel.model as lsl >>> import tensorflow_probability.substrates.jax.distributions as tfd >>> from liesel.optim import NegLogProbLoss, PositionSplit >>> y = lsl.Var.new_obs( ... jnp.arange(3.0), ... lsl.Dist(tfd.Normal, loc=0.0, scale=1.0), ... name="y", ... ) >>> model = lsl.Model([y]) >>> split = PositionSplit.from_model(model, position_keys=["y"]) >>> loss = NegLogProbLoss(model, split) >>> loss.position([]) == {} True >>> repr(loss) 'NegLogProbLoss(validation_strategy=log_lik)'
Methods
loss_monitor(params, carry)Computes validation loss.
loss_train(params, carry)Computes full-data negative log posterior.
loss_train_batched(params, carry)Computes mini-batch negative log posterior.
position(position_keys)Extracts an initial optimizer position from the model.
Attributes
Liesel model evaluated by the loss.
Scalar multiplier applied to the negative log probability.
If
True, divide losses by the training sample size.Train/validation/test split.
Validation objective.