Implement shared event-trajectory reasoning backbone
This commit is contained in:
483
models.py
483
models.py
@@ -7,39 +7,197 @@ import torch.nn.functional as F
|
||||
|
||||
from backbones import (
|
||||
AgeSinusoidalEncoding,
|
||||
GPTBlock,
|
||||
GaussianRBFTimeBasis,
|
||||
SharedEventTrajectoryCore,
|
||||
TimeRoPE,
|
||||
TokenAutoDiscretization,
|
||||
)
|
||||
from targets import PAD_IDX
|
||||
|
||||
|
||||
TRAJ_MIXER_ARCHITECTURE = "traj_mixer_v2"
|
||||
EVENT_TRAJECTORY_ARCHITECTURE = "event_trajectory_shared_v1"
|
||||
|
||||
|
||||
def validate_traj_mixer_config(config: Mapping[str, object]) -> None:
|
||||
actual = config.get("model_architecture")
|
||||
if actual != TRAJ_MIXER_ARCHITECTURE:
|
||||
@dataclass(frozen=True)
|
||||
class EventTrajectoryModelSize:
|
||||
d_model: int
|
||||
n_trajectory: int
|
||||
|
||||
@property
|
||||
def trajectory_dim(self) -> int:
|
||||
return self.d_model // self.n_trajectory
|
||||
|
||||
@property
|
||||
def traj_hidden(self) -> int:
|
||||
return 4 * self.n_trajectory
|
||||
|
||||
|
||||
MODEL_SIZE_PRESETS = {
|
||||
"nano": EventTrajectoryModelSize(d_model=256, n_trajectory=8),
|
||||
"small": EventTrajectoryModelSize(d_model=512, n_trajectory=8),
|
||||
"medium": EventTrajectoryModelSize(d_model=768, n_trajectory=12),
|
||||
"huge": EventTrajectoryModelSize(d_model=1024, n_trajectory=16),
|
||||
}
|
||||
MODEL_SIZE_NAMES = tuple(MODEL_SIZE_PRESETS)
|
||||
|
||||
|
||||
def resolve_model_size(model_size: str) -> EventTrajectoryModelSize:
|
||||
if not isinstance(model_size, str):
|
||||
raise ValueError(
|
||||
"This branch only accepts models trained with the TrajMixer "
|
||||
f"architecture marker {TRAJ_MIXER_ARCHITECTURE!r}; got {actual!r}."
|
||||
f"model_size must be a string, got {type(model_size).__name__}"
|
||||
)
|
||||
normalized = model_size.strip().lower()
|
||||
try:
|
||||
return MODEL_SIZE_PRESETS[normalized]
|
||||
except KeyError as exc:
|
||||
choices = ", ".join(MODEL_SIZE_NAMES)
|
||||
raise ValueError(
|
||||
f"Unknown model_size {model_size!r}; expected one of: {choices}"
|
||||
) from exc
|
||||
|
||||
|
||||
def _required_config_int(
|
||||
config: Mapping[str, object],
|
||||
key: str,
|
||||
) -> int:
|
||||
raw_value = config.get(key)
|
||||
if isinstance(raw_value, bool):
|
||||
raise ValueError(f"Config field {key!r} must be an integer")
|
||||
try:
|
||||
value = int(raw_value)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError(
|
||||
f"Config field {key!r} must be present and integer-valued; "
|
||||
f"got {raw_value!r}"
|
||||
) from exc
|
||||
if isinstance(raw_value, float) and not raw_value.is_integer():
|
||||
raise ValueError(f"Config field {key!r} must be an integer")
|
||||
return value
|
||||
|
||||
|
||||
def validate_event_trajectory_config(config: Mapping[str, object]) -> None:
|
||||
actual = config.get("model_architecture")
|
||||
if actual != EVENT_TRAJECTORY_ARCHITECTURE:
|
||||
raise ValueError(
|
||||
"This branch only accepts models trained with the shared "
|
||||
"event-trajectory architecture marker "
|
||||
f"{EVENT_TRAJECTORY_ARCHITECTURE!r}; got {actual!r}."
|
||||
)
|
||||
raw_model_size = config.get("model_size")
|
||||
if not isinstance(raw_model_size, str):
|
||||
raise ValueError(
|
||||
"Config field 'model_size' must be one of: "
|
||||
+ ", ".join(MODEL_SIZE_NAMES)
|
||||
)
|
||||
model_size = raw_model_size.strip().lower()
|
||||
preset = resolve_model_size(model_size)
|
||||
d_model = _required_config_int(config, "d_model")
|
||||
n_trajectory = _required_config_int(config, "n_trajectory")
|
||||
n_reasoning_rounds = _required_config_int(
|
||||
config,
|
||||
"n_reasoning_rounds",
|
||||
)
|
||||
trajectory_dim = _required_config_int(config, "trajectory_dim")
|
||||
traj_hidden = _required_config_int(config, "traj_hidden")
|
||||
if n_reasoning_rounds <= 0:
|
||||
raise ValueError(
|
||||
"n_reasoning_rounds must be positive"
|
||||
)
|
||||
expected_values = {
|
||||
"d_model": preset.d_model,
|
||||
"n_trajectory": preset.n_trajectory,
|
||||
"trajectory_dim": preset.trajectory_dim,
|
||||
"traj_hidden": preset.traj_hidden,
|
||||
}
|
||||
actual_values = {
|
||||
"d_model": d_model,
|
||||
"n_trajectory": n_trajectory,
|
||||
"trajectory_dim": trajectory_dim,
|
||||
"traj_hidden": traj_hidden,
|
||||
}
|
||||
mismatches = [
|
||||
f"{key}: expected {expected}, got {actual_values[key]}"
|
||||
for key, expected in expected_values.items()
|
||||
if actual_values[key] != expected
|
||||
]
|
||||
if mismatches:
|
||||
raise ValueError(
|
||||
f"Config does not match model_size={model_size!r}: "
|
||||
+ "; ".join(mismatches)
|
||||
)
|
||||
|
||||
|
||||
def validate_traj_mixer_state_dict(state_dict: Mapping[str, object]) -> None:
|
||||
def _checkpoint_scalar_int(
|
||||
state_dict: Mapping[str, object],
|
||||
key: str,
|
||||
) -> int:
|
||||
value = state_dict[key]
|
||||
if not isinstance(value, torch.Tensor) or value.numel() != 1:
|
||||
raise ValueError(
|
||||
f"Checkpoint architecture field {key!r} must be a scalar tensor"
|
||||
)
|
||||
return int(value.detach().cpu().item())
|
||||
|
||||
|
||||
def validate_event_trajectory_state_dict(
|
||||
state_dict: Mapping[str, object],
|
||||
*,
|
||||
expected_d_model: int | None = None,
|
||||
expected_n_trajectory: int | None = None,
|
||||
expected_n_reasoning_rounds: int | None = None,
|
||||
) -> None:
|
||||
required_keys = {
|
||||
"blocks.0.mlp.group_align",
|
||||
"blocks.0.mlp.gate_proj",
|
||||
"blocks.0.mlp.value_proj",
|
||||
"blocks.0.mlp.output_proj",
|
||||
"architecture_d_model",
|
||||
"architecture_n_trajectory",
|
||||
"architecture_n_reasoning_rounds",
|
||||
"event_projection.weight",
|
||||
"trajectory_prototypes",
|
||||
"query_projection.weight",
|
||||
"reasoning_core.cross_attention.q_proj.weight",
|
||||
"reasoning_core.cross_attention.k_proj.weight",
|
||||
"reasoning_core.cross_attention.v_proj.weight",
|
||||
"reasoning_core.traj_mixer.gate_proj",
|
||||
"reasoning_core.traj_mixer.value_proj",
|
||||
"reasoning_core.traj_mixer.output_proj",
|
||||
"reasoning_core.attn_scale",
|
||||
"reasoning_core.mixer_scale",
|
||||
}
|
||||
missing = sorted(required_keys.difference(state_dict))
|
||||
if missing:
|
||||
raise ValueError(
|
||||
"Checkpoint is not a TrajMixer checkpoint; missing required "
|
||||
"Checkpoint is not a shared event-trajectory checkpoint; "
|
||||
"missing required "
|
||||
f"parameters: {', '.join(missing)}"
|
||||
)
|
||||
checkpoint_values = {
|
||||
"d_model": _checkpoint_scalar_int(
|
||||
state_dict,
|
||||
"architecture_d_model",
|
||||
),
|
||||
"n_trajectory": _checkpoint_scalar_int(
|
||||
state_dict,
|
||||
"architecture_n_trajectory",
|
||||
),
|
||||
"n_reasoning_rounds": _checkpoint_scalar_int(
|
||||
state_dict,
|
||||
"architecture_n_reasoning_rounds",
|
||||
),
|
||||
}
|
||||
expected_values = {
|
||||
"d_model": expected_d_model,
|
||||
"n_trajectory": expected_n_trajectory,
|
||||
"n_reasoning_rounds": expected_n_reasoning_rounds,
|
||||
}
|
||||
mismatches = [
|
||||
f"{name}: checkpoint={checkpoint_values[name]}, expected={expected}"
|
||||
for name, expected in expected_values.items()
|
||||
if expected is not None and checkpoint_values[name] != expected
|
||||
]
|
||||
if mismatches:
|
||||
raise ValueError(
|
||||
"Checkpoint architecture does not match the constructed model: "
|
||||
+ "; ".join(mismatches)
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -173,10 +331,8 @@ class DeepHealth(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size: int,
|
||||
n_embd: int,
|
||||
n_head: int,
|
||||
n_hist_layer: int,
|
||||
n_tab_layer: int,
|
||||
model_size: str,
|
||||
n_reasoning_rounds: int,
|
||||
n_types: int,
|
||||
n_cont_types: int,
|
||||
n_categories: int,
|
||||
@@ -201,11 +357,21 @@ class DeepHealth(nn.Module):
|
||||
"dist_mode must be either 'exponential', 'weibull' or 'mixed'")
|
||||
if extra_pool_reduce not in {"mean", "sum"}:
|
||||
raise ValueError("extra_pool_reduce must be either 'mean' or 'sum'")
|
||||
self.token_embedding = nn.Embedding(vocab_size, n_embd, padding_idx=0)
|
||||
if n_reasoning_rounds <= 0:
|
||||
raise ValueError(
|
||||
"n_reasoning_rounds must be positive, got "
|
||||
f"{n_reasoning_rounds}"
|
||||
)
|
||||
size_config = resolve_model_size(model_size)
|
||||
normalized_model_size = model_size.strip().lower()
|
||||
d_model = size_config.d_model
|
||||
n_trajectory = size_config.n_trajectory
|
||||
|
||||
self.token_embedding = nn.Embedding(vocab_size, d_model, padding_idx=0)
|
||||
self.gender_embedding = nn.Embedding(
|
||||
2, n_embd) # Assuming binary gender
|
||||
2, d_model) # Assuming binary gender
|
||||
self.tokenizer = OtherInfoTokenizer(
|
||||
n_embd=n_embd,
|
||||
n_embd=d_model,
|
||||
n_types=n_types,
|
||||
n_cont_types=n_cont_types,
|
||||
n_categories=n_categories,
|
||||
@@ -217,70 +383,101 @@ class DeepHealth(nn.Module):
|
||||
self.time_mode = time_mode
|
||||
self.dist_mode = dist_mode
|
||||
self.extra_pool_reduce = extra_pool_reduce
|
||||
self.n_embd = n_embd
|
||||
self.model_size = normalized_model_size
|
||||
self.d_model = d_model
|
||||
self.n_trajectory = n_trajectory
|
||||
self.trajectory_dim = d_model // n_trajectory
|
||||
self.traj_hidden = 4 * n_trajectory
|
||||
self.n_reasoning_rounds = n_reasoning_rounds
|
||||
self.vocab_size = vocab_size
|
||||
self.register_buffer(
|
||||
"architecture_d_model",
|
||||
torch.tensor(d_model, dtype=torch.int64),
|
||||
)
|
||||
self.register_buffer(
|
||||
"architecture_n_trajectory",
|
||||
torch.tensor(n_trajectory, dtype=torch.int64),
|
||||
)
|
||||
self.register_buffer(
|
||||
"architecture_n_reasoning_rounds",
|
||||
torch.tensor(n_reasoning_rounds, dtype=torch.int64),
|
||||
)
|
||||
nn.init.normal_(self.token_embedding.weight, mean=0.0, std=0.02)
|
||||
nn.init.zeros_(self.token_embedding.weight[0])
|
||||
nn.init.normal_(self.gender_embedding.weight, mean=0.0, std=0.02)
|
||||
if dist_mode == "weibull":
|
||||
self.rho_head = nn.Linear(n_embd, vocab_size)
|
||||
self.rho_head = nn.Linear(d_model, vocab_size)
|
||||
nn.init.zeros_(self.rho_head.weight)
|
||||
nn.init.constant_(self.rho_head.bias, 0.5413)
|
||||
|
||||
if dist_mode == "mixed":
|
||||
self.death_idx = vocab_size - 1
|
||||
self.rho_death_head = nn.Linear(n_embd, 1)
|
||||
self.rho_death_head = nn.Linear(d_model, 1)
|
||||
nn.init.zeros_(self.rho_death_head.weight)
|
||||
nn.init.constant_(self.rho_death_head.bias, 0.5413)
|
||||
|
||||
if time_mode == "absolute":
|
||||
self.age_encoding = AgeSinusoidalEncoding(n_embd)
|
||||
self.blocks = nn.ModuleList([
|
||||
GPTBlock(
|
||||
n_embd=n_embd,
|
||||
n_head=n_head,
|
||||
use_time_rope=False,
|
||||
use_rbf_bias=False,
|
||||
mlp_dropout=dropout,
|
||||
) for _ in range(n_hist_layer)
|
||||
])
|
||||
self.rope = None
|
||||
self.rbf = None
|
||||
elif time_mode == "relative":
|
||||
self.age_encoding = None
|
||||
self.blocks = nn.ModuleList([
|
||||
GPTBlock(
|
||||
n_embd=n_embd,
|
||||
n_head=n_head,
|
||||
use_time_rope=True,
|
||||
use_rbf_bias=True,
|
||||
mlp_dropout=dropout,
|
||||
) for _ in range(n_hist_layer)
|
||||
])
|
||||
self.rope = TimeRoPE(n_embd // n_head)
|
||||
self.rbf = GaussianRBFTimeBasis(n_bases=16, max_time_diff=40.0)
|
||||
|
||||
self.final_ln = nn.LayerNorm(n_embd)
|
||||
self.risk_head = nn.Linear(n_embd, vocab_size, bias=False)
|
||||
if target_mode == "next_token":
|
||||
self.risk_head.weight = self.token_embedding.weight
|
||||
self.query_token = nn.Parameter(torch.zeros(n_embd))
|
||||
# Event and query time are encoded once before shared reasoning. In
|
||||
# relative mode, cross-attention additionally uses TimeRoPE and RBF.
|
||||
self.age_encoding = AgeSinusoidalEncoding(d_model)
|
||||
self.event_projection = nn.Linear(d_model, d_model, bias=False)
|
||||
self.event_norm = nn.LayerNorm(d_model)
|
||||
self.query_projection = nn.Linear(d_model, d_model, bias=False)
|
||||
self.trajectory_prototypes = nn.Parameter(
|
||||
torch.empty(n_trajectory, self.trajectory_dim)
|
||||
)
|
||||
self.query_token = nn.Parameter(torch.empty(d_model))
|
||||
nn.init.normal_(self.event_projection.weight, mean=0.0, std=0.02)
|
||||
nn.init.normal_(self.query_projection.weight, mean=0.0, std=0.02)
|
||||
nn.init.normal_(self.trajectory_prototypes, mean=0.0, std=0.02)
|
||||
nn.init.normal_(self.query_token, mean=0.0, std=0.02)
|
||||
|
||||
def _make_history_attn_mask(
|
||||
use_relative_time = time_mode == "relative"
|
||||
self.reasoning_core = SharedEventTrajectoryCore(
|
||||
d_model=d_model,
|
||||
n_trajectory=n_trajectory,
|
||||
n_reasoning_rounds=n_reasoning_rounds,
|
||||
dropout=dropout,
|
||||
n_rbf_bases=16,
|
||||
use_time_rope=use_relative_time,
|
||||
use_rbf_bias=use_relative_time,
|
||||
)
|
||||
if use_relative_time:
|
||||
self.rope = TimeRoPE(self.trajectory_dim)
|
||||
self.rbf = GaussianRBFTimeBasis(
|
||||
n_bases=16,
|
||||
max_time_diff=40.0,
|
||||
)
|
||||
else:
|
||||
self.rope = None
|
||||
self.rbf = None
|
||||
|
||||
self.final_ln = nn.LayerNorm(d_model)
|
||||
self.risk_head = nn.Linear(d_model, vocab_size, bias=False)
|
||||
if target_mode == "next_token":
|
||||
self.risk_head.weight = self.token_embedding.weight
|
||||
|
||||
def _make_event_invalid_mask(
|
||||
self,
|
||||
padding_mask: torch.Tensor,
|
||||
time_seq: torch.Tensor,
|
||||
dtype: torch.dtype,
|
||||
event_valid_mask: torch.Tensor,
|
||||
event_time: torch.Tensor,
|
||||
query_time: torch.Tensor,
|
||||
query_position: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
valid_key = padding_mask[:, None, :] # (B, 1, L)
|
||||
visible_by_time = time_seq[:, None, :] <= time_seq[:, :, None]
|
||||
valid = valid_key & visible_by_time
|
||||
return torch.zeros(
|
||||
valid.shape,
|
||||
device=valid.device,
|
||||
dtype=dtype,
|
||||
).masked_fill(~valid, -1e4)[:, None, :, :]
|
||||
valid_key = event_valid_mask[:, None, :]
|
||||
key_time = event_time[:, None, :]
|
||||
query_time = query_time[:, :, None]
|
||||
if query_position is None:
|
||||
visible_by_time = key_time <= query_time
|
||||
else:
|
||||
key_position = torch.arange(
|
||||
event_time.size(1),
|
||||
device=event_time.device,
|
||||
).view(1, 1, -1)
|
||||
visible_by_time = (key_time < query_time) | (
|
||||
(key_time == query_time)
|
||||
& (key_position <= query_position[:, :, None])
|
||||
)
|
||||
return ~(valid_key & visible_by_time)
|
||||
|
||||
def _pool_other_by_time(
|
||||
self,
|
||||
@@ -384,8 +581,8 @@ class DeepHealth(nn.Module):
|
||||
padding_mask = padding_mask.to(device=event_seq.device, dtype=torch.bool)
|
||||
|
||||
event_len = event_seq.size(1)
|
||||
h_disease = self.token_embedding(event_seq)
|
||||
t_disease = time_seq
|
||||
event_features = self.token_embedding(event_seq)
|
||||
event_time = time_seq
|
||||
|
||||
if other_time.shape != other_type.shape:
|
||||
raise ValueError(
|
||||
@@ -393,64 +590,120 @@ class DeepHealth(nn.Module):
|
||||
f"{tuple(other_time.shape)} vs {tuple(other_type.shape)}"
|
||||
)
|
||||
other_time = other_time.to(device=event_seq.device, dtype=time_seq.dtype)
|
||||
h_other, other_mask = self.tokenizer(
|
||||
other_features, other_mask = self.tokenizer(
|
||||
other_type=other_type,
|
||||
other_value=other_value,
|
||||
other_value_kind=other_value_kind,
|
||||
)
|
||||
h_other = h_other.to(device=event_seq.device)
|
||||
other_features = other_features.to(device=event_seq.device)
|
||||
other_mask = other_mask.to(device=event_seq.device, dtype=torch.bool)
|
||||
|
||||
h_disease = torch.cat([h_disease, h_other], dim=1)
|
||||
t_disease = torch.cat([t_disease, other_time], dim=1)
|
||||
padding_mask = torch.cat([padding_mask, other_mask], dim=1)
|
||||
h_disease = h_disease * padding_mask.unsqueeze(-1).to(h_disease.dtype)
|
||||
event_features = torch.cat([event_features, other_features], dim=1)
|
||||
event_time = torch.cat([event_time, other_time], dim=1)
|
||||
event_valid_mask = torch.cat([padding_mask, other_mask], dim=1)
|
||||
|
||||
batch_size = event_seq.size(0)
|
||||
sex_context = self.gender_embedding(sex)[:, None, :]
|
||||
event_features = (
|
||||
event_features
|
||||
+ sex_context
|
||||
+ self.age_encoding(event_time)
|
||||
)
|
||||
event_features = event_features * event_valid_mask.unsqueeze(-1).to(
|
||||
event_features.dtype
|
||||
)
|
||||
event_memory = self.event_norm(
|
||||
self.event_projection(event_features)
|
||||
)
|
||||
event_memory = event_memory * event_valid_mask.unsqueeze(-1).to(
|
||||
event_memory.dtype
|
||||
)
|
||||
|
||||
if mode == "all_future":
|
||||
batch_size = event_seq.size(0)
|
||||
query = self.query_token.view(1, 1, -1).expand(batch_size, 1, -1)
|
||||
h_disease = torch.cat([h_disease, query], dim=1)
|
||||
t_disease = torch.cat([t_disease, t_query[:, None]], dim=1)
|
||||
query_mask = torch.ones(
|
||||
query_time = t_query[:, None]
|
||||
query_position = None
|
||||
query_features = (
|
||||
self.query_token.view(1, 1, -1)
|
||||
+ sex_context
|
||||
+ self.age_encoding(query_time)
|
||||
)
|
||||
query_valid_mask = torch.ones(
|
||||
batch_size,
|
||||
1,
|
||||
dtype=torch.bool,
|
||||
device=event_seq.device,
|
||||
)
|
||||
padding_mask = torch.cat([padding_mask, query_mask], dim=1)
|
||||
else:
|
||||
# Each event position is an independent parallel query. Including
|
||||
# its event feature preserves token-level next-step semantics.
|
||||
# Equal-time memory is additionally position-causal so a token
|
||||
# cannot read a later token that may be its Delphi2M target.
|
||||
query_time = event_time
|
||||
query_position = torch.arange(
|
||||
event_time.size(1),
|
||||
device=event_time.device,
|
||||
).view(1, -1).expand(batch_size, -1)
|
||||
query_features = event_features
|
||||
query_valid_mask = event_valid_mask
|
||||
|
||||
sex_emb = self.gender_embedding(sex)[:, None, :]
|
||||
h_disease = h_disease + sex_emb
|
||||
h_disease = h_disease * padding_mask.unsqueeze(-1).to(h_disease.dtype)
|
||||
|
||||
rope_cache = None
|
||||
rbf_cache = None
|
||||
if self.time_mode == "absolute":
|
||||
h_disease = h_disease + self.age_encoding(t_disease)
|
||||
h_disease = h_disease * padding_mask.unsqueeze(-1).to(h_disease.dtype)
|
||||
elif self.time_mode == "relative":
|
||||
rope_cache = self.rope.precompute_cache(t_disease)
|
||||
rbf_cache = self.rbf.precompute_cache(t_disease)
|
||||
|
||||
attn_mask = self._make_history_attn_mask(
|
||||
padding_mask=padding_mask,
|
||||
time_seq=t_disease,
|
||||
dtype=h_disease.dtype,
|
||||
n_query = query_time.size(1)
|
||||
query_context = self.query_projection(query_features).reshape(
|
||||
batch_size,
|
||||
n_query,
|
||||
self.n_trajectory,
|
||||
self.trajectory_dim,
|
||||
)
|
||||
for block in self.blocks:
|
||||
h_disease = block(
|
||||
h_disease,
|
||||
rope_cache=rope_cache,
|
||||
rbf_cache=rbf_cache,
|
||||
attn_mask=attn_mask,
|
||||
trajectory_state = (
|
||||
self.trajectory_prototypes.view(
|
||||
1,
|
||||
1,
|
||||
self.n_trajectory,
|
||||
self.trajectory_dim,
|
||||
)
|
||||
h_disease = h_disease * padding_mask.unsqueeze(-1).to(h_disease.dtype)
|
||||
+ query_context
|
||||
)
|
||||
event_invalid_mask = self._make_event_invalid_mask(
|
||||
event_valid_mask=event_valid_mask,
|
||||
event_time=event_time,
|
||||
query_time=query_time,
|
||||
query_position=query_position,
|
||||
)
|
||||
|
||||
h_disease = self.final_ln(h_disease)
|
||||
h_disease = h_disease * padding_mask.unsqueeze(-1).to(h_disease.dtype)
|
||||
event_rope_cache = None
|
||||
query_rope_cache = None
|
||||
rbf_cache = None
|
||||
if self.time_mode == "relative":
|
||||
if self.rope is None or self.rbf is None:
|
||||
raise RuntimeError("Relative-time modules are not initialized")
|
||||
event_rope_cache = self.rope.precompute_cache(event_time)
|
||||
query_rope_cache = self.rope.precompute_cache(query_time)
|
||||
rbf_cache = self.rbf.precompute_cross_cache(
|
||||
query_time,
|
||||
event_time,
|
||||
)
|
||||
|
||||
event_key_value = self.reasoning_core.project_event_memory(
|
||||
event_memory,
|
||||
event_rope_cache=event_rope_cache,
|
||||
)
|
||||
for _ in range(self.n_reasoning_rounds):
|
||||
trajectory_state = self.reasoning_core(
|
||||
trajectory_state=trajectory_state,
|
||||
event_key_value=event_key_value,
|
||||
event_invalid_mask=event_invalid_mask,
|
||||
query_rope_cache=query_rope_cache,
|
||||
rbf_cache=rbf_cache,
|
||||
)
|
||||
|
||||
hidden_sequence = self.final_ln(
|
||||
trajectory_state.reshape(batch_size, n_query, self.d_model)
|
||||
)
|
||||
hidden_sequence = hidden_sequence * query_valid_mask.unsqueeze(-1).to(
|
||||
hidden_sequence.dtype
|
||||
)
|
||||
|
||||
if mode == "all_future":
|
||||
hidden = h_disease[:, -1, :]
|
||||
hidden = hidden_sequence[:, 0, :]
|
||||
if return_output:
|
||||
return DeepHealthOutput(
|
||||
hidden=hidden,
|
||||
@@ -465,13 +718,13 @@ class DeepHealth(nn.Module):
|
||||
)
|
||||
return hidden
|
||||
if return_output:
|
||||
h_event = h_disease[:, :event_len, :]
|
||||
t_event = t_disease[:, :event_len]
|
||||
event_mask = padding_mask[:, :event_len]
|
||||
h_event = hidden_sequence[:, :event_len, :]
|
||||
t_event = event_time[:, :event_len]
|
||||
event_mask = event_valid_mask[:, :event_len]
|
||||
h_extra, t_extra, extra_mask = self._pool_other_by_time(
|
||||
h_other=h_disease[:, event_len:, :],
|
||||
other_time=t_disease[:, event_len:],
|
||||
other_mask=padding_mask[:, event_len:],
|
||||
h_other=hidden_sequence[:, event_len:, :],
|
||||
other_time=event_time[:, event_len:],
|
||||
other_mask=event_valid_mask[:, event_len:],
|
||||
)
|
||||
return DeepHealthOutput(
|
||||
hidden=torch.cat([h_event, h_extra], dim=1),
|
||||
@@ -479,7 +732,7 @@ class DeepHealth(nn.Module):
|
||||
padding_mask=torch.cat([event_mask, extra_mask], dim=1),
|
||||
event_len=event_len,
|
||||
)
|
||||
return h_disease[:, :event_len, :]
|
||||
return hidden_sequence[:, :event_len, :]
|
||||
|
||||
def forward_next_token(self, **kwargs) -> torch.Tensor:
|
||||
return self._forward_shared(mode="next_token", **kwargs)
|
||||
|
||||
Reference in New Issue
Block a user