refactor: isolate Delphi2M next-token pipeline
This commit is contained in:
95
losses.py
95
losses.py
@@ -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."
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user