Remove legacy event and mixed distribution paths

This commit is contained in:
2026-08-01 14:23:18 +08:00
parent dfb22adf2d
commit de6f9b75b9
22 changed files with 370 additions and 463 deletions

View File

@@ -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."
)