netmap.model.zinbautoencoder.ZINBLoss¶
- class netmap.model.zinbautoencoder.ZINBLoss(*args: Any, **kwargs: Any)[source]¶
Bases:
Module- __init__(scale_factor=1.0, eps=1e-10, ridge_lambda=0.0)[source]¶
Zero-Inflated Negative Binomial loss module.
- forward(y_true, y_pred, theta, pi)[source]¶
Compute the ZINB negative log-likelihood loss.
- Parameters:
y_true (torch.Tensor) – Ground truth counts (non-negative integers).
y_pred (torch.Tensor) – Predicted mean values (mu).
theta (torch.Tensor) – Dispersion parameter.
pi (torch.Tensor) – Zero-inflation probability in (0, 1).
- Returns:
- Mean ZINB negative log-likelihood, optionally with
ridge penalty on pi.
- Return type: