refactor: isolate Delphi2M next-token pipeline

This commit is contained in:
2026-07-25 14:22:36 +08:00
parent 15ace878f4
commit 315f552301
17 changed files with 330 additions and 1817 deletions

View File

@@ -145,95 +145,6 @@ class Delphi2MLoss(nn.Module):
return total_loss
class UniqueTimeSetExponentialLoss(nn.Module):
"""Next distinct timestamp event-set supervision with sum reduction."""
def __init__(
self,
ignored_idx: Iterable[int] = (PAD_IDX, CHECKUP_IDX),
t_min: float = 1.0 / 365.25,
max_exp_input: float = 60.0,
exclude_ignored_from_intensity: bool = True,
):
super().__init__()
self.ignored_idx = [int(x) for x in ignored_idx]
self.t_min = float(t_min)
self.max_exp_input = float(max_exp_input)
self.exclude_ignored_from_intensity = bool(exclude_ignored_from_intensity)
def forward(
self,
logits: torch.Tensor,
target_multi_hot: torch.Tensor,
target_dt_unique: torch.Tensor,
readout_mask: torch.Tensor,
return_components: bool = False,
) -> torch.Tensor | tuple[torch.Tensor, dict[str, torch.Tensor]]:
if logits.dim() != 3:
raise ValueError(f"logits must be (B, L, K), got {tuple(logits.shape)}")
bsz, seq_len, vocab_size = logits.shape
if target_multi_hot.shape != (bsz, seq_len, vocab_size):
raise ValueError(
"target_multi_hot must match logits shape, "
f"got {tuple(target_multi_hot.shape)} vs {tuple(logits.shape)}"
)
if target_dt_unique.shape != (bsz, seq_len):
raise ValueError(
f"target_dt_unique must be {(bsz, seq_len)}, got {tuple(target_dt_unique.shape)}"
)
if readout_mask.shape != (bsz, seq_len):
raise ValueError(f"readout_mask must be {(bsz, seq_len)}, got {tuple(readout_mask.shape)}")
ignore_mask = _make_ignore_mask(vocab_size, self.ignored_idx, logits.device)
num_targets = target_multi_hot[:, :, ~ignore_mask].sum(dim=-1)
valid_mask = readout_mask.bool() & (num_targets > 0)
if not valid_mask.any():
total_loss = _zero_loss_like(logits)
if return_components:
return total_loss, {
"observed": total_loss.detach(),
"penalty": total_loss.detach(),
"total": total_loss.detach(),
}
return total_loss
logits_safe = torch.nan_to_num(
logits[valid_mask],
nan=0.0,
posinf=self.max_exp_input,
neginf=-self.max_exp_input,
)
target_valid = target_multi_hot[valid_mask].to(logits_safe.dtype)
target_valid[:, ignore_mask] = 0.0
observed_term = (logits_safe * target_valid).sum(dim=-1)
penalty_scale = target_valid.sum(dim=-1)
logits_for_lse = logits_safe
if self.exclude_ignored_from_intensity:
logits_for_lse = logits_safe.masked_fill(ignore_mask.unsqueeze(0), float("-inf"))
dt_clamped = torch.clamp(target_dt_unique[valid_mask], min=self.t_min)
log_lambda_total = torch.logsumexp(logits_for_lse, dim=-1)
log_penalty = log_lambda_total + dt_clamped.log()
penalty = torch.exp(torch.clamp(log_penalty, max=self.max_exp_input))
observed_loss = -observed_term
penalty_loss = penalty_scale * penalty
total_loss = (observed_loss + penalty_loss).mean()
if return_components:
return total_loss, {
"observed": observed_loss.mean().detach(),
"penalty": penalty_loss.mean().detach(),
"total": total_loss.detach(),
}
return total_loss
class ExponentialLoss(nn.Module):
"""Query-conditioned all-future-event exponential point-process loss."""
@@ -385,10 +296,8 @@ class MixedLoss(nn.Module):
def build_loss(name: str, **kwargs) -> nn.Module:
name = name.lower()
if name in {"delphi2m", "d2m", "next_token"}:
if name == "delphi2m":
return Delphi2MLoss(**kwargs)
if name in {"uts", "unique_time_set", "unique_time_exponential"}:
return UniqueTimeSetExponentialLoss(**kwargs)
if name in {"exponential", "query_exponential"}:
return ExponentialLoss(**kwargs)
if name in {"weibull", "query_weibull"}:
@@ -396,5 +305,5 @@ def build_loss(name: str, **kwargs) -> nn.Module:
if name in {"mixed", "query_mixed"}:
return MixedLoss(**kwargs)
raise ValueError(
f"Unknown loss {name!r}. Available: delphi2m, uts, exponential, weibull, mixed."
f"Unknown loss {name!r}. Available: delphi2m, exponential, weibull, mixed."
)