Text Generation
Transformers
Safetensors
modilify_mk2
diffusion
mixture-of-experts
trust-remote-code
conversational
custom_code
Instructions to use modilify/Modilify-Mk2-preview with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use modilify/Modilify-Mk2-preview with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="modilify/Modilify-Mk2-preview", trust_remote_code=True) messages = [ {"role": "user", "content": "Who are you?"}, ] pipe(messages)# pip install -U transformers accelerate # Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("modilify/Modilify-Mk2-preview", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use modilify/Modilify-Mk2-preview with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "modilify/Modilify-Mk2-preview" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "modilify/Modilify-Mk2-preview", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker
docker model run hf.co/modilify/Modilify-Mk2-preview
- SGLang
How to use modilify/Modilify-Mk2-preview with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "modilify/Modilify-Mk2-preview" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "modilify/Modilify-Mk2-preview", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "modilify/Modilify-Mk2-preview" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "modilify/Modilify-Mk2-preview", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }' - Docker Model Runner
How to use modilify/Modilify-Mk2-preview with Docker Model Runner:
docker model run hf.co/modilify/Modilify-Mk2-preview
Download latent_deliberation.py from modilify/Modilify-Mk2-preview: direct link, hf CLI and curl.
- Browser
- Download file 26.2 kB
-
https://huggingface.co/modilify/Modilify-Mk2-preview/resolve/main/latent_deliberation.py
- Command line
-
hf download hf://modilify/Modilify-Mk2-preview/latent_deliberation.py
-
curl -L -o latent_deliberation.py https://huggingface.co/modilify/Modilify-Mk2-preview/resolve/main/latent_deliberation.py
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) | |
| 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 | |
| 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), | |
| ), | |
| ) | |
| 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) | |
| 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 | |