Revert "Implement shared event-trajectory reasoning backbone"

This reverts commit 06f29c0f0a.
This commit is contained in:
2026-07-23 16:05:00 +08:00
parent 22faee7c51
commit 85352dae0f
11 changed files with 753 additions and 1392 deletions

View File

@@ -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):