| from __future__ import annotations |
|
|
| import math |
| import sys |
| from dataclasses import dataclass, field |
| from pathlib import Path |
| from typing import Any |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
|
|
|
|
| SCRIPT_DIR = Path(__file__).resolve().parent |
| V3P5_SCRIPT_DIR = SCRIPT_DIR.parents[1] / "v3p5_static" / "scripts" |
| if str(V3P5_SCRIPT_DIR) not in sys.path: |
| sys.path.insert(0, str(V3P5_SCRIPT_DIR)) |
|
|
| from model_v3p5 import ( |
| DrugSetEncoder, |
| FieldSpecificCategoricalHead, |
| PatientStaticEncoder, |
| ProportionalOddsOrdinalHead, |
| build_service_attention_prior, |
| ) |
|
|
|
|
| PWE_BIN_EDGES_DAYS = (7.0, 14.0, 30.0, 60.0, 90.0, 365.0, math.inf) |
| PWE_FINITE_LOWER_DAYS = (0.0, 7.0, 14.0, 30.0, 60.0, 90.0, 365.0) |
|
|
|
|
| @dataclass |
| class SCTMv4p3Config: |
| cat_vocab_size: int |
| drug_vocab_size: int |
| n_cat_fields: int |
| n_numeric_fields: int |
| n_ordinal_fields: int |
| cbe_dim: int |
| max_drug_atoms: int |
| n_events: int |
| d_model: int = 192 |
| n_heads: int = 6 |
| n_layers: int = 4 |
| dropout: float = 0.05 |
| n_missing: int = 5 |
| n_drug_classes: int = 4 |
| n_service_states: int = 8 |
| n_active_states: int = 5 |
| n_pwe_causes: int = 3 |
| n_pwe_bins: int = 7 |
| n_eras: int = 3 |
| n_conditional_status: int = 4 |
| static_vocab_size: int = 0 |
| n_static_fields: int = 0 |
| year_min: int = 2010 |
| year_max: int = 2026 |
| missing_observed_id: int = 0 |
| missing_era_unavailable_id: int = 1 |
| missing_visit_missing_id: int = 3 |
| missing_structural_inactive_id: int = 2 |
| missing_no_clinical_id: int = 4 |
| existing_missing_visit_missing_id: int = 2 |
| existing_missing_pad_id: int = 3 |
| unknown_cat_value_id: int = 1 |
| contact_injection_init: float = 0.05 |
| two_year_days: float = 730.0 |
| conditional_specs: tuple[dict[str, Any], ...] = field(default_factory=tuple) |
| color_cat_field_index: int = -1 |
| color_value_to_class: tuple[tuple[int, int], ...] = field(default_factory=tuple) |
|
|
|
|
| def era_tokens(visit_year: torch.Tensor) -> torch.Tensor: |
| year = visit_year.long() |
| return torch.where(year <= 2013, torch.zeros_like(year), torch.where(year <= 2018, torch.ones_like(year), torch.full_like(year, 2))) |
|
|
|
|
| def year_features(visit_year: torch.Tensor, year_min: int, year_max: int) -> torch.Tensor: |
| year = visit_year.float().clamp(min=year_min, max=year_max) |
| scale = max(1.0, float(year_max - year_min)) |
| phase = (year - year_min) / scale * (2.0 * math.pi) |
| return torch.stack([torch.sin(phase), torch.cos(phase)], dim=-1) |
|
|
|
|
| def previous_interval_log(time_since_start_days: torch.Tensor) -> torch.Tensor: |
| prev_gap = torch.zeros_like(time_since_start_days) |
| prev_gap[:, 1:] = (time_since_start_days[:, 1:] - time_since_start_days[:, :-1]).clamp(min=0.0) |
| return torch.log1p(prev_gap) |
|
|
|
|
| class AxisSlotTabEncoderV4p3(nn.Module): |
| """v4.3 intra-visit encoder with conditional-field status and 3-era tokens.""" |
|
|
| def __init__(self, config: SCTMv4p3Config) -> None: |
| super().__init__() |
| self.config = config |
| d = config.d_model |
| self.cat_field_emb = nn.Embedding(config.n_cat_fields, d) |
| self.cat_value_emb = nn.Embedding(config.cat_vocab_size, d, padding_idx=0) |
| self.missing_emb = nn.Embedding(config.n_missing, d) |
| self.conditional_status_emb = nn.Embedding(config.n_conditional_status, d) |
| self.numeric_field_emb = nn.Embedding(max(1, config.n_numeric_fields), d) |
| self.numeric_value_proj = nn.Linear(1, d) |
| self.ordinal_field_emb = nn.Embedding(max(1, config.n_ordinal_fields), d) |
| self.ordinal_cbe_proj = nn.Linear(config.cbe_dim, d) |
| self.drug_encoder = DrugSetEncoder(config) |
| self.service_emb = nn.Embedding(config.n_service_states, d) |
| self.year_proj = nn.Linear(2, d) |
| self.era_emb = nn.Embedding(config.n_eras, d) |
| self.axis_slot_emb = nn.Parameter(torch.randn(1, 1, 7, d) * 0.02) |
| self.axis_attn = nn.MultiheadAttention(d, config.n_heads, dropout=config.dropout, batch_first=True) |
| self.axis_norm = nn.LayerNorm(d) |
| self.axis_ff = nn.Sequential(nn.Linear(d, d * 2), nn.GELU(), nn.Dropout(config.dropout), nn.Linear(d * 2, d)) |
| self.out_norm = nn.LayerNorm(d) |
| self.register_buffer("cat_field_ids", torch.arange(config.n_cat_fields), persistent=False) |
| self.register_buffer("numeric_field_ids", torch.arange(max(1, config.n_numeric_fields)), persistent=False) |
| self.register_buffer("ordinal_field_ids", torch.arange(max(1, config.n_ordinal_fields)), persistent=False) |
|
|
| def conditional_status(self, cat_value_ids: torch.Tensor, missing_ids: torch.Tensor) -> torch.Tensor: |
| """Return ACTIVE/STRUCTURAL_INACTIVE/PARENT_UNKNOWN/DATA_INCONSISTENT codes.""" |
|
|
| status = torch.zeros_like(cat_value_ids, dtype=torch.long) |
| if not self.config.conditional_specs: |
| return status |
| missing = missing_ids.clamp(min=0, max=self.config.n_missing - 1) |
| observed = missing.eq(self.config.missing_observed_id) |
| unknown = self.config.unknown_cat_value_id |
| for spec in self.config.conditional_specs: |
| parent_idx = int(spec["parent_idx"]) |
| child_indices = [int(x) for x in spec["child_indices"]] |
| active_ids = torch.tensor([int(x) for x in spec["active_value_ids"]], device=cat_value_ids.device, dtype=cat_value_ids.dtype) |
| if parent_idx < 0 or parent_idx >= cat_value_ids.shape[-1] or not child_indices or active_ids.numel() == 0: |
| continue |
| parent_val = cat_value_ids[:, :, parent_idx] |
| parent_observed = observed[:, :, parent_idx] & parent_val.ne(unknown) |
| parent_active = (parent_val[:, :, None] == active_ids[None, None, :]).any(dim=-1) & parent_observed |
| parent_unknown = ~parent_observed |
| for child_idx in child_indices: |
| if child_idx < 0 or child_idx >= cat_value_ids.shape[-1]: |
| continue |
| child_val = cat_value_ids[:, :, child_idx] |
| child_observed = observed[:, :, child_idx] & child_val.ne(unknown) |
| code = torch.zeros_like(parent_val, dtype=torch.long) |
| code = torch.where(parent_unknown, torch.full_like(code, 2), code) |
| code = torch.where((~parent_unknown) & (~parent_active) & (~child_observed), torch.full_like(code, 1), code) |
| code = torch.where((~parent_unknown) & (~parent_active) & child_observed, torch.full_like(code, 3), code) |
| status[:, :, child_idx] = code |
| return status |
|
|
| def forward( |
| self, |
| cat_value_ids: torch.Tensor, |
| missing_ids: torch.Tensor, |
| numeric_values: torch.Tensor, |
| numeric_mask: torch.Tensor, |
| ordinal_cbe: torch.Tensor, |
| ordinal_mask: torch.Tensor, |
| drug_name_ids: torch.Tensor, |
| drug_class_ids: torch.Tensor, |
| drug_mask: torch.Tensor, |
| service_state: torch.Tensor, |
| visit_year: torch.Tensor, |
| ) -> tuple[torch.Tensor, torch.Tensor]: |
| bsz, seq_len, _ = cat_value_ids.shape |
| device = cat_value_ids.device |
| cat_fields = self.cat_field_emb(self.cat_field_ids.to(device))[None, None, :, :] |
| missing = missing_ids.clamp(min=0, max=self.config.n_missing - 1) |
| conditional = self.conditional_status(cat_value_ids, missing) |
| cat = ( |
| cat_fields |
| + self.cat_value_emb(cat_value_ids.clamp(min=0, max=self.config.cat_vocab_size - 1)) |
| + self.missing_emb(missing) |
| + self.conditional_status_emb(conditional.clamp(0, self.config.n_conditional_status - 1)) |
| ) |
| cat_vec = cat.mean(dim=2) |
|
|
| if self.config.n_numeric_fields: |
| num_fields = self.numeric_field_emb(self.numeric_field_ids[: self.config.n_numeric_fields].to(device))[None, None, :, :] |
| num = num_fields + self.numeric_value_proj(numeric_values.unsqueeze(-1)) |
| num = num * numeric_mask.unsqueeze(-1).float() |
| denom = numeric_mask.sum(dim=2).clamp(min=1).unsqueeze(-1).float() |
| num_vec = num.sum(dim=2) / denom |
| else: |
| num_vec = torch.zeros_like(cat_vec) |
|
|
| if self.config.n_ordinal_fields: |
| ord_fields = self.ordinal_field_emb(self.ordinal_field_ids[: self.config.n_ordinal_fields].to(device))[None, None, :, :] |
| ord_vecs = ord_fields + self.ordinal_cbe_proj(ordinal_cbe) |
| ord_vecs = ord_vecs * ordinal_mask.unsqueeze(-1).float() |
| denom = ordinal_mask.sum(dim=2).clamp(min=1).unsqueeze(-1).float() |
| ord_vec = ord_vecs.sum(dim=2) / denom |
| else: |
| ord_vec = torch.zeros_like(cat_vec) |
|
|
| drug_vec = self.drug_encoder(drug_name_ids, drug_class_ids, drug_mask) |
| service_vec = self.service_emb(service_state.clamp(min=0, max=self.config.n_service_states - 1)) |
| year_vec = self.year_proj(year_features(visit_year, self.config.year_min, self.config.year_max)) |
| era_vec = self.era_emb(era_tokens(visit_year).clamp(0, self.config.n_eras - 1)) |
| slots = torch.stack([cat_vec, num_vec, ord_vec, drug_vec, service_vec, year_vec, era_vec], dim=2) |
| slots = slots + self.axis_slot_emb |
| flat_slots = slots.reshape(bsz * seq_len, 7, -1) |
| attended, _ = self.axis_attn(flat_slots, flat_slots, flat_slots, need_weights=False) |
| flat_slots = self.axis_norm(flat_slots + attended) |
| flat_slots = self.axis_norm(flat_slots + self.axis_ff(flat_slots)) |
| visit_vec = flat_slots.mean(dim=1).reshape(bsz, seq_len, -1) |
| return self.out_norm(visit_vec), conditional |
|
|
|
|
| class EraTimeStateSelfAttention(nn.Module): |
| def __init__(self, config: SCTMv4p3Config, service_prior_bias: torch.Tensor | None = None) -> None: |
| super().__init__() |
| if config.d_model % config.n_heads != 0: |
| raise ValueError("d_model must be divisible by n_heads") |
| self.config = config |
| self.head_dim = config.d_model // config.n_heads |
| self.qkv = nn.Linear(config.d_model, config.d_model * 3) |
| self.out = nn.Linear(config.d_model, config.d_model) |
| self.dropout = nn.Dropout(config.dropout) |
| self.time_alpha = nn.Parameter(torch.zeros(config.n_heads)) |
| init = torch.zeros(config.n_heads, config.n_eras, config.n_eras, config.n_service_states, config.n_service_states) |
| if service_prior_bias is not None: |
| if service_prior_bias.shape != (config.n_service_states, config.n_service_states): |
| raise ValueError("service_prior_bias must be [n_service_states, n_service_states]") |
| prior = service_prior_bias.float().to(device=init.device) |
| init = init + prior[None, None, None, :, :] |
| self.state_era_bias = nn.Parameter(init) |
|
|
| def forward( |
| self, |
| x: torch.Tensor, |
| service_state: torch.Tensor, |
| visit_year: torch.Tensor, |
| time_since_start_days: torch.Tensor, |
| valid_mask: torch.Tensor, |
| ) -> torch.Tensor: |
| bsz, seq_len, d_model = x.shape |
| qkv = self.qkv(x).view(bsz, seq_len, 3, self.config.n_heads, self.head_dim) |
| q, k, v = qkv[:, :, 0].transpose(1, 2), qkv[:, :, 1].transpose(1, 2), qkv[:, :, 2].transpose(1, 2) |
| scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim) |
| tdiff = (time_since_start_days[:, :, None] - time_since_start_days[:, None, :]).abs().clamp(min=0) |
| time_penalty = -F.softplus(self.time_alpha)[None, :, None, None] * torch.log1p(tdiff)[:, None, :, :] |
| scores = scores + time_penalty |
|
|
| qs = service_state.clamp(0, self.config.n_service_states - 1) |
| ks = qs |
| qe = era_tokens(visit_year).clamp(0, self.config.n_eras - 1) |
| ke = qe |
| qid = qe * self.config.n_service_states + qs |
| kid = ke * self.config.n_service_states + ks |
| bias_table = self.state_era_bias.reshape( |
| self.config.n_heads, |
| self.config.n_eras * self.config.n_service_states, |
| self.config.n_eras * self.config.n_service_states, |
| ) |
| bias = bias_table[:, qid[:, :, None], kid[:, None, :]].permute(1, 0, 2, 3) |
| scores = scores + bias |
|
|
| causal = torch.triu(torch.ones(seq_len, seq_len, dtype=torch.bool, device=x.device), diagonal=1) |
| scores = scores.masked_fill(causal[None, None, :, :], torch.finfo(scores.dtype).min) |
| scores = scores.masked_fill(~valid_mask[:, None, None, :], torch.finfo(scores.dtype).min) |
| attn = torch.softmax(scores, dim=-1) |
| attn = torch.nan_to_num(attn, nan=0.0) |
| y = torch.matmul(self.dropout(attn), v).transpose(1, 2).contiguous().view(bsz, seq_len, d_model) |
| return self.out(y) |
|
|
|
|
| class EraTrajDecoderLayer(nn.Module): |
| def __init__(self, config: SCTMv4p3Config, service_prior_bias: torch.Tensor | None = None) -> None: |
| super().__init__() |
| d = config.d_model |
| self.attn = EraTimeStateSelfAttention(config, service_prior_bias=service_prior_bias) |
| self.norm1 = nn.LayerNorm(d) |
| self.norm2 = nn.LayerNorm(d) |
| self.ff = nn.Sequential(nn.Linear(d, d * 4), nn.GELU(), nn.Dropout(config.dropout), nn.Linear(d * 4, d)) |
| self.dropout = nn.Dropout(config.dropout) |
|
|
| def forward( |
| self, |
| x: torch.Tensor, |
| service_state: torch.Tensor, |
| visit_year: torch.Tensor, |
| time_since_start_days: torch.Tensor, |
| valid_mask: torch.Tensor, |
| ) -> torch.Tensor: |
| x = x + self.dropout(self.attn(self.norm1(x), service_state, visit_year, time_since_start_days, valid_mask)) |
| x = x + self.dropout(self.ff(self.norm2(x))) |
| return x |
|
|
|
|
| class EraTrajDecoder(nn.Module): |
| def __init__(self, config: SCTMv4p3Config, service_prior_bias: torch.Tensor | None = None) -> None: |
| super().__init__() |
| self.config = config |
| self.start_visit = nn.Parameter(torch.zeros(1, 1, config.d_model)) |
| self.layers = nn.ModuleList([EraTrajDecoderLayer(config, service_prior_bias=service_prior_bias) for _ in range(config.n_layers)]) |
| self.final_norm = nn.LayerNorm(config.d_model) |
|
|
| def causal_decode( |
| self, |
| u: torch.Tensor, |
| service_state: torch.Tensor, |
| visit_year: torch.Tensor, |
| time_since_start_days: torch.Tensor, |
| valid_mask: torch.Tensor, |
| ) -> torch.Tensor: |
| x = u |
| for layer in self.layers: |
| x = layer(x, service_state, visit_year, time_since_start_days, valid_mask) |
| return self.final_norm(x) |
|
|
| def history_inputs( |
| self, |
| u: torch.Tensor, |
| service_state: torch.Tensor, |
| visit_year: torch.Tensor, |
| time_since_start_days: torch.Tensor, |
| ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: |
| bsz = u.shape[0] |
| shifted_u = torch.cat([self.start_visit.expand(bsz, 1, -1), u[:, :-1, :]], dim=1) |
| shifted_service = torch.cat([torch.zeros_like(service_state[:, :1]), service_state[:, :-1]], dim=1) |
| shifted_year = torch.cat([visit_year[:, :1], visit_year[:, :-1]], dim=1) |
| shifted_time = torch.cat([torch.zeros_like(time_since_start_days[:, :1]), time_since_start_days[:, :-1]], dim=1) |
| return shifted_u, shifted_service, shifted_year, shifted_time |
|
|
| def forward( |
| self, |
| u: torch.Tensor, |
| service_state: torch.Tensor, |
| visit_year: torch.Tensor, |
| time_since_start_days: torch.Tensor, |
| valid_mask: torch.Tensor, |
| ) -> tuple[torch.Tensor, torch.Tensor]: |
| h_u, h_s, h_y, h_t = self.history_inputs(u, service_state, visit_year, time_since_start_days) |
| z_history = self.causal_decode(h_u, h_s, h_y, h_t, valid_mask) |
| z_post = self.causal_decode(u, service_state, visit_year, time_since_start_days, valid_mask) |
| return z_history, z_post |
|
|
|
|
| class ContactInjector(nn.Module): |
| def __init__(self, config: SCTMv4p3Config) -> None: |
| super().__init__() |
| self.config = config |
| d = config.d_model |
| self.arrival_proj = nn.Linear(4, d) |
| self.era_emb = nn.Embedding(config.n_eras, d) |
| self.norm = nn.LayerNorm(d) |
| self.scale = nn.Parameter(torch.tensor(float(config.contact_injection_init))) |
|
|
| def forward(self, z_history: torch.Tensor, batch: dict[str, torch.Tensor]) -> torch.Tensor: |
| prev_log = previous_interval_log(batch["time_since_start_days"]) |
| year_feat = year_features(batch["visit_year"], self.config.year_min, self.config.year_max) |
| if "visit_indices" in batch: |
| visit_idx = batch["visit_indices"].float().clamp(min=0.0) / 256.0 |
| else: |
| visit_idx = torch.zeros_like(prev_log) |
| arrival = torch.cat([prev_log.unsqueeze(-1), year_feat, visit_idx.unsqueeze(-1)], dim=-1) |
| contact = self.arrival_proj(arrival) + self.era_emb(era_tokens(batch["visit_year"]).clamp(0, self.config.n_eras - 1)) |
| return self.norm(z_history + self.scale * contact) |
|
|
|
|
| class SCTMv4p3(nn.Module): |
| """SCTM v4.3 L1/L2-aware architecture. |
| |
| The raw v3p5 tensor remains the compatibility contract. v4.3-specific |
| targets (3-cause K, 5-state missingness, conditional structural inactivity) |
| are derived in the loss from that tensor and the metadata-built config. |
| """ |
|
|
| def __init__( |
| self, |
| config: SCTMv4p3Config, |
| field_value_mask: torch.Tensor, |
| service_prior_bias: torch.Tensor | None = None, |
| ) -> None: |
| super().__init__() |
| if config.n_pwe_bins != len(PWE_BIN_EDGES_DAYS): |
| raise ValueError("v4.3 uses the 7-bin data-verified PWE grid") |
| self.config = config |
| self.tab_encoder = AxisSlotTabEncoderV4p3(config) |
| self.traj_decoder = EraTrajDecoder(config, service_prior_bias=service_prior_bias) |
| self.contact_injector = ContactInjector(config) |
| self.static_encoder = PatientStaticEncoder(config) |
| self.static_injection_scale = nn.Parameter(torch.zeros(1)) |
| d = config.d_model |
| self.next_delta_condition = nn.Linear(1, d) |
| self.pwe_log_lambda_pre_head = nn.Linear(d, config.n_pwe_causes * config.n_pwe_bins) |
| self.pwe_log_lambda_contact_head = nn.Linear(d, config.n_pwe_causes * config.n_pwe_bins) |
| self.pwe_log_lambda_post_head = nn.Linear(d, config.n_pwe_causes * config.n_pwe_bins) |
| self.active_state_head = nn.Linear(d, config.n_active_states) |
| self.missingness_head = nn.Linear(d, config.n_cat_fields * config.n_missing) |
| self.deployable_risk_pre_head = nn.Linear(d, config.n_events) |
| self.deployable_risk_contact_head = nn.Sequential( |
| nn.Linear(d * 2, d), |
| nn.GELU(), |
| nn.Dropout(config.dropout), |
| nn.Linear(d, config.n_events), |
| ) |
| self.deployable_risk_post_head = nn.Linear(d, config.n_events) |
| self.field_head = FieldSpecificCategoricalHead(d, config.n_cat_fields, config.cat_vocab_size, field_value_mask) |
| self.numeric_head = nn.Linear(d, max(1, config.n_numeric_fields)) |
| self.ordinal_head = ProportionalOddsOrdinalHead(d, config.n_ordinal_fields, config.cbe_dim) |
| self.ontology_event_head = nn.Linear(d, config.n_events) |
| self.history_event_head = nn.Linear(d, config.n_events) |
| self.contact_event_head = nn.Linear(d, config.n_events) |
| self.color_signal_head = nn.Linear(d, 4) |
| self.rollout_prefix_proj = nn.Sequential( |
| nn.Linear(config.n_pwe_causes + config.n_active_states + config.n_events + 2, d), |
| nn.GELU(), |
| nn.LayerNorm(d), |
| ) |
| |
| |
| |
| nn.init.zeros_(self.tab_encoder.conditional_status_emb.weight) |
| nn.init.zeros_(self.tab_encoder.era_emb.weight) |
| self._init_pwe_hazard_bias() |
|
|
| def _init_pwe_hazard_bias(self) -> None: |
| |
| |
| prior = torch.full((self.config.n_pwe_causes, self.config.n_pwe_bins), -10.0) |
| prior[0, :] = -4.4 |
| prior[1, :] = -9.0 |
| prior[2, :] = -10.5 |
| with torch.no_grad(): |
| for head in (self.pwe_log_lambda_pre_head, self.pwe_log_lambda_contact_head, self.pwe_log_lambda_post_head): |
| head.bias.copy_(prior.reshape(-1)) |
|
|
| def _teacher_forced_next_context(self, z_post: torch.Tensor, batch: dict[str, torch.Tensor]) -> torch.Tensor: |
| """Condition generated next-visit observations on sampled/teacher-forced Δt only. |
| |
| The previous version also injected the true `service_state[t+1]`, which |
| made `active_state_head` a label copier. A_t is now predicted from |
| z_post plus the generated/teacher-forced time interval, and no next |
| service label enters any deployable risk path. |
| """ |
|
|
| delta_ctx = self.next_delta_condition(batch["delta_t_next_log"].unsqueeze(-1).float()) |
| return z_post + delta_ctx |
|
|
| def _heads_from_states(self, z_history: torch.Tensor, z_contact: torch.Tensor, z_post: torch.Tensor, batch: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: |
| bsz, seq_len, _ = z_post.shape |
| tf_ctx = self._teacher_forced_next_context(z_post, batch) |
| pwe_pre = self.pwe_log_lambda_pre_head(z_history).view(bsz, seq_len, self.config.n_pwe_causes, self.config.n_pwe_bins) |
| pwe_contact = self.pwe_log_lambda_contact_head(z_contact).view(bsz, seq_len, self.config.n_pwe_causes, self.config.n_pwe_bins) |
| pwe_post = self.pwe_log_lambda_post_head(z_post).view(bsz, seq_len, self.config.n_pwe_causes, self.config.n_pwe_bins) |
| deployable_pre = self.deployable_risk_pre_head(z_history) |
| deployable_contact = self.deployable_risk_contact_head(torch.cat([z_history, z_contact], dim=-1)) |
| deployable_post = self.deployable_risk_post_head(z_post) |
| ontology_event_logits = self.ontology_event_head(tf_ctx) |
| out = { |
| "z_history": z_history, |
| "z_contact": z_contact, |
| "z_post": z_post, |
| "z_minus": z_history, |
| "z_plus": z_post, |
| "pwe_log_lambda_pre": pwe_pre, |
| "pwe_log_lambda_contact": pwe_contact, |
| "pwe_log_lambda_post": pwe_post, |
| "pwe_log_lambda": pwe_post, |
| "active_state_logits": self.active_state_head(tf_ctx), |
| "missingness_logits": self.missingness_head(tf_ctx).view(bsz, seq_len, self.config.n_cat_fields, self.config.n_missing), |
| "risk_event_logits": deployable_contact, |
| "deployable_risk_logits": deployable_contact, |
| "deployable_risk_logits_pre": deployable_pre, |
| "deployable_risk_logits_contact": deployable_contact, |
| "deployable_risk_logits_post": deployable_post, |
| "field_logits": self.field_head(tf_ctx), |
| "numeric_mu": self.numeric_head(tf_ctx)[..., : self.config.n_numeric_fields], |
| "ordinal_cum_logits": self.ordinal_head(tf_ctx), |
| "ontology_event_logits": ontology_event_logits, |
| "history_event_logits": self.history_event_head(z_history), |
| "contact_event_logits": self.contact_event_head(z_contact), |
| "color_cum_logits": self.color_signal_head(tf_ctx), |
| |
| |
| |
| "event_logits": deployable_contact, |
| "future_risk_logits": self.history_event_head(z_history), |
| } |
| pwe_ci = pwe_closed_form_cif(pwe_post, torch.tensor(PWE_BIN_EDGES_DAYS[:-1], dtype=z_post.dtype, device=z_post.device)) |
| out["pwe_bin_survival_probs"] = pwe_ci["survival"] |
| out["pwe_bin_cif_probs"] = pwe_ci["cif"] |
| return out |
|
|
| def _rollout_prefix_embedding(self, out: dict[str, torch.Tensor], temperature: float) -> torch.Tensor: |
| cause = F.gumbel_softmax(out["pwe_log_lambda"].amax(dim=-1), tau=max(temperature, 1.0e-3), hard=False, dim=-1) |
| active = F.gumbel_softmax(out["active_state_logits"], tau=max(temperature, 1.0e-3), hard=False, dim=-1) |
| event = torch.sigmoid(out["risk_event_logits"]) |
| delta_mu = torch.log1p(F.softplus(out["pwe_log_lambda"]).sum(dim=(2, 3)).reciprocal().clamp(max=365.0)) |
| next_contact_prob = cause[..., :1] |
| features = torch.cat([cause, active, event, delta_mu.unsqueeze(-1), next_contact_prob], dim=-1) |
| return self.rollout_prefix_proj(features) |
|
|
| def forward( |
| self, |
| batch: dict[str, torch.Tensor], |
| rollout_steps: int = 1, |
| teacher_forcing_rate: float = 1.0, |
| gumbel_temperature: float = 1.0, |
| ) -> dict[str, torch.Tensor]: |
| u, conditional_status = self.tab_encoder( |
| batch["cat_value_ids"], |
| batch["missing_ids"], |
| batch["numeric_values"], |
| batch["numeric_mask"], |
| batch["ordinal_cbe"], |
| batch["ordinal_mask"], |
| batch["drug_name_ids"], |
| batch["drug_class_ids"], |
| batch["drug_mask"], |
| batch["service_state"], |
| batch["visit_year"], |
| ) |
| z_history, z_post = self.traj_decoder(u, batch["service_state"], batch["visit_year"], batch["time_since_start_days"], batch["valid_mask"]) |
| static_vec = self.static_encoder(batch.get("static_value_ids"), z_post) |
| static_injection = self.static_injection_scale * static_vec |
| z_history = z_history + static_injection[:, None, :] |
| z_post = z_post + static_injection[:, None, :] |
| z_contact = self.contact_injector(z_history, batch) |
| out = self._heads_from_states(z_history, z_contact, z_post, batch) |
| out["u"] = u |
| out["conditional_status"] = conditional_status |
| out["static_vec"] = static_vec |
| out["static_injection_scale"] = self.static_injection_scale |
| out["contact_injection_scale"] = self.contact_injector.scale |
|
|
| if rollout_steps > 1: |
| |
| |
| |
| |
| rollout_logits: list[torch.Tensor] = [] |
| rollout_source = out |
| eps = float(max(0.0, min(1.0, teacher_forcing_rate))) |
| for _ in range(2, int(rollout_steps) + 1): |
| prefix = self._rollout_prefix_embedding(rollout_source, gumbel_temperature) |
| shifted_prefix = torch.cat([u[:, :1, :], prefix[:, :-1, :]], dim=1) |
| u_mix = eps * u + (1.0 - eps) * shifted_prefix |
| rz_history, rz_post = self.traj_decoder( |
| u_mix, |
| batch["service_state"], |
| batch["visit_year"], |
| batch["time_since_start_days"], |
| batch["valid_mask"], |
| ) |
| rz_history = rz_history + static_injection[:, None, :] |
| rz_post = rz_post + static_injection[:, None, :] |
| rz_contact = self.contact_injector(rz_history, batch) |
| r_out = self._heads_from_states(rz_history, rz_contact, rz_post, batch) |
| rollout_source = r_out |
| rollout_logits.append(r_out["pwe_log_lambda"]) |
| out["_last_rollout_context"] = r_out["z_post"] |
| out["rollout_pwe_log_lambda"] = torch.stack(rollout_logits, dim=0) if rollout_logits else out["pwe_log_lambda"].new_zeros(0) |
| return out |
|
|
| @torch.no_grad() |
| def predict_fixed_horizon_risk( |
| self, |
| batch: dict[str, torch.Tensor], |
| horizons_days: torch.Tensor | list[float], |
| mode: str = "contact", |
| ) -> dict[str, torch.Tensor]: |
| out = self.forward(batch) |
| if mode in {"pre_visit", "history", "pre"}: |
| log_lambda = out["pwe_log_lambda_pre"] |
| elif mode in {"contact", "contact_present"}: |
| log_lambda = out["pwe_log_lambda_contact"] |
| elif mode in {"field_end", "post", "post_visit"}: |
| log_lambda = out["pwe_log_lambda_post"] |
| else: |
| raise ValueError(f"unknown fixed-horizon risk mode: {mode}") |
| h = torch.as_tensor(horizons_days, dtype=log_lambda.dtype, device=log_lambda.device) |
| return pwe_closed_form_cif(log_lambda, h) |
|
|
|
|
| def pwe_closed_form_cif(log_lambda_kj: torch.Tensor, horizons_days: torch.Tensor) -> dict[str, torch.Tensor]: |
| """Closed-form PWE survival and cause-specific cumulative incidence. |
| |
| Args: |
| log_lambda_kj: [..., K, J] raw log-rate parameters. |
| horizons_days: [H] positive horizons in days. |
| """ |
|
|
| if horizons_days.ndim == 0: |
| horizons_days = horizons_days[None] |
| |
| |
| |
| |
| calc_dtype = torch.float32 if log_lambda_kj.dtype in {torch.float16, torch.bfloat16} else log_lambda_kj.dtype |
| lambda_kj = F.softplus(log_lambda_kj.to(calc_dtype)) + 1.0e-8 |
| lambda_j = lambda_kj.sum(dim=-2) |
| lower = torch.tensor(PWE_FINITE_LOWER_DAYS, dtype=lambda_kj.dtype, device=lambda_kj.device) |
| upper = torch.tensor(PWE_BIN_EDGES_DAYS[:-1], dtype=lambda_kj.dtype, device=lambda_kj.device) |
| h = horizons_days.to(dtype=lambda_kj.dtype, device=lambda_kj.device).clamp(min=0.0) |
| finite_exposure = (torch.minimum(h[:, None], upper[None, :]) - lower[:-1][None, :]).clamp(min=0.0) |
| last_exposure = (h - lower[-1]).clamp(min=0.0)[:, None] |
| exposure = torch.cat([finite_exposure, last_exposure], dim=1) |
| base_ndim = lambda_j.ndim - 1 |
| hazard_exposure = lambda_j.unsqueeze(-2) * exposure.reshape((1,) * base_ndim + exposure.shape) |
| cum_before = torch.cumsum(hazard_exposure, dim=-1) - hazard_exposure |
| survival_before = torch.exp(-cum_before) |
| event_in_bin = -torch.expm1(-hazard_exposure) |
| share = lambda_kj / lambda_j.unsqueeze(-2).clamp(min=1.0e-8) |
| cif = (share.unsqueeze(-3) * survival_before.unsqueeze(-2) * event_in_bin.unsqueeze(-2)).sum(dim=-1) |
| survival = torch.exp(-hazard_exposure.sum(dim=-1)).clamp(min=0.0, max=1.0) |
| max_cif_total = (1.0 - survival).unsqueeze(-1) |
| cif_total = cif.sum(dim=-1, keepdim=True) |
| scale = torch.where(cif_total > max_cif_total, max_cif_total / cif_total.clamp(min=1.0e-12), torch.ones_like(cif_total)) |
| cif = (cif * scale).clamp(min=0.0, max=1.0) |
| return {"survival": survival, "cif": cif} |
|
|
|
|
| __all__ = [ |
| "SCTMv4p3", |
| "SCTMv4p3Config", |
| "PWE_BIN_EDGES_DAYS", |
| "build_service_attention_prior", |
| "era_tokens", |
| "pwe_closed_form_cif", |
| "previous_interval_log", |
| ] |
|
|