from __future__ import annotations import sys from dataclasses import dataclass from pathlib import Path from typing import Any import torch import torch.nn as nn SCRIPT_DIR = Path(__file__).resolve().parent ROOT_DIR = SCRIPT_DIR.parents[1] V4P4_SCRIPT_DIR = ROOT_DIR / "v4p4_world_model" / "scripts" if str(V4P4_SCRIPT_DIR) not in sys.path: sys.path.insert(0, str(V4P4_SCRIPT_DIR)) from model_v4p4 import SCTMv4p4, SCTMv4p4Config, pwe_closed_form_cif # noqa: E402 @dataclass class SCTMv5Config(SCTMv4p4Config): n_action_slots: int = 16 action_value_vocab_size: int = 256 n_action_availability: int = 4 class ActionEncoder(nn.Module): """Slot-wise action encoder for v5 observed-action and strategy-conditioned modeling.""" def __init__(self, config: SCTMv5Config) -> None: super().__init__() d = int(config.d_model) self.config = config self.type_emb = nn.Embedding(config.n_action_slots, d) self.value_emb = nn.Embedding(config.action_value_vocab_size, d, padding_idx=0) self.availability_emb = nn.Embedding(config.n_action_availability, d) self.slot_emb = nn.Embedding(config.n_action_slots, d) self.norm = nn.LayerNorm(d) self.proj = nn.Sequential(nn.Linear(d, d), nn.GELU(), nn.Dropout(config.dropout), nn.Linear(d, d)) self.register_buffer("default_type_ids", torch.arange(config.n_action_slots, dtype=torch.long), persistent=False) def forward(self, batch: dict[str, torch.Tensor]) -> torch.Tensor: if "action_value_ids" not in batch: bsz, seq_len = batch["valid_mask"].shape device = batch["valid_mask"].device type_ids = self.default_type_ids.to(device).view(1, 1, -1).expand(bsz, seq_len, -1) value_ids = torch.zeros_like(type_ids) availability = torch.zeros_like(type_ids) mask = torch.zeros_like(type_ids, dtype=torch.bool) else: value_ids = batch["action_value_ids"].long().clamp(min=0, max=self.config.action_value_vocab_size - 1) bsz, seq_len, n_slots = value_ids.shape type_ids = batch.get("action_type_ids") if type_ids is None: type_ids = self.default_type_ids.to(value_ids.device).view(1, 1, -1).expand(bsz, seq_len, -1) type_ids = type_ids.long().clamp(min=0, max=self.config.n_action_slots - 1) availability = batch.get("action_available_at") if availability is None: availability = torch.zeros_like(value_ids) availability = availability.long().clamp(min=0, max=self.config.n_action_availability - 1) mask = batch.get("action_mask") if mask is None: mask = value_ids.ne(0) mask = mask.bool() if n_slots != self.config.n_action_slots: raise ValueError(f"Expected {self.config.n_action_slots} action slots, found {n_slots}") slot_ids = self.default_type_ids.to(value_ids.device).view(1, 1, -1).expand_as(value_ids) emb = self.type_emb(type_ids) + self.value_emb(value_ids) + self.availability_emb(availability) + self.slot_emb(slot_ids) emb = self.norm(emb) masked = emb * mask.unsqueeze(-1).to(emb.dtype) denom = mask.sum(dim=-1, keepdim=True).clamp(min=1).to(emb.dtype) pooled = masked.sum(dim=-2) / denom return self.proj(pooled) def zero_init_last_linear(module: nn.Module) -> None: for child in reversed(list(module.modules())): if isinstance(child, nn.Linear): nn.init.zeros_(child.weight) nn.init.zeros_(child.bias) return class SCTMv5(SCTMv4p4): """SCTM-v5 action-conditioned target-trial-aware care-process world model. The v5 branch leaves v4 untouched. It reuses the v4 backbone and adds action-conditioned residual heads plus a behavior-policy head for propensity/overlap diagnostics. The observed-action likelihood is distinct from causal estimands; target-trial scripts define those separately. """ def __init__( self, config: SCTMv5Config, field_value_mask: torch.Tensor, service_prior_bias: torch.Tensor | None = None, ) -> None: super().__init__(config, field_value_mask=field_value_mask, service_prior_bias=service_prior_bias) self.config: SCTMv5Config = config d = int(config.d_model) self.action_encoder = ActionEncoder(config) self.action_pwe_delta_head = nn.Sequential(nn.LayerNorm(d), nn.Linear(d, d), nn.GELU(), nn.Dropout(config.dropout), nn.Linear(d, config.n_pwe_causes * config.n_pwe_bins)) self.action_active_delta_head = nn.Sequential(nn.LayerNorm(d), nn.Linear(d, d), nn.GELU(), nn.Dropout(config.dropout), nn.Linear(d, config.n_active_states)) self.action_missing_delta_head = nn.Sequential(nn.LayerNorm(d), nn.Linear(d, d), nn.GELU(), nn.Dropout(config.dropout), nn.Linear(d, config.n_cat_fields * config.n_missing)) self.action_field_delta_head = nn.Sequential(nn.LayerNorm(d), nn.Linear(d, d), nn.GELU(), nn.Dropout(config.dropout), nn.Linear(d, config.n_cat_fields * config.cat_vocab_size)) self.action_numeric_delta_head = nn.Sequential(nn.LayerNorm(d), nn.Linear(d, d), nn.GELU(), nn.Dropout(config.dropout), nn.Linear(d, config.n_numeric_fields)) self.action_ordinal_delta_head = nn.Sequential(nn.LayerNorm(d), nn.Linear(d, d), nn.GELU(), nn.Dropout(config.dropout), nn.Linear(d, config.n_ordinal_fields * config.cbe_dim)) self.action_event_delta_head = nn.Sequential(nn.LayerNorm(d), nn.Linear(d, d), nn.GELU(), nn.Dropout(config.dropout), nn.Linear(d, config.n_events)) self.behavior_policy_head = nn.Sequential( nn.LayerNorm(d), nn.Linear(d, d), nn.GELU(), nn.Dropout(config.dropout), nn.Linear(d, config.n_action_slots * config.action_value_vocab_size), ) for module in ( self.action_pwe_delta_head, self.action_active_delta_head, self.action_missing_delta_head, self.action_field_delta_head, self.action_numeric_delta_head, self.action_ordinal_delta_head, self.action_event_delta_head, ): zero_init_last_linear(module) def _action_residual_l2(self) -> torch.Tensor: total = None for module in ( self.action_pwe_delta_head, self.action_active_delta_head, self.action_missing_delta_head, self.action_field_delta_head, self.action_numeric_delta_head, self.action_ordinal_delta_head, self.action_event_delta_head, ): for param in module.parameters(): value = (param.float() ** 2).mean() total = value if total is None else total + value assert total is not None return total def forward( self, batch: dict[str, torch.Tensor], *args: Any, compute_action_conditioned: bool = True, **kwargs: Any, ) -> dict[str, torch.Tensor]: out = super().forward(batch, *args, **kwargs) if not compute_action_conditioned: return out z_post = out["z_post"] bsz, seq_len, _ = z_post.shape action_ctx = self.action_encoder(batch) z_action = z_post + action_ctx time_ctx_action = self._teacher_forced_next_time_context(z_post, batch) + action_ctx active_ids = self._teacher_forced_next_active_ids(batch) obs_ctx_action = self._condition_on_next_active(time_ctx_action, active_ids) pwe_delta = self.action_pwe_delta_head(z_action).view(bsz, seq_len, self.config.n_pwe_causes, self.config.n_pwe_bins) active_delta = self.action_active_delta_head(time_ctx_action) missing_delta = self.action_missing_delta_head(obs_ctx_action).view(bsz, seq_len, self.config.n_cat_fields, self.config.n_missing) field_delta = self.action_field_delta_head(obs_ctx_action).view(bsz, seq_len, self.config.n_cat_fields, self.config.cat_vocab_size) numeric_delta = self.action_numeric_delta_head(obs_ctx_action).view(bsz, seq_len, self.config.n_numeric_fields) ordinal_delta = self.action_ordinal_delta_head(obs_ctx_action).view(bsz, seq_len, self.config.n_ordinal_fields, self.config.cbe_dim) event_delta = self.action_event_delta_head(obs_ctx_action) # Propensity diagnostics must not see the action-bearing clinical fields. # z_contact is the conservative current-contact state before field/end-of-visit actions. behavior_policy_logits = self.behavior_policy_head(out["z_contact"]).view( bsz, seq_len, self.config.n_action_slots, self.config.action_value_vocab_size, ) out.update( { "action_context": action_ctx, "z_action": z_action, "pwe_log_lambda_action": out["pwe_log_lambda_post"] + pwe_delta, "active_state_logits_action": self.active_state_head(time_ctx_action) + active_delta, "missingness_logits_action": self.missingness_head(obs_ctx_action).view(bsz, seq_len, self.config.n_cat_fields, self.config.n_missing) + missing_delta, "field_logits_action": self.field_head(obs_ctx_action) + field_delta, "numeric_mu_action": self.numeric_head(obs_ctx_action) + numeric_delta, "ordinal_cum_logits_action": self.ordinal_head(obs_ctx_action) + ordinal_delta, "event_generation_logits_action": self.event_generation_head(obs_ctx_action) + event_delta, "behavior_policy_logits": behavior_policy_logits, "v5_residual_l2": self._action_residual_l2(), } ) return out @torch.no_grad() def _write_generated_visit(self, batch: dict[str, torch.Tensor], generated: dict[str, torch.Tensor], position: int, alive: torch.Tensor) -> None: """Write generated visit tensors and clear future observed-action tensors. v4 writes only visit/observation tensors. In v5, leaving the original action tensors in place would make free-running rollout condition on the real future actions from the sampled trajectory. Generated visits do not have observed actions unless an explicit strategy/policy writer supplies them, so the safe default is a no-observed-action landmark. """ super()._write_generated_visit(batch, generated, position, alive) next_pos = position + 1 if "action_value_ids" not in batch: return bsz = batch["action_value_ids"].shape[0] device = batch["action_value_ids"].device action_mask = alive.reshape(bsz, 1) if "action_type_ids" in batch: default_types = self.action_encoder.default_type_ids.to(device).view(1, -1).expand(bsz, -1) batch["action_type_ids"][:, next_pos, :] = torch.where( action_mask, default_types.to(batch["action_type_ids"].dtype), batch["action_type_ids"][:, next_pos, :], ) batch["action_value_ids"][:, next_pos, :] = torch.where( action_mask, torch.zeros_like(batch["action_value_ids"][:, next_pos, :]), batch["action_value_ids"][:, next_pos, :], ) if "action_mask" in batch: batch["action_mask"][:, next_pos, :] = torch.where( action_mask, torch.zeros_like(batch["action_mask"][:, next_pos, :]), batch["action_mask"][:, next_pos, :], ) if "action_available_at" in batch: batch["action_available_at"][:, next_pos, :] = torch.where( action_mask, torch.zeros_like(batch["action_available_at"][:, next_pos, :]), batch["action_available_at"][:, next_pos, :], ) @torch.no_grad() def sample_next_visit( self, batch: dict[str, torch.Tensor], out: dict[str, torch.Tensor], position: int, *, deterministic: bool = False, generator: torch.Generator | None = None, forced_cause: torch.Tensor | None = None, forced_delta_days: torch.Tensor | float | None = None, forced_active_state: torch.Tensor | int | None = None, max_time_days: float | None = 3650.0, rao_blackwell_rare: bool = True, missingness_logit_bias: torch.Tensor | None = None, missingness_prior_probs: torch.Tensor | None = None, missingness_prior_blend: float = 0.0, ) -> dict[str, torch.Tensor]: """Sample visit[t+1] using v5 action-conditioned dynamics. The inherited v4 rollout loop calls this method dynamically. Overriding it keeps closed-loop v5 evaluation action-conditioned whenever action tensors are present in the generated batch. """ if position < 0 or position >= batch["valid_mask"].shape[1] - 1: raise ValueError("position must leave room for a generated next visit") bsz = batch["valid_mask"].shape[0] device = batch["valid_mask"].device z_post_t = out["z_post"][:, position, :] action_ctx_t = out.get("action_context") if action_ctx_t is None: action_ctx = self.action_encoder(batch)[:, position, :] else: action_ctx = action_ctx_t[:, position, :] z_action = z_post_t + action_ctx pwe_log_lambda = out.get("pwe_log_lambda_action") if pwe_log_lambda is None: pwe_delta = self.action_pwe_delta_head(z_action).view(bsz, self.config.n_pwe_causes, self.config.n_pwe_bins) pwe_log_lambda_t = out["pwe_log_lambda_post"][:, position, :, :] + pwe_delta else: pwe_log_lambda_t = pwe_log_lambda[:, position, :, :] pwe = self.sample_pwe_event_time( pwe_log_lambda_t, deterministic=deterministic, generator=generator, max_time_days=max_time_days, rao_blackwell_rare=rao_blackwell_rare, ) cause = pwe["cause"] if forced_cause is not None: cause = forced_cause.to(device=device, dtype=torch.long).reshape(bsz).clamp(0, self.config.n_pwe_causes - 1) pwe["cause"] = cause pwe["absorbed"] = cause.ne(0) delta_days = pwe["delta_days"].to(device=device) if forced_delta_days is not None: if torch.is_tensor(forced_delta_days): delta_days = forced_delta_days.to(device=device, dtype=delta_days.dtype).reshape(bsz) else: delta_days = torch.full((bsz,), float(forced_delta_days), dtype=delta_days.dtype, device=device) pwe["delta_days"] = delta_days delta_log = torch.log1p(delta_days.clamp(min=1.0e-6)).to(z_post_t.dtype) time_ctx = z_post_t + self.next_delta_condition(delta_log[:, None].float()) + action_ctx active_logits = self.active_state_head(time_ctx) + self.action_active_delta_head(time_ctx) active_state = self._sample_categorical(active_logits, deterministic=deterministic, generator=generator).clamp(0, self.config.n_active_states - 1) if forced_active_state is not None: if torch.is_tensor(forced_active_state): active_state = forced_active_state.to(device=device, dtype=torch.long).reshape(bsz).clamp(0, self.config.n_active_states - 1) else: active_state = torch.full((bsz,), int(forced_active_state), dtype=torch.long, device=device).clamp(0, self.config.n_active_states - 1) service_state = torch.where( cause.eq(1), torch.full_like(active_state, 7), torch.where(cause.eq(2), torch.full_like(active_state, 6), active_state), ) next_contact = cause.eq(0) active_ids = torch.where(next_contact, active_state, torch.full_like(active_state, self.config.n_active_states)) obs_ctx = self._condition_on_next_active(time_ctx, active_ids) missing_logits = self.missingness_head(obs_ctx).view(bsz, self.config.n_cat_fields, self.config.n_missing) missing_logits = missing_logits + self.action_missing_delta_head(obs_ctx).view(bsz, self.config.n_cat_fields, self.config.n_missing) if missingness_logit_bias is not None: bias = missingness_logit_bias.to(device=device, dtype=missing_logits.dtype).view(1, 1, self.config.n_missing) missing_logits = missing_logits + bias blocked_contact_missing_ids = ( int(self.config.missing_no_clinical_id), int(self.config.missing_visit_missing_id), ) if any(0 <= missing_id < self.config.n_missing for missing_id in blocked_contact_missing_ids): contact_missing_logits = missing_logits.clone() for missing_id in blocked_contact_missing_ids: if 0 <= missing_id < self.config.n_missing: contact_missing_logits[..., missing_id] = -1.0e9 missing_logits = torch.where(next_contact[:, None, None], contact_missing_logits, missing_logits) if missingness_prior_probs is not None and float(missingness_prior_blend) > 0.0: prior = missingness_prior_probs.to(device=device, dtype=torch.float32).view(1, self.config.n_cat_fields, self.config.n_missing) prior = prior.clamp_min(0.0) contact_prior = prior.clone() for missing_id in blocked_contact_missing_ids: if 0 <= missing_id < self.config.n_missing: contact_prior[..., missing_id] = 0.0 denom = contact_prior.sum(dim=-1, keepdim=True) fallback = torch.zeros_like(contact_prior) fallback[..., self.config.missing_observed_id] = 1.0 contact_prior = torch.where(denom.gt(0.0), contact_prior / denom.clamp_min(1.0e-12), fallback) model_probs = torch.softmax(missing_logits.float(), dim=-1) blend = min(max(float(missingness_prior_blend), 0.0), 1.0) mixed_probs = ((1.0 - blend) * model_probs + blend * contact_prior).clamp_min(1.0e-12) mixed_logits = torch.log(mixed_probs).to(missing_logits.dtype) missing_logits = torch.where(next_contact[:, None, None], mixed_logits, missing_logits) missing_ids = self._sample_categorical(missing_logits, deterministic=deterministic, generator=generator) missing_ids = torch.where( next_contact[:, None], missing_ids, torch.full_like(missing_ids, self.config.missing_no_clinical_id), ) field_logits = self.field_head(obs_ctx[:, None, :]).squeeze(1) field_logits = field_logits + self.action_field_delta_head(obs_ctx).view(bsz, self.config.n_cat_fields, self.config.cat_vocab_size) cat_value_ids = self._sample_categorical(field_logits, deterministic=deterministic, generator=generator) cat_value_ids = torch.where( missing_ids.eq(self.config.missing_observed_id), cat_value_ids, torch.full_like(cat_value_ids, self.config.unknown_cat_value_id), ) cat_value_ids, missing_ids = self._apply_generated_observation_constraints(cat_value_ids, missing_ids, next_contact) numeric_mu = self.numeric_head(obs_ctx)[..., : self.config.n_numeric_fields] + self.action_numeric_delta_head(obs_ctx) if self.config.n_numeric_fields: numeric_mask = next_contact[:, None].expand(bsz, self.config.n_numeric_fields) numeric_values = torch.where(numeric_mask, numeric_mu, torch.zeros_like(numeric_mu)) else: numeric_values = torch.zeros((bsz, 0), dtype=obs_ctx.dtype, device=device) numeric_mask = torch.zeros((bsz, 0), dtype=torch.bool, device=device) ordinal_logits = self.ordinal_head(obs_ctx[:, None, :]).squeeze(1) ordinal_logits = ordinal_logits + self.action_ordinal_delta_head(obs_ctx).view(bsz, self.config.n_ordinal_fields, self.config.cbe_dim) ordinal_mask = next_contact[:, None].expand(bsz, self.config.n_ordinal_fields) ordinal_cbe = torch.sigmoid(ordinal_logits).ge(0.5) & ordinal_mask[:, :, None] generation_event_logits = self.event_generation_head(obs_ctx) + self.action_event_delta_head(obs_ctx) event_prob = torch.sigmoid(generation_event_logits) if deterministic: event_labels = event_prob.ge(0.5).to(event_prob.dtype) else: event_labels = torch.bernoulli(event_prob.float(), generator=generator).to(event_prob.dtype) event_labels = torch.where(cause[:, None].eq(0), event_labels, torch.zeros_like(event_labels)) terminal_label = torch.where( cause.eq(1), torch.ones_like(cause), torch.where(cause.eq(2), torch.full_like(cause, 2), torch.zeros_like(cause)), ) current_time = batch["time_since_start_days"][:, position].to(delta_days.dtype) next_time = current_time + delta_days start_year = batch["visit_year"][:, 0].long() visit_year = (start_year + torch.floor(next_time / 365.25).long()).clamp(min=self.config.year_min, max=self.config.year_max) next_visit_index = batch.get("visit_indices", torch.zeros_like(batch["service_state"]))[:, position].long() + 1 return { **pwe, "delta_t_next_log_current": delta_log, "cat_value_ids": cat_value_ids.long(), "missing_ids": missing_ids.long(), "numeric_values": numeric_values.to(batch["numeric_values"].dtype), "numeric_mask": numeric_mask, "ordinal_cbe": ordinal_cbe.to(batch["ordinal_cbe"].dtype), "ordinal_mask": ordinal_mask, "drug_name_ids": torch.where(next_contact[:, None], batch["drug_name_ids"][:, position, :], torch.zeros_like(batch["drug_name_ids"][:, position, :])), "drug_class_ids": torch.where(next_contact[:, None], batch["drug_class_ids"][:, position, :], torch.zeros_like(batch["drug_class_ids"][:, position, :])), "drug_mask": torch.where(next_contact[:, None], batch["drug_mask"][:, position, :], torch.zeros_like(batch["drug_mask"][:, position, :])), "service_state": service_state.long(), "terminal_label": terminal_label.long(), "event_labels": event_labels.to(batch["event_labels"].dtype), "time_since_start_days": next_time.to(batch["time_since_start_days"].dtype), "visit_year": visit_year.to(batch["visit_year"].dtype), "visit_indices": next_visit_index.to(batch.get("visit_indices", batch["service_state"]).dtype), "active_state_logits": active_logits, "missingness_logits": missing_logits, "event_generation_logits": generation_event_logits, "action_context": action_ctx, } __all__ = [ "ActionEncoder", "SCTMv5", "SCTMv5Config", "pwe_closed_form_cif", ]