Skip to content

tausurv.nn.losses.deephit

class DeepHitRankingLoss(*, sigma: float = 0.1, reduction: str = 'mean')

DeepHit pair-ranking regularizer as an nn.Module.

Class form of tausurv.nn.functional.deephit_ranking.

  • sigma — float, default 0.1 — Soft-step bandwidth.
  • reduction — none, default “mean”
class DeepHitLoss(*, alpha: float = 0.5, sigma: float = 0.1, reduction: str = 'mean')

Combined DeepHit loss (PMF NLL + ranking) as an nn.Module.

Class form of tausurv.nn.functional.deephit_loss. alpha=1 recovers pure PMFLoss.

  • alpha — float in [0, 1], default 0.5 — Weight on the NLL term.
  • sigma — float, default 0.1 — Ranking-loss bandwidth.
  • reduction — none, default “mean”