451 lines
17 KiB
Python
451 lines
17 KiB
Python
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
|