Files
DeepHealth/backbones.py

451 lines
17 KiB
Python
Raw Normal View History

import math
import torch
import torch.nn as nn
import torch.nn.functional as F
class TimeRoPE(nn.Module):
def __init__(self, dim: int, base: float = 10000.0):
super().__init__()
assert dim % 2 == 0, "RoPE dim must be even"
self.dim = dim
# inv_freq is not trainable, but should move with device.
inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))
self.register_buffer("inv_freq", inv_freq, persistent=False)
def precompute_cache(self, tau: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
t = tau.unsqueeze(-1) # (B, L, 1)
angles = t * self.inv_freq # (B, L, dim//2)
# Pre-expand for heads and interleave once (avoids N_layers repeats)
cos = angles.cos().unsqueeze(1).repeat_interleave(2, dim=-1)
sin = angles.sin().unsqueeze(1).repeat_interleave(2, dim=-1)
return cos, sin # (B, 1, L, dim)
@staticmethod
def _rotate_half(x: torch.Tensor) -> torch.Tensor:
"""Rotate pairs: ``[-x2, x1, -x4, x3, ...]``."""
x1 = x[..., 0::2]
x2 = x[..., 1::2]
return torch.stack((-x2, x1), dim=-1).flatten(-2)
@staticmethod
def apply_single_from_cache(
x: torch.Tensor,
rope_cache: tuple[torch.Tensor, torch.Tensor],
) -> torch.Tensor:
cos, sin = rope_cache
return x * cos + TimeRoPE._rotate_half(x) * sin
@staticmethod
def apply_from_cache(
q: torch.Tensor,
k: torch.Tensor,
rope_cache: tuple[torch.Tensor, torch.Tensor],
) -> tuple[torch.Tensor, torch.Tensor]:
cos, sin = rope_cache # each (B, 1, L, dim)
q_rot = q * cos + TimeRoPE._rotate_half(q) * sin
k_rot = k * cos + TimeRoPE._rotate_half(k) * sin
return q_rot, k_rot
def forward(
self,
tau: torch.Tensor,
q: torch.Tensor,
k: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
cache = self.precompute_cache(tau)
return self.apply_from_cache(q, k, cache)
class GaussianRBFTimeBasis(nn.Module):
def __init__(
self,
n_bases: int = 16,
max_time_diff: float = 40.0,
):
super().__init__()
self.n_bases = n_bases
# Evenly spaced RBF centres for non-negative linear time differences.
# Causal masking enforces query_time >= key_time, so diff is >= 0.
centers = torch.linspace(0.0, max_time_diff, n_bases)
self.register_buffer("centers", centers,
persistent=False) # (n_bases,)
# Learnable log-widths (initialized to center spacing on linear scale).
init_width = max(max_time_diff / max(n_bases - 1, 1), 1e-3)
init_log_width = math.log(init_width)
self.log_widths = nn.Parameter(torch.full((n_bases,), init_log_width))
def precompute_cache(self, tau: torch.Tensor) -> torch.Tensor:
time_coord = tau.float() # (B, L)
# Pairwise signed difference: query_i - key_j.
diff = time_coord.unsqueeze(
2) - time_coord.unsqueeze(1) # (B, L_q, L_k)
# Gaussian RBF: exp(-0.5 * ((diff - c) / w)^2)
diff = diff.unsqueeze(-1) # (B, L, L, 1)
widths = self.log_widths.exp() # (n_bases,)
rbf_acts = torch.exp(
-0.5 * ((diff - self.centers) / widths).square()
# (B, L, L, n_bases)
)
return rbf_acts
def precompute_cross_cache(
self,
query_tau: torch.Tensor,
key_tau: torch.Tensor,
) -> torch.Tensor:
"""Return RBF activations for query-time minus event-time."""
diff = query_tau.float().unsqueeze(2) - key_tau.float().unsqueeze(1)
widths = self.log_widths.exp()
return torch.exp(
-0.5
* (
(diff.unsqueeze(-1) - self.centers)
/ widths
).square()
)
class TrajectoryCrossAttention(nn.Module):
"""Shared trajectory-slot queries reading a fixed event memory."""
def __init__(
self,
d_model: int,
n_trajectory: int,
n_rbf_bases: int = 16,
use_time_rope: bool = False,
use_rbf_bias: bool = False,
):
super().__init__()
if d_model <= 0 or n_trajectory <= 0:
raise ValueError("d_model and n_trajectory must be positive")
if d_model % n_trajectory != 0:
raise ValueError(
"d_model must be divisible by n_trajectory, got "
f"{d_model} and {n_trajectory}"
)
self.d_model = d_model
self.n_trajectory = n_trajectory
self.trajectory_dim = d_model // n_trajectory
self.scale = self.trajectory_dim ** -0.5
self.use_time_rope = use_time_rope
self.use_rbf_bias = use_rbf_bias
# q_proj acts on each slot independently and is shared across slots.
self.q_proj = nn.Linear(
self.trajectory_dim,
self.trajectory_dim,
bias=False,
)
self.k_proj = nn.Linear(d_model, d_model, bias=False)
self.v_proj = nn.Linear(d_model, d_model, bias=False)
if use_rbf_bias:
self.rbf_proj = nn.Linear(
n_rbf_bases,
n_trajectory,
bias=False,
)
self.time_bias_scale = nn.Parameter(torch.tensor(0.0))
else:
self.rbf_proj = None
self.register_parameter("time_bias_scale", None)
self.reset_parameters()
def reset_parameters(self) -> None:
nn.init.normal_(self.q_proj.weight, mean=0.0, std=0.02)
nn.init.normal_(self.k_proj.weight, mean=0.0, std=0.02)
nn.init.normal_(self.v_proj.weight, mean=0.0, std=0.02)
if self.rbf_proj is not None:
# The scalar gate starts at zero, so the relative-time bias still
# starts disabled. A nonzero projection is necessary for the gate
# itself to receive a gradient on the first optimization step.
nn.init.xavier_uniform_(self.rbf_proj.weight)
def project_event_memory(
self,
event_memory: torch.Tensor,
event_rope_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Project K/V once for reuse by every reasoning round."""
batch_size, memory_len, _ = event_memory.shape
key = self.k_proj(event_memory).reshape(
batch_size,
memory_len,
self.n_trajectory,
self.trajectory_dim,
).transpose(1, 2)
value = self.v_proj(event_memory).reshape(
batch_size,
memory_len,
self.n_trajectory,
self.trajectory_dim,
).transpose(1, 2)
if self.use_time_rope:
if event_rope_cache is None:
raise ValueError(
"event_rope_cache is required when TimeRoPE is enabled"
)
key = TimeRoPE.apply_single_from_cache(key, event_rope_cache)
return key, value
def forward(
self,
trajectory_state: torch.Tensor,
event_key_value: tuple[torch.Tensor, torch.Tensor],
event_invalid_mask: torch.Tensor,
query_rope_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
rbf_cache: torch.Tensor | None = None,
) -> torch.Tensor:
"""Read memory for states shaped ``(B, Q, H, Dh)``."""
if trajectory_state.ndim != 4:
raise ValueError(
"trajectory_state must have shape (B, Q, H, Dh), got "
f"{tuple(trajectory_state.shape)}"
)
batch_size, n_query, n_trajectory, trajectory_dim = (
trajectory_state.shape
)
if (n_trajectory, trajectory_dim) != (
self.n_trajectory,
self.trajectory_dim,
):
raise ValueError(
"Unexpected trajectory shape: "
f"{(n_trajectory, trajectory_dim)}"
)
key, value = event_key_value
memory_len = key.size(2)
if event_invalid_mask.shape != (batch_size, n_query, memory_len):
raise ValueError(
"event_invalid_mask must have shape "
f"{(batch_size, n_query, memory_len)}, got "
f"{tuple(event_invalid_mask.shape)}"
)
query = self.q_proj(trajectory_state).transpose(1, 2)
if self.use_time_rope:
if query_rope_cache is None:
raise ValueError(
"query_rope_cache is required when TimeRoPE is enabled"
)
query = TimeRoPE.apply_single_from_cache(query, query_rope_cache)
query = query.transpose(1, 2)
scores = torch.einsum("bqhd,bhld->bqhl", query, key) * self.scale
if self.use_rbf_bias:
if rbf_cache is None or self.rbf_proj is None:
raise ValueError(
"rbf_cache is required when relative time bias is enabled"
)
time_bias = self.rbf_proj(rbf_cache).permute(0, 1, 3, 2)
scores = scores + self.time_bias_scale.tanh() * time_bias
mask = event_invalid_mask.unsqueeze(2)
min_value = torch.finfo(scores.dtype).min
masked_scores = scores.masked_fill(mask, min_value)
weights = torch.softmax(masked_scores.float(), dim=-1).to(scores.dtype)
weights = weights.masked_fill(mask, 0.0)
denominator = weights.sum(dim=-1, keepdim=True)
weights = weights / denominator.clamp_min(
torch.finfo(weights.dtype).eps
)
return torch.einsum("bqhl,bhld->bqhd", weights, value)
2026-07-22 11:52:44 +08:00
class SharedTrajectoryMixer(nn.Module):
"""SwiGLU interaction along the trajectory axis only."""
2026-07-22 11:52:44 +08:00
def __init__(self, n_trajectory: int, trajectory_dim: int):
super().__init__()
if n_trajectory <= 0 or trajectory_dim <= 0:
raise ValueError("trajectory dimensions must be positive")
self.n_trajectory = n_trajectory
self.trajectory_dim = trajectory_dim
self.traj_hidden = 4 * n_trajectory
2026-07-22 11:52:44 +08:00
self.gate_proj = nn.Parameter(
torch.empty(trajectory_dim, n_trajectory, self.traj_hidden)
2026-07-22 11:52:44 +08:00
)
self.value_proj = nn.Parameter(
torch.empty(trajectory_dim, n_trajectory, self.traj_hidden)
2026-07-22 11:52:44 +08:00
)
self.output_proj = nn.Parameter(
torch.empty(trajectory_dim, self.traj_hidden, n_trajectory)
2026-07-22 11:52:44 +08:00
)
self.reset_parameters()
def reset_parameters(self) -> None:
for feature_idx in range(self.trajectory_dim):
2026-07-22 11:52:44 +08:00
nn.init.xavier_uniform_(self.gate_proj[feature_idx])
nn.init.xavier_uniform_(self.value_proj[feature_idx])
nn.init.normal_(self.output_proj, mean=0.0, std=1e-3)
def forward(self, state: torch.Tensor) -> torch.Tensor:
if state.shape[-2:] != (self.n_trajectory, self.trajectory_dim):
2026-07-22 11:52:44 +08:00
raise ValueError(
"Expected trailing trajectory shape "
f"{(self.n_trajectory, self.trajectory_dim)}, got "
f"{tuple(state.shape[-2:])}"
2026-07-22 11:52:44 +08:00
)
gate = torch.einsum("...hr,rhk->...kr", state, self.gate_proj)
value = torch.einsum("...hr,rhk->...kr", state, self.value_proj)
2026-07-22 11:52:44 +08:00
hidden = F.silu(gate) * value
return torch.einsum("...kr,rkh->...hr", hidden, self.output_proj)
class SharedEventTrajectoryCore(nn.Module):
"""One parameter-shared reasoning core reused across all rounds."""
def __init__(
self,
d_model: int,
n_trajectory: int,
n_reasoning_rounds: int,
dropout: float = 0.0,
n_rbf_bases: int = 16,
use_time_rope: bool = False,
use_rbf_bias: bool = False,
):
super().__init__()
if n_reasoning_rounds <= 0:
raise ValueError("n_reasoning_rounds must be positive")
if d_model <= 0 or n_trajectory <= 0:
raise ValueError("d_model and n_trajectory must be positive")
if d_model % n_trajectory != 0:
raise ValueError("d_model must equal n_trajectory * trajectory_dim")
trajectory_dim = d_model // n_trajectory
self.norm_attn = nn.LayerNorm(trajectory_dim)
self.cross_attention = TrajectoryCrossAttention(
d_model=d_model,
n_trajectory=n_trajectory,
n_rbf_bases=n_rbf_bases,
use_time_rope=use_time_rope,
use_rbf_bias=use_rbf_bias,
)
self.norm_mixer = nn.LayerNorm(trajectory_dim)
self.traj_mixer = SharedTrajectoryMixer(
n_trajectory=n_trajectory,
trajectory_dim=trajectory_dim,
)
initial_scale = 1.0 / math.sqrt(n_reasoning_rounds)
self.attn_scale = nn.Parameter(torch.tensor(initial_scale))
self.mixer_scale = nn.Parameter(torch.tensor(initial_scale))
self.dropout = nn.Dropout(dropout)
def project_event_memory(
self,
event_memory: torch.Tensor,
event_rope_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
return self.cross_attention.project_event_memory(
event_memory,
event_rope_cache=event_rope_cache,
2026-07-22 11:52:44 +08:00
)
def forward(
self,
trajectory_state: torch.Tensor,
event_key_value: tuple[torch.Tensor, torch.Tensor],
event_invalid_mask: torch.Tensor,
query_rope_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
rbf_cache: torch.Tensor | None = None,
) -> torch.Tensor:
readout = self.cross_attention(
trajectory_state=self.norm_attn(trajectory_state),
event_key_value=event_key_value,
event_invalid_mask=event_invalid_mask,
query_rope_cache=query_rope_cache,
rbf_cache=rbf_cache,
)
updated = (
trajectory_state
+ self.attn_scale * self.dropout(readout)
)
mixed = self.traj_mixer(self.norm_mixer(updated))
return updated + self.mixer_scale * self.dropout(mixed)
class TokenAutoDiscretization(nn.Module):
def __init__(
self,
n_cont_types: int,
n_bins: int,
n_embd: int,
):
super().__init__()
if n_cont_types <= 0:
raise ValueError(f"n_cont_types must be > 0, got {n_cont_types}")
if n_bins <= 1:
raise ValueError(f"n_bins must be > 1, got {n_bins}")
if n_embd <= 0:
raise ValueError(f"n_embd must be > 0, got {n_embd}")
self.n_cont_types = n_cont_types
self.n_bins = n_bins
self.n_embd = n_embd
self.weight = nn.Parameter(torch.empty(n_cont_types, n_bins))
self.bias = nn.Parameter(torch.empty(n_cont_types, n_bins))
self.bin_emb = nn.Parameter(torch.empty(n_cont_types, n_bins, n_embd))
self.reset_parameters()
def reset_parameters(self) -> None:
nn.init.normal_(self.weight, mean=0.0, std=0.02)
nn.init.zeros_(self.bias)
nn.init.normal_(self.bin_emb, mean=0.0, std=0.02)
def forward(
self,
cont_type_idx: torch.LongTensor, # (N,)
value: torch.Tensor, # (N,)
) -> torch.Tensor:
if cont_type_idx.dim() != 1:
raise ValueError(
f"cont_type_idx must be 1D, got {tuple(cont_type_idx.shape)}"
)
if value.dim() != 1:
raise ValueError(f"value must be 1D, got {tuple(value.shape)}")
if cont_type_idx.numel() != value.numel():
raise ValueError("cont_type_idx and value must have the same length")
w = self.weight[cont_type_idx] # (N, n_bins)
b = self.bias[cont_type_idx] # (N, n_bins)
e = self.bin_emb[cont_type_idx] # (N, n_bins, D)
logits = value.unsqueeze(-1) * w + b
probs = torch.softmax(logits, dim=-1)
return torch.einsum("nb,nbd->nd", probs, e)
class AgeSinusoidalEncoding(nn.Module):
def __init__(self, embedding_dim: int):
super().__init__()
if embedding_dim % 2 != 0:
raise ValueError(
f"Embedding dimension must be an even number, but got {embedding_dim}")
self.embedding_dim = embedding_dim
i = torch.arange(0, self.embedding_dim, 2, dtype=torch.float32)
divisor = torch.pow(10000, i / self.embedding_dim)
self.register_buffer('divisor', divisor)
self.linear = nn.Linear(embedding_dim, embedding_dim, bias=False)
def forward(self, t: torch.Tensor) -> torch.Tensor:
t_years = t
# Broadcast (B, L, 1) against (1, 1, D/2) to get (B, L, D/2)
args = t_years.unsqueeze(-1) / self.divisor.view(1, 1, -1)
# Interleave cos and sin along the last dimension
output = torch.zeros(t.shape[0], t.shape[1],
self.embedding_dim, device=t.device)
output[:, :, 0::2] = torch.cos(args)
output[:, :, 1::2] = torch.sin(args)
output = self.linear(output)
return output