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

@@ -156,7 +156,7 @@ class DeepHealth(nn.Module):
n_value_kinds: int = 3,
n_bins: int = 16,
target_mode: str = "next_token", # "next_token" or "all_future"
time_mode: str = "relative", # "relative" or "absolute"
time_mode: str = "absolute", # next_token requires absolute
dist_mode: str = "exponential", # "exponential", "weibull" or "mixed"
extra_pool_reduce: str = "mean",
dropout: float = 0.0,
@@ -169,6 +169,11 @@ class DeepHealth(nn.Module):
if time_mode not in ["relative", "absolute"]:
raise ValueError(
"time_mode must be either 'relative' or 'absolute'")
if target_mode == "next_token" and time_mode != "absolute":
raise ValueError(
"next_token is reserved for Delphi2M reproduction and "
"requires time_mode='absolute'"
)
if dist_mode not in ["exponential", "weibull", "mixed"]:
raise ValueError(
"dist_mode must be either 'exponential', 'weibull' or 'mixed'")
@@ -461,15 +466,8 @@ class DeepHealth(nn.Module):
)
return h_disease[:, :event_len, :]
def forward_next_token(self, **kwargs) -> torch.Tensor:
return self._forward_shared(mode="next_token", **kwargs)
def forward_all_future(self, **kwargs) -> torch.Tensor:
return self._forward_shared(mode="all_future", **kwargs)
def forward(self, target_mode: str | None = None, **kwargs) -> torch.Tensor:
mode = self.target_mode if target_mode is None else target_mode
return self._forward_shared(mode=mode, **kwargs)
def forward(self, **kwargs) -> torch.Tensor:
return self._forward_shared(mode=self.target_mode, **kwargs)
def calc_risk(self, x: torch.Tensor) -> torch.Tensor:
return self.risk_head(x)