Modilify-Mk2-preview / latent_deliberation.py
ydy9038074's picture
Publish Modilify Mk2 Preview step 1250 schema25
b88f761 verified
Raw History Blame Contribute Delete
26.2 kB
"""Schema25 Torch GDN2 trajectory state, processor, and memory readers."""
from __future__ import annotations
from collections.abc import Sequence
from dataclasses import dataclass, replace
import weakref
import math
import torch
from torch import nn
from torch.nn import functional as F
from .gdn2_trajectory import GDN2TrajectoryMemory, GDN2TrajectoryState
COMMIT_REASON_NONE = 0
COMMIT_REASON_NORMAL = 1
COMMIT_REASON_FORCED_JUMP = 2
COMMIT_REASON_TERMINAL = 3
def _sdpa_mask_value(dtype: torch.dtype) -> float:
"""Additive SDPA mask that stays finite on MPS fp16/bf16."""
if dtype in (torch.float16, torch.bfloat16):
return -1.0e4
return -1.0e9
def _fp32_scaled_dot_product_attention(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
*,
attn_mask: torch.Tensor | None = None,
) -> torch.Tensor:
"""Run latent-memory attention reductions in FP32, then restore dtype.
These attention maps are small compared with the frozen decoder, while
their outputs feed recurrent trajectory and persistent-memory paths. A
BF16 reduction error therefore compounds across denoise/commit steps and
is much more expensive than the modest FP32 workspace.
"""
output_dtype = query.dtype
stable_mask = attn_mask
if stable_mask is not None and stable_mask.is_floating_point():
stable_mask = stable_mask.float()
query_fp32 = query.float()
key_fp32 = key.float()
use_batched_mm = query.ndim == 4 and key.ndim == 4 and value.ndim == 4
if use_batched_mm:
query_length = int(query.shape[-2])
key_length = int(key.shape[-2])
scores = torch.bmm(
query_fp32.reshape(-1, query_length, query.shape[-1]),
key_fp32.reshape(-1, key_length, key.shape[-1]).transpose(1, 2),
).reshape(*query.shape[:-2], query_length, key_length)
else:
scores = torch.matmul(query_fp32, key_fp32.transpose(-2, -1))
scores = scores / math.sqrt(max(query.shape[-1], 1))
if stable_mask is not None:
if stable_mask.dtype == torch.bool:
scores = scores.masked_fill(~stable_mask, _sdpa_mask_value(torch.float32))
else:
scores = scores + stable_mask
probabilities = torch.softmax(scores, dim=-1)
# All-masked rows are uniform under a finite mask, but keep a NaN
# barrier for any remaining -inf path that MPS softmax cannot invert.
probabilities = torch.nan_to_num(probabilities, nan=0.0)
if use_batched_mm:
output = torch.bmm(
probabilities.reshape(-1, query.shape[-2], key.shape[-2]),
value.float().reshape(-1, key.shape[-2], value.shape[-1]),
).reshape(*query.shape[:-2], query.shape[-2], value.shape[-1])
else:
output = torch.matmul(probabilities, value.float())
return output.to(dtype=output_dtype)
@dataclass
class LatentDeliberationState:
"""Persistent slots plus per-canvas trajectory clocks. No token latents."""
memory_slots: torch.Tensor
confidence: torch.Tensor
entropy: torch.Tensor
ponder_steps: torch.Tensor
stagnation_steps: torch.Tensor
gdn2: GDN2TrajectoryState
@classmethod
def empty(
cls,
*,
batch_size: int,
canvas_length: int,
device: torch.device,
) -> "LatentDeliberationState":
persistent = torch.zeros(batch_size, 16, 128, 128, device=device, dtype=torch.float32)
return cls(
memory_slots=persistent,
confidence=torch.zeros(
batch_size, canvas_length, device=device, dtype=torch.float32
),
entropy=torch.zeros(
batch_size, canvas_length, device=device, dtype=torch.float32
),
ponder_steps=torch.zeros(batch_size, device=device, dtype=torch.int32),
stagnation_steps=torch.zeros(batch_size, device=device, dtype=torch.int32),
gdn2=GDN2TrajectoryState(
cells=torch.zeros(batch_size, canvas_length, 16, 64, 64,
device=device, dtype=torch.float32),
row=torch.zeros(batch_size, 16, 64, 64,
device=device, dtype=torch.float32),
persistent=persistent,
seen=torch.zeros(batch_size, canvas_length,
device=device, dtype=torch.bool),
),
)
@dataclass
class LatentProcessorOutput:
context: torch.Tensor
state: LatentDeliberationState
def advance_trajectory_clocks(
ponder_steps: torch.Tensor,
stagnation_steps: torch.Tensor,
*,
commit_lengths: torch.LongTensor,
active_rows: torch.BoolTensor,
) -> tuple[torch.IntTensor, torch.IntTensor]:
"""Advance useful-ponder and stagnation clocks for each row."""
if not (
ponder_steps.shape == stagnation_steps.shape == commit_lengths.shape
== active_rows.shape
):
raise ValueError("Trajectory clock inputs must share shape [batch].")
committed = commit_lengths.gt(0)
waiting = active_rows & ~committed
next_ponder = torch.where(
committed, torch.zeros_like(ponder_steps), ponder_steps + waiting.to(torch.int32)
)
next_stagnation = torch.where(
committed,
torch.zeros_like(stagnation_steps),
stagnation_steps + waiting.to(torch.int32),
)
return next_ponder.to(torch.int32), next_stagnation.to(torch.int32)
def should_force_trajectory_jump(
stagnation_steps: torch.Tensor,
*,
progress_scores: torch.Tensor | None = None,
min_progress: float = 0.0,
stagnation_threshold: int,
ponder_steps: torch.Tensor | None = None,
max_ponder_steps: int | None = None,
) -> torch.BoolTensor:
jump = stagnation_steps.ge(stagnation_threshold)
if progress_scores is not None:
jump = jump & progress_scores.le(float(min_progress))
if ponder_steps is not None and max_ponder_steps is not None and max_ponder_steps > 0:
jump = jump | ponder_steps.ge(max_ponder_steps)
return jump.to(torch.bool)
class _RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1.0e-6) -> None:
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.eps = eps
def forward(self, hidden: torch.Tensor) -> torch.Tensor:
rms = hidden.float().square().mean(dim=-1, keepdim=True).add(self.eps).rsqrt()
return (hidden.float() * rms * self.weight.float()).to(dtype=hidden.dtype)
class _SwiGLU(nn.Module):
def __init__(self, dim: int, hidden: int) -> None:
super().__init__()
self.gate = nn.Linear(dim, hidden, bias=False)
self.up = nn.Linear(dim, hidden, bias=False)
self.down = nn.Linear(hidden, dim, bias=False)
def forward(self, hidden: torch.Tensor) -> torch.Tensor:
return self.down(F.silu(self.gate(hidden)) * self.up(hidden))
class _RankAttention(nn.Module):
"""Sequence attention in a rank-``kv_rank`` subspace, then map back to ``dim``."""
def __init__(self, dim: int, num_heads: int, kv_rank: int) -> None:
super().__init__()
if kv_rank % num_heads:
raise ValueError("`kv_rank` must be divisible by `num_heads`.")
self.num_heads = num_heads
self.kv_rank = kv_rank
self.head_dim = kv_rank // num_heads
self.q_proj = nn.Linear(dim, kv_rank, bias=False)
self.k_proj = nn.Linear(dim, kv_rank, bias=False)
self.v_proj = nn.Linear(dim, kv_rank, bias=False)
self.o_proj = nn.Linear(kv_rank, dim, bias=False)
self.q_norm = _RMSNorm(dim)
self.k_norm = _RMSNorm(dim)
def forward(
self,
query: torch.Tensor,
keys: torch.Tensor,
values: torch.Tensor,
attn_mask: torch.Tensor | None = None,
) -> torch.Tensor:
batch, queries, _dim = query.shape
key_len = keys.shape[1]
heads = self.num_heads
head_dim = self.head_dim
query = self.q_proj(self.q_norm(query)).view(batch, queries, heads, head_dim).transpose(1, 2)
keys = self.k_proj(self.k_norm(keys)).view(batch, key_len, heads, head_dim).transpose(1, 2)
values = self.v_proj(values).view(batch, key_len, heads, head_dim).transpose(1, 2)
mask = attn_mask
if mask is not None and mask.ndim == 2:
mask = mask.view(1, 1, queries, key_len)
elif mask is not None and mask.ndim == 3:
mask = mask.unsqueeze(1)
context = _fp32_scaled_dot_product_attention(
query, keys, values, attn_mask=mask
)
context = context.transpose(1, 2).reshape(batch, queries, self.kv_rank)
return self.o_proj(context)
class DecoderMemoryBus(nn.Module):
"""Read memory through a per-head gated residual."""
def __init__(
self,
hidden_size: int,
num_heads: int,
num_readers: int,
memory_dim: int,
kv_rank: int,
*,
relative_bias: bool = False,
address_with_identity: bool = False,
max_relative_span: int = 256,
) -> None:
super().__init__()
if num_readers > 0:
if kv_rank % num_heads:
raise ValueError("Memory bus rank must be divisible by heads.")
if hidden_size % num_heads:
raise ValueError("Memory bus hidden size must be divisible by heads.")
self.hidden_size = hidden_size
self.num_heads = num_heads
self.num_readers = num_readers
self.kv_rank = kv_rank
self.head_dim = kv_rank // num_heads if num_heads else kv_rank
self.relative_bias = relative_bias
self.address_with_identity = address_with_identity
self.memory_norm = _RMSNorm(memory_dim)
self.address_norm = _RMSNorm(memory_dim)
self.memory_to_hidden = (
nn.Identity()
if memory_dim == hidden_size
else nn.Linear(memory_dim, hidden_size, bias=False)
)
self.k_proj = nn.Linear(hidden_size, kv_rank, bias=False)
self.v_proj = nn.Linear(hidden_size, kv_rank, bias=False)
self.q_norm = _RMSNorm(hidden_size)
self.q_proj = nn.ModuleList(
[nn.Linear(hidden_size, kv_rank, bias=False) for _ in range(num_readers)]
)
self.o_proj = nn.ModuleList(
[nn.Linear(kv_rank, hidden_size, bias=False) for _ in range(num_readers)]
)
self.alpha = nn.Parameter(torch.zeros(max(num_readers, 1), max(num_heads, 1)))
span = max(2 * max_relative_span - 1, 1)
self.rel_bias = nn.Parameter(torch.zeros(max(num_heads, 1), span))
self.max_relative_span = max_relative_span
def prepare_kv(
self,
memory: torch.Tensor,
slot_identity: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor] | None:
if self.num_readers <= 0:
return None
if self.address_with_identity:
if slot_identity is None:
raise ValueError("Persistent bus requires slot identity on keys.")
mapped_keys = self.memory_to_hidden(self.address_norm(memory + slot_identity))
mapped_values = self.memory_to_hidden(self.memory_norm(memory))
else:
mapped_keys = mapped_values = self.memory_to_hidden(self.memory_norm(memory))
batch, slots, _dim = mapped_keys.shape
heads = self.num_heads
head_dim = self.head_dim
keys = self.k_proj(mapped_keys).view(batch, slots, heads, head_dim).transpose(1, 2)
values = self.v_proj(mapped_values).view(batch, slots, heads, head_dim).transpose(1, 2)
return keys, values
def _relative_mask(
self,
queries: int,
keys: int,
device: torch.device,
dtype: torch.dtype,
positions: torch.Tensor | None = None,
) -> torch.Tensor | None:
if not self.relative_bias:
return None
if positions is not None:
if positions.shape[1] != queries or queries != keys:
raise ValueError("Working memory positions must match both canvas axes.")
relative = (
positions[:, :, None] - positions[:, None, :] + (keys - 1)
).clamp(0, self.rel_bias.shape[1] - 1)
return self.rel_bias[:, relative].permute(1, 0, 2, 3).to(dtype=dtype)
q = torch.arange(queries, device=device)
k = torch.arange(keys, device=device)
rel = (q[:, None] - k[None, :] + (keys - 1)).clamp(0, self.rel_bias.shape[1] - 1)
return self.rel_bias[:, rel].to(dtype=dtype)
def read(
self,
hidden: torch.Tensor,
reader_index: int,
keys: torch.Tensor,
values: torch.Tensor,
positions: torch.Tensor | None = None,
*, key_seen: torch.Tensor | None = None,
) -> torch.Tensor:
batch, canvas, _dim = hidden.shape
heads = self.num_heads
head_dim = self.head_dim
query = self.q_proj[reader_index](self.q_norm(hidden))
query = query.view(batch, canvas, heads, head_dim).transpose(1, 2)
bias = self._relative_mask(
canvas, keys.shape[2], hidden.device, query.dtype, positions
)
if bias is not None and bias.ndim == 3:
bias = bias.unsqueeze(0)
if key_seen is not None:
key_mask = torch.zeros((batch, 1, 1, keys.shape[2]), device=hidden.device, dtype=torch.float32)
key_mask = key_mask.masked_fill(~key_seen[:, None, None, :], -1e9)
bias = key_mask if bias is None else bias.float() + key_mask
context = _fp32_scaled_dot_product_attention(
query, keys, values, attn_mask=bias
)
if key_seen is not None:
context = torch.where(key_seen.any(-1)[:, None, None, None], context, torch.zeros_like(context))
scale = torch.tanh(self.alpha[reader_index]).to(dtype=hidden.dtype).view(1, heads, 1, 1)
context = context * scale
context = context.transpose(1, 2).reshape(batch, canvas, self.kv_rank)
return hidden + self.o_proj[reader_index](context)
@torch.no_grad()
def reset_identity_parameters(self) -> None:
self.alpha.zero_()
self.rel_bias.zero_()
def slice_latent_state(
state: LatentDeliberationState, rows: slice | torch.Tensor
) -> LatentDeliberationState:
return LatentDeliberationState(
memory_slots=state.memory_slots[rows],
confidence=state.confidence[rows],
entropy=state.entropy[rows],
ponder_steps=state.ponder_steps[rows],
stagnation_steps=state.stagnation_steps[rows],
gdn2=GDN2TrajectoryState(
cells=state.gdn2.cells[rows], row=state.gdn2.row[rows],
persistent=state.gdn2.persistent[rows], seen=state.gdn2.seen[rows],
),
)
def cat_latent_states(states: Sequence[LatentDeliberationState]) -> LatentDeliberationState:
def cat(name: str) -> torch.Tensor:
return torch.cat([getattr(state, name) for state in states], dim=0)
return LatentDeliberationState(
memory_slots=cat("memory_slots"),
confidence=cat("confidence"),
entropy=cat("entropy"),
ponder_steps=cat("ponder_steps"),
stagnation_steps=cat("stagnation_steps"),
gdn2=GDN2TrajectoryState(
cells=torch.cat([state.gdn2.cells for state in states], dim=0),
row=torch.cat([state.gdn2.row for state in states], dim=0),
persistent=torch.cat([state.gdn2.persistent for state in states], dim=0),
seen=torch.cat([state.gdn2.seen for state in states], dim=0),
),
)
def infer_commit_reason(
commit_lengths: torch.Tensor,
*,
jump_rows: torch.Tensor | None = None,
commit_token_ids: torch.Tensor | None = None,
terminal_token_ids: Sequence[int] = (),
) -> torch.Tensor:
"""Return per-row commit-reason codes. No hard skip; writer sees the label."""
reasons = torch.full(
commit_lengths.shape,
COMMIT_REASON_NONE,
device=commit_lengths.device,
dtype=torch.long,
)
committed = commit_lengths.gt(0)
default = (
COMMIT_REASON_NORMAL
)
reasons = torch.where(committed, torch.full_like(reasons, default), reasons)
if jump_rows is not None:
reasons = torch.where(
committed & jump_rows.to(dtype=torch.bool),
torch.full_like(reasons, COMMIT_REASON_FORCED_JUMP),
reasons,
)
if commit_token_ids is not None and terminal_token_ids:
positions = torch.arange(
commit_token_ids.shape[1], device=commit_token_ids.device
)[None, :]
selected = positions.lt(commit_lengths[:, None])
terminal = torch.zeros_like(committed)
for token_id in terminal_token_ids:
terminal |= (commit_token_ids.eq(int(token_id)) & selected).any(dim=-1)
reasons = torch.where(
committed & terminal,
torch.full_like(reasons, COMMIT_REASON_TERMINAL),
reasons,
)
return reasons
class _CanvasBlock(nn.Module):
def __init__(self, width: int, heads: int, rank: int, ffn: int,
window: int, global_attention: bool) -> None:
super().__init__()
self.norm = _RMSNorm(width)
self.attn = _RankAttention(width, heads, rank)
self.ff_norm = _RMSNorm(width)
self.ff = _SwiGLU(width, ffn)
self.window = window
self.global_attention = global_attention
def forward(self, hidden: torch.Tensor, seen: torch.Tensor,
offsets: torch.Tensor) -> torch.Tensor:
allowed = seen[:, None, :].expand(-1, hidden.shape[1], -1)
if not self.global_attention:
allowed = allowed & ((offsets[:, :, None] - offsets[:, None, :]).abs() < self.window)
additive = torch.zeros(allowed.shape, device=hidden.device, dtype=torch.float32)
additive = additive.masked_fill(~allowed, -1e9)
normed = self.norm(hidden)
hidden = hidden + self.attn(normed, normed, normed, attn_mask=additive)
hidden = hidden + self.ff(self.ff_norm(hidden))
return torch.where(seen[..., None], hidden, torch.zeros_like(hidden))
class _WorkingBus(DecoderMemoryBus):
def prepare_kv(self, memory: tuple[torch.Tensor, torch.Tensor],
slot_identity: torch.Tensor | None = None):
del slot_identity
working, seen = memory
pair = super().prepare_kv(working)
return None if pair is None else (*pair, seen)
def read(self, hidden: torch.Tensor, reader_index: int,
keys: torch.Tensor, values: torch.Tensor, seen: torch.Tensor,
positions: torch.Tensor | None = None) -> torch.Tensor:
written = super().read(hidden, reader_index, keys, values, positions, key_seen=seen)
return torch.where(seen[..., None], written, hidden)
class _PersistentBus(nn.Module):
def __init__(self, memory: nn.Module, readers: int) -> None:
super().__init__()
object.__setattr__(self, "_memory_ref", weakref.ref(memory))
self.num_readers = readers
self.alpha = nn.Parameter(torch.zeros(max(readers, 1), 1))
def reset_identity_parameters(self) -> None:
with torch.no_grad():
self.alpha.zero_()
def prepare_kv(self, memory: torch.Tensor,
seen: torch.Tensor | None = None):
if self.num_readers <= 0:
return None
return memory, seen
def read(self, hidden: torch.Tensor, reader_index: int,
memory: torch.Tensor, seen: torch.Tensor | None) -> torch.Tensor:
delta = self._memory_ref().read_shared(memory, hidden)
if seen is not None:
delta = delta * seen[..., None].to(delta.dtype)
return hidden + torch.tanh(self.alpha[reader_index]).to(hidden.dtype) * delta
class LatentDeliberationTransformer(nn.Module):
def __init__(self, *, hidden_size: int,
latent_dim: int = 2816, ffn_dim: int = 7168,
num_layers: int = 4,
num_heads: int = 16, local_attention_window: int = 128,
tape_probes: int = 4,
history_kv_rank: int = 1024, num_memory_readers: int = 0,
num_working_readers: int | None = None,
num_persistent_readers: int | None = None,
working_last_block_global: bool = True,
commit_sequence_dim: int | None = None,
max_canvas_length: int = 256) -> None:
super().__init__()
if hidden_size != latent_dim:
raise ValueError("Schema25 GDN2 requires hidden_size == latent_dim.")
self.hidden_size = hidden_size
self.latent_dim = latent_dim
self.tape_probes = tape_probes
self.packet_dim = int(commit_sequence_dim or history_kv_rank)
if self.packet_dim % num_heads:
raise ValueError("Commit packet rank must be divisible by attention heads.")
self.trajectory = GDN2TrajectoryMemory(
hidden_size, probes=tape_probes, persistent_observation_dim=self.packet_dim,
)
self.blocks = nn.ModuleList([
_CanvasBlock(hidden_size, num_heads, history_kv_rank, ffn_dim,
local_attention_window,
bool(working_last_block_global and i == num_layers - 1))
for i in range(num_layers)
])
self.output_norm = _RMSNorm(hidden_size)
readers = num_memory_readers if num_working_readers is None else num_working_readers
persistent_readers = (
num_memory_readers if num_persistent_readers is None else num_persistent_readers
)
bus_rank = history_kv_rank if history_kv_rank % num_heads == 0 else num_heads
self.working_memory_bus = _WorkingBus(
hidden_size, num_heads, readers, hidden_size, bus_rank,
relative_bias=True, max_relative_span=max_canvas_length,
)
self.persistent_memory_bus = _PersistentBus(
self.trajectory.persistent, persistent_readers,
)
self.experience_in = nn.Linear(hidden_size * 3, self.packet_dim, bias=False)
self.reason_embed = nn.Embedding(6, self.packet_dim)
def reset_identity_parameters(self) -> None:
self.working_memory_bus.reset_identity_parameters()
self.persistent_memory_bus.reset_identity_parameters()
def forward(self, *, token_embeddings: torch.Tensor, confidence: torch.Tensor,
entropy: torch.Tensor, state: LatentDeliberationState,
canvas_head: torch.Tensor | None = None) -> LatentProcessorOutput:
batch, canvas, width = token_embeddings.shape
if state.memory_slots.shape != state.gdn2.persistent.shape:
raise ValueError("Persistent GDN2 state shape differs from memory slots.")
memory = replace(state.gdn2, persistent=state.memory_slots)
hidden = self.trajectory.read(memory, token_embeddings)
offsets = torch.arange(canvas, device=token_embeddings.device)[None, :].expand(batch, -1)
if canvas_head is not None:
offsets = (offsets - canvas_head[:, None]) % canvas
for block in self.blocks:
hidden = block(hidden, memory.seen, offsets)
hidden = self.output_norm(hidden) * memory.seen[..., None].to(hidden.dtype)
next_state = replace(state, confidence=confidence.float(), entropy=entropy.float(),
gdn2=memory)
return LatentProcessorOutput(hidden, next_state)
def observe_state(self, state: LatentDeliberationState,
heavy: torch.Tensor, working: torch.Tensor,
live: torch.Tensor, head: torch.Tensor) -> LatentDeliberationState:
source = heavy.detach() + working
updated = self.trajectory.observe(state.gdn2, source, live, head)
return replace(state, gdn2=updated)
def commit_write(self, *, memory: torch.Tensor, working_state: torch.Tensor,
heavy_hidden: torch.Tensor,
committed_token_embeddings: torch.Tensor,
commit_lengths: torch.Tensor,
commit_reason: torch.Tensor | None = None,
canvas_head: torch.Tensor | None = None):
batch, canvas, width = working_state.shape
count = int(commit_lengths.max().item())
if count <= 0:
return memory
index = torch.arange(count, device=working_state.device)[None, :].expand(batch, -1)
if canvas_head is not None:
index = (index + canvas_head[:, None]) % canvas
selected_working = working_state.gather(
1, index[..., None].expand(-1, -1, width)
)
selected_heavy = heavy_hidden.detach().gather(
1, index[..., None].expand(-1, -1, width)
)
if committed_token_embeddings.shape != selected_working.shape:
raise ValueError("Committed embeddings do not match the prefix.")
# Canvas processing already supplies bidirectional spatial context.
# Separate normalized role channels feed the ordered GDN2 writer directly.
# Unit-floor normalization keeps a zero Working state at zero without
# amplifying its derivative by 1/sqrt(eps) on the first denoise.
roles = tuple(F.rms_norm(value.float(), (width,), eps=1.0).to(value.dtype)
for value in (selected_heavy, selected_working,
committed_token_embeddings.detach()))
packet = self.experience_in(torch.cat(roles, dim=-1))
reason = torch.zeros(batch, device=packet.device, dtype=torch.long) if commit_reason is None else commit_reason.long()
reason = torch.where((reason == 1) | (reason == 2), 5, reason).clamp(0, 5)
packet = packet + self.reason_embed(reason)[:, None]
valid = torch.arange(count, device=packet.device)[None, :] < commit_lengths[:, None]
written = self.trajectory.persistent.write_sequence(memory, packet, valid)
return written