Remove legacy event and mixed distribution paths
This commit is contained in:
78
losses.py
78
losses.py
@@ -8,7 +8,7 @@ import torch.nn.functional as F
|
||||
|
||||
|
||||
PAD_IDX = 0
|
||||
CHECKUP_IDX = 1
|
||||
RESERVED_IDX = 1
|
||||
NO_EVENT_IDX = 2
|
||||
|
||||
|
||||
@@ -52,7 +52,7 @@ class Delphi2MLoss(nn.Module):
|
||||
super().__init__()
|
||||
self.t_min = float(t_min)
|
||||
self.ignored_tokens = (
|
||||
[PAD_IDX, CHECKUP_IDX]
|
||||
[PAD_IDX, RESERVED_IDX]
|
||||
if ignored_tokens is None
|
||||
else [int(x) for x in ignored_tokens]
|
||||
)
|
||||
@@ -150,7 +150,7 @@ class ExponentialLoss(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
ignored_idx: Iterable[int] = (PAD_IDX, CHECKUP_IDX),
|
||||
ignored_idx: Iterable[int] = (PAD_IDX, RESERVED_IDX),
|
||||
eps: float = 1e-8,
|
||||
):
|
||||
super().__init__()
|
||||
@@ -183,7 +183,7 @@ class WeibullLoss(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
ignored_idx: Iterable[int] = (PAD_IDX, CHECKUP_IDX),
|
||||
ignored_idx: Iterable[int] = (PAD_IDX, RESERVED_IDX),
|
||||
eps: float = 1e-8,
|
||||
):
|
||||
super().__init__()
|
||||
@@ -232,78 +232,14 @@ class WeibullLoss(nn.Module):
|
||||
return (-observed + penalty).mean()
|
||||
|
||||
|
||||
class MixedLoss(nn.Module):
|
||||
"""Exponential diseases plus one Weibull death endpoint."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
death_idx: int,
|
||||
ignored_idx: Iterable[int] = (PAD_IDX, CHECKUP_IDX),
|
||||
eps: float = 1e-8,
|
||||
):
|
||||
super().__init__()
|
||||
self.death_idx = int(death_idx)
|
||||
self.ignored_idx = tuple(int(i) for i in ignored_idx)
|
||||
self.eps = eps
|
||||
|
||||
def forward(
|
||||
self,
|
||||
logits: torch.Tensor,
|
||||
death_rho: torch.Tensor,
|
||||
targets: torch.Tensor,
|
||||
dt: torch.Tensor,
|
||||
exposure: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
_, vocab_size = logits.shape
|
||||
dtype = logits.dtype
|
||||
rate = F.softplus(logits) + self.eps
|
||||
|
||||
if death_rho.dim() == 2:
|
||||
death_rho = death_rho.squeeze(-1)
|
||||
death_rho = death_rho.to(device=logits.device, dtype=dtype).clamp_min(self.eps)
|
||||
|
||||
valid_vocab = _valid_vocab_mask(vocab_size, self.ignored_idx, logits.device)
|
||||
valid_disease_vocab = valid_vocab.clone()
|
||||
valid_disease_vocab[self.death_idx] = False
|
||||
|
||||
t_exp = exposure.to(dtype).clamp_min(self.eps)
|
||||
disease_penalty = t_exp * rate[:, valid_disease_vocab].sum(dim=-1)
|
||||
death_rate = rate[:, self.death_idx]
|
||||
death_penalty = death_rate * torch.pow(t_exp, death_rho)
|
||||
penalty = disease_penalty + death_penalty
|
||||
|
||||
target_valid = torch.ones_like(targets, dtype=torch.bool, device=logits.device)
|
||||
for idx in self.ignored_idx:
|
||||
target_valid &= targets != idx
|
||||
|
||||
disease_event_mask = target_valid & (targets != self.death_idx)
|
||||
safe_targets = targets.clamp(min=0, max=vocab_size - 1)
|
||||
disease_log_rate = rate.log().gather(1, safe_targets)
|
||||
observed_disease = (disease_log_rate * disease_event_mask.to(dtype)).sum(dim=-1)
|
||||
|
||||
death_event_mask = target_valid & (targets == self.death_idx)
|
||||
death_observed = death_event_mask.any(dim=1)
|
||||
death_dt = (dt.to(dtype).clamp_min(self.eps) * death_event_mask.to(dtype)).sum(dim=1)
|
||||
death_log_intensity = (
|
||||
death_rate.log()
|
||||
+ death_rho.log()
|
||||
+ (death_rho - 1.0) * death_dt.clamp_min(self.eps).log()
|
||||
)
|
||||
observed_death = death_log_intensity * death_observed.to(dtype)
|
||||
|
||||
return (-observed_disease - observed_death + penalty).mean()
|
||||
|
||||
|
||||
def build_loss(name: str, **kwargs) -> nn.Module:
|
||||
name = name.lower()
|
||||
if name == "delphi2m":
|
||||
return Delphi2MLoss(**kwargs)
|
||||
if name in {"exponential", "query_exponential"}:
|
||||
if name == "exponential":
|
||||
return ExponentialLoss(**kwargs)
|
||||
if name in {"weibull", "query_weibull"}:
|
||||
if name == "weibull":
|
||||
return WeibullLoss(**kwargs)
|
||||
if name in {"mixed", "query_mixed"}:
|
||||
return MixedLoss(**kwargs)
|
||||
raise ValueError(
|
||||
f"Unknown loss {name!r}. Available: delphi2m, exponential, weibull, mixed."
|
||||
f"Unknown loss {name!r}. Available: delphi2m, exponential, weibull."
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user