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) class SharedTrajectoryMixer(nn.Module): """SwiGLU interaction along the trajectory axis only.""" 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 self.gate_proj = nn.Parameter( torch.empty(trajectory_dim, n_trajectory, self.traj_hidden) ) self.value_proj = nn.Parameter( torch.empty(trajectory_dim, n_trajectory, self.traj_hidden) ) self.output_proj = nn.Parameter( torch.empty(trajectory_dim, self.traj_hidden, n_trajectory) ) self.reset_parameters() def reset_parameters(self) -> None: for feature_idx in range(self.trajectory_dim): 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): raise ValueError( "Expected trailing trajectory shape " f"{(self.n_trajectory, self.trajectory_dim)}, got " f"{tuple(state.shape[-2:])}" ) gate = torch.einsum("...hr,rhk->...kr", state, self.gate_proj) value = torch.einsum("...hr,rhk->...kr", state, self.value_proj) 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, ) 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