Revert "Implement shared event-trajectory reasoning backbone"
This reverts commit 06f29c0f0a.
This commit is contained in:
390
backbones.py
390
backbones.py
@@ -29,14 +29,6 @@ class TimeRoPE(nn.Module):
|
||||
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,
|
||||
@@ -93,280 +85,232 @@ class GaussianRBFTimeBasis(nn.Module):
|
||||
)
|
||||
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."""
|
||||
|
||||
class TemporalAttention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
d_model: int,
|
||||
n_trajectory: int,
|
||||
n_embd: int,
|
||||
n_head: int,
|
||||
n_rbf_bases: int = 16,
|
||||
use_time_rope: bool = False,
|
||||
use_rbf_bias: bool = False,
|
||||
dropout: float = 0.0,
|
||||
use_time_rope: bool = True,
|
||||
use_rbf_bias: bool = True,
|
||||
):
|
||||
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
|
||||
assert n_embd % n_head == 0, "n_embd must be divisible by n_head"
|
||||
self.n_head = n_head
|
||||
self.d_head = n_embd // n_head
|
||||
self.scale = 1.0 / math.sqrt(self.d_head)
|
||||
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)
|
||||
# QKV projection (fused for efficiency)
|
||||
self.qkv = nn.Linear(n_embd, 3 * n_embd, bias=False)
|
||||
# Output projection
|
||||
self.out_proj = nn.Linear(n_embd, n_embd, bias=False)
|
||||
|
||||
# Layer-specific projection from shared RBF basis activations to per-head attention bias.
|
||||
self.rbf_proj = nn.Linear(n_rbf_bases, n_head, bias=False)
|
||||
self.time_bias_scale = nn.Parameter(torch.tensor(0.0))
|
||||
|
||||
self.resid_drop = nn.Dropout(dropout)
|
||||
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
|
||||
"""Match the previous version's GPT-style weight initialization."""
|
||||
nn.init.normal_(self.qkv.weight, mean=0.0, std=0.02)
|
||||
nn.init.normal_(self.out_proj.weight, mean=0.0, std=0.02)
|
||||
nn.init.zeros_(self.rbf_proj.weight)
|
||||
|
||||
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,
|
||||
x: torch.Tensor,
|
||||
rope_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
rbf_cache: torch.Tensor | None = None,
|
||||
attn_mask: 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
|
||||
assert rope_cache is not None, "rope_cache must be provided when use_time_rope is True"
|
||||
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
|
||||
assert rbf_cache is not None, "rbf_cache must be provided when use_rbf_bias is True"
|
||||
|
||||
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
|
||||
B, L, _ = x.shape
|
||||
H, D = self.n_head, self.d_head
|
||||
|
||||
# --- QKV ----------------------------------------------------------
|
||||
qkv = self.qkv(x).reshape(B, L, 3, H, D).permute(2, 0, 3, 1, 4)
|
||||
q, k, v = qkv.unbind(0) # each (B, H, L, D)
|
||||
|
||||
# --- Apply RoPE (from shared cache) --------------------------------
|
||||
if self.use_time_rope:
|
||||
q, k = TimeRoPE.apply_from_cache(q, k, rope_cache)
|
||||
|
||||
# Build additive attention bias mask: time bias + causal/padding mask.
|
||||
time_bias = None
|
||||
if self.use_rbf_bias:
|
||||
time_bias = self.rbf_proj(rbf_cache).permute(
|
||||
0, 3, 1, 2) # (B, H, L, L)
|
||||
time_bias = self.time_bias_scale.tanh() * time_bias
|
||||
|
||||
if time_bias is not None and attn_mask is not None:
|
||||
attn_bias = time_bias + attn_mask.to(time_bias.dtype)
|
||||
elif time_bias is not None:
|
||||
attn_bias = time_bias
|
||||
elif attn_mask is not None:
|
||||
attn_bias = attn_mask
|
||||
else:
|
||||
attn_bias = None
|
||||
|
||||
out = F.scaled_dot_product_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
attn_mask=attn_bias,
|
||||
dropout_p=0.0,
|
||||
is_causal=False,
|
||||
scale=self.scale,
|
||||
)
|
||||
return torch.einsum("bqhl,bhld->bqhd", weights, value)
|
||||
|
||||
# --- Aggregate & project out --------------------------------------
|
||||
out = out.transpose(1, 2).reshape(B, L, H * D)
|
||||
return self.resid_drop(self.out_proj(out))
|
||||
|
||||
|
||||
class SharedTrajectoryMixer(nn.Module):
|
||||
"""SwiGLU interaction along the trajectory axis only."""
|
||||
class TrajMixer(nn.Module):
|
||||
"""Lightweight gated interaction across latent residual-space groups.
|
||||
|
||||
def __init__(self, n_trajectory: int, trajectory_dim: int):
|
||||
The groups are contiguous partitions of the post-``W_O`` residual
|
||||
representation. They are deliberately not treated as attention heads.
|
||||
All operations are position-wise, so the sequence dimension remains fully
|
||||
parallel and no temporal information can leak between positions here.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
n_embd: int,
|
||||
n_head: int = 10,
|
||||
dropout: float = 0.0,
|
||||
):
|
||||
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
|
||||
if n_embd <= 0:
|
||||
raise ValueError(f"n_embd must be > 0, got {n_embd}")
|
||||
if n_head <= 0:
|
||||
raise ValueError(f"n_head must be > 0, got {n_head}")
|
||||
if n_embd % n_head != 0:
|
||||
raise ValueError(
|
||||
f"n_embd must be divisible by n_head, got {n_embd} and {n_head}"
|
||||
)
|
||||
self.n_embd = n_embd
|
||||
# The residual-group count is tied to n_head, but the resulting groups
|
||||
# are still residual-space partitions rather than attention heads.
|
||||
self.n_group = n_head
|
||||
self.d_group = n_embd // n_head
|
||||
self.hidden_group = 4 * n_head
|
||||
|
||||
# Per-group feature alignment: [group, input feature, output feature].
|
||||
self.group_align = nn.Parameter(
|
||||
torch.empty(self.n_group, self.d_group, self.d_group)
|
||||
)
|
||||
|
||||
# Per-feature cross-group projections. The feature index is kept
|
||||
# independent, exactly as specified by the TrajMixer baseline.
|
||||
self.gate_proj = nn.Parameter(
|
||||
torch.empty(trajectory_dim, n_trajectory, self.traj_hidden)
|
||||
torch.empty(self.d_group, self.n_group, self.hidden_group)
|
||||
)
|
||||
self.value_proj = nn.Parameter(
|
||||
torch.empty(trajectory_dim, n_trajectory, self.traj_hidden)
|
||||
torch.empty(self.d_group, self.n_group, self.hidden_group)
|
||||
)
|
||||
self.output_proj = nn.Parameter(
|
||||
torch.empty(trajectory_dim, self.traj_hidden, n_trajectory)
|
||||
torch.empty(self.d_group, self.hidden_group, self.n_group)
|
||||
)
|
||||
self.drop = nn.Dropout(dropout)
|
||||
self.reset_parameters()
|
||||
|
||||
def reset_parameters(self) -> None:
|
||||
for feature_idx in range(self.trajectory_dim):
|
||||
with torch.no_grad():
|
||||
identity = torch.eye(
|
||||
self.d_group,
|
||||
dtype=self.group_align.dtype,
|
||||
device=self.group_align.device,
|
||||
)
|
||||
self.group_align.copy_(identity.unsqueeze(0).expand_as(self.group_align))
|
||||
|
||||
# Initialise each feature-specific matrix independently so Xavier's
|
||||
# fan-in/fan-out calculation sees a two-dimensional matrix.
|
||||
for feature_idx in range(self.d_group):
|
||||
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):
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Map ``(B, L, n_embd)`` to an equally shaped residual update."""
|
||||
if x.ndim != 3:
|
||||
raise ValueError(f"TrajMixer expects a 3D tensor, got shape {tuple(x.shape)}")
|
||||
if x.size(-1) != self.n_embd:
|
||||
raise ValueError(
|
||||
"Expected trailing trajectory shape "
|
||||
f"{(self.n_trajectory, self.trajectory_dim)}, got "
|
||||
f"{tuple(state.shape[-2:])}"
|
||||
f"Expected hidden size {self.n_embd}, got {x.size(-1)}"
|
||||
)
|
||||
gate = torch.einsum("...hr,rhk->...kr", state, self.gate_proj)
|
||||
value = torch.einsum("...hr,rhk->...kr", state, self.value_proj)
|
||||
|
||||
batch_size, seq_len, _ = x.shape
|
||||
grouped = x.reshape(
|
||||
batch_size, seq_len, self.n_group, self.d_group
|
||||
)
|
||||
aligned = torch.einsum(
|
||||
"blgd,gde->blge", grouped, self.group_align
|
||||
)
|
||||
|
||||
gate = torch.einsum(
|
||||
"blgr,rgh->blhr", aligned, self.gate_proj
|
||||
)
|
||||
value = torch.einsum(
|
||||
"blgr,rgh->blhr", aligned, self.value_proj
|
||||
)
|
||||
hidden = F.silu(gate) * value
|
||||
return torch.einsum("...kr,rkh->...hr", hidden, self.output_proj)
|
||||
mixed = torch.einsum(
|
||||
"blhr,rhg->blgr", hidden, self.output_proj
|
||||
)
|
||||
return self.drop(mixed.reshape(batch_size, seq_len, self.n_embd))
|
||||
|
||||
|
||||
class SharedEventTrajectoryCore(nn.Module):
|
||||
"""One parameter-shared reasoning core reused across all rounds."""
|
||||
|
||||
class GPTBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
d_model: int,
|
||||
n_trajectory: int,
|
||||
n_reasoning_rounds: int,
|
||||
dropout: float = 0.0,
|
||||
n_rbf_bases: int = 16,
|
||||
n_embd: int,
|
||||
n_head: int,
|
||||
|
||||
attn_dropout: float = 0.0,
|
||||
mlp_dropout: float = 0.0,
|
||||
use_time_rope: bool = False,
|
||||
use_rbf_bias: bool = False,
|
||||
n_rbf_bases: int = 16,
|
||||
):
|
||||
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,
|
||||
self.attn = TemporalAttention(
|
||||
n_embd=n_embd,
|
||||
n_head=n_head,
|
||||
n_rbf_bases=n_rbf_bases,
|
||||
dropout=attn_dropout,
|
||||
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,
|
||||
self.mlp = TrajMixer(
|
||||
n_embd=n_embd,
|
||||
n_head=n_head,
|
||||
dropout=mlp_dropout,
|
||||
)
|
||||
self.ln1 = nn.LayerNorm(n_embd)
|
||||
self.ln2 = nn.LayerNorm(n_embd)
|
||||
|
||||
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,
|
||||
x: torch.Tensor,
|
||||
rope_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
rbf_cache: torch.Tensor | None = None,
|
||||
attn_mask: 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)
|
||||
x = x + self.attn(self.ln1(x), rope_cache, rbf_cache, attn_mask)
|
||||
x = x + self.mlp(self.ln2(x))
|
||||
return x
|
||||
|
||||
|
||||
class TokenAutoDiscretization(nn.Module):
|
||||
|
||||
Reference in New Issue
Block a user