diff --git "a/modeling_nested_mamba.py" "b/modeling_nested_mamba.py" new file mode 100644--- /dev/null +++ "b/modeling_nested_mamba.py" @@ -0,0 +1,2055 @@ +"""Inference-only architecture for the three-level nested byte Mamba-2 model.""" +from __future__ import annotations + +import copy +import math +import threading +from collections import deque +from typing import Dict, List, Optional, Tuple + +import torch +import torch.nn as nn +import torch.nn.functional as F + +try: + from mamba_ssm.utils.generation import InferenceParams +except Exception: + InferenceParams = None + +MAMBA_IMPORT_ERROR: Optional[Exception] = None +try: + from mamba_ssm import Mamba, Mamba2 +except Exception as exc: + Mamba = None + Mamba2 = None + MAMBA_IMPORT_ERROR = exc + + +class BLTGlobalBlock(nn.Module): + def __init__( + self, + dim: int, + use_mamba: bool = True, + ff_mult: int = 2, + layer_idx: int = 0, + mamba_version: int = 2, + mamba_d_state: int = 64, + mamba2_headdim: int = 0, + ): + super().__init__() + self.norm = nn.LayerNorm(dim) + self.mamba_version = int(mamba_version) + self.mamba_d_state = int(mamba_d_state) + self.mamba2_headdim = int(mamba2_headdim) + if self.mamba_version not in (1, 2): + raise ValueError(f"mamba_version must be 1 or 2, got {self.mamba_version}") + if self.mamba_d_state < 1: + raise ValueError("mamba_d_state must be positive") + if use_mamba and self.mamba_version == 2: + if Mamba2 is None: + raise RuntimeError( + "Mamba-2 was requested but mamba_ssm.Mamba2 is unavailable" + ) from MAMBA_IMPORT_ERROR + inner_dim = int(dim) * 2 + if self.mamba2_headdim <= 0: + self.mamba2_headdim = next( + (candidate for candidate in (64, 128, 100, 80, 50, 40, 32, 25, 20, 16, 10, 8, 5, 4, 2, 1) + if inner_dim % candidate == 0), + 1, + ) + if inner_dim % self.mamba2_headdim: + raise ValueError( + f"Mamba-2 inner dimension {inner_dim} (2*dim) must be divisible by " + f"mamba2_headdim={self.mamba2_headdim}" + ) + self.mixer = Mamba2( + d_model=dim, + d_state=self.mamba_d_state, + d_conv=4, + expand=2, + headdim=self.mamba2_headdim, + layer_idx=layer_idx, + ) + self.kind = "mamba" + elif use_mamba and Mamba is not None: + try: + self.mixer = Mamba( + d_model=dim, + d_state=self.mamba_d_state, + d_conv=4, + expand=2, + layer_idx=layer_idx, + ) + except TypeError: + self.mixer = Mamba( + d_model=dim, d_state=self.mamba_d_state, d_conv=4, expand=2 + ) + self.mixer.layer_idx = layer_idx + self.kind = "mamba" + else: + self.mixer = nn.GRU(dim, dim, batch_first=True) + self.kind = "gru" + self.ff = nn.Sequential( + nn.LayerNorm(dim), + nn.Linear(dim, dim * ff_mult), + nn.GELU(), + nn.Linear(dim * ff_mult, dim), + ) + + def forward(self, x: torch.Tensor, inference_params=None) -> torch.Tensor: + h = self.norm(x) + if self.kind == "gru": + h, _ = self.mixer(h) + elif inference_params is None: + h = self.mixer(h) + else: + h = self.mixer(h, inference_params=inference_params) + x = x + h + return x + self.ff(x) + + +class CausalConv1d(nn.Module): + def __init__(self, dim: int, kernel_size: int): + super().__init__() + self.left_pad = max(0, int(kernel_size) - 1) + self.conv = nn.Conv1d(dim, dim, kernel_size=kernel_size, padding=0, groups=1) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if self.left_pad: + x = F.pad(x, (self.left_pad, 0)) + return self.conv(x) + + +class LearnedPoolEncoder(nn.Module): + """Causal, learned replacement for the old mean patch autoencoder. + + ``close_probability`` is used to decide whether an eligible pool ends; a + maximum length still forces a close, so the global stream has a bounded + memory footprint. The same probabilities also participate in the pooled + representation, which gives the close head a learning signal from the LM + objective (rather than making it a detached routing heuristic). + """ + + def __init__(self, dim: int, *, detach_close_content_gradient: bool = True): + super().__init__() + self.detach_close_content_gradient = bool(detach_close_content_gradient) + # A recurrent encoder is impractical for the 660k-token training + # windows this trainer supports (and cuDNN rejects some such shapes). + # This is still a learned causal encoder, but its convolutional form + # is safe for very long windows and compatible with AMP/cuDNN. + self.context_norm = nn.LayerNorm(dim) + self.context = nn.Conv1d(dim, dim, kernel_size=3, padding=0) + self.value = nn.Sequential(nn.LayerNorm(dim), nn.Linear(dim, dim), nn.GELU(), nn.Linear(dim, dim)) + self.weight = nn.Sequential(nn.LayerNorm(dim), nn.Linear(dim, 1)) + self.close = nn.Sequential(nn.LayerNorm(dim), nn.Linear(dim, dim), nn.GELU(), nn.Linear(dim, 1)) + self.close_value = nn.Linear(1, dim, bias=False) + # Start with stable max-sized pools; the head can learn earlier closes + # without immediately multiplying the global sequence length at init. + nn.init.constant_(self.close[-1].bias, -4.0) + + def contextualize(self, local_h: torch.Tensor) -> torch.Tensor: + h = self.context_norm(local_h).transpose(1, 2) + h = self.context(F.pad(h, (2, 0))).transpose(1, 2) + return local_h + F.gelu(h) + + def close_probability(self, contextual_h: torch.Tensor) -> torch.Tensor: + return torch.sigmoid(self.close_logits(contextual_h)) + + def close_logits(self, contextual_h: torch.Tensor) -> torch.Tensor: + return self.close(contextual_h).squeeze(-1) + + def pool( + self, + contextual_h: torch.Tensor, + patch_ids: torch.Tensor, + patch_counts: torch.Tensor, + token_mask: torch.Tensor, + max_patches: int, + ) -> torch.Tensor: + b, _, c = contextual_h.shape + value = self.value(contextual_h) + # Preserve existing checkpoint forward values while preventing the LM + # objective from training the boundary probability as a content side + # channel. Boundary heads receive their own causal routing objective. + close_prob = self.close_probability(contextual_h) + if self.detach_close_content_gradient: + close_prob = close_prob.detach() + close_prob = close_prob.unsqueeze(-1) + weight = torch.sigmoid(self.weight(contextual_h)) * token_mask.unsqueeze(-1).to(contextual_h.dtype) + encoded = (value + self.close_value(close_prob)) * weight + pooled = contextual_h.new_zeros((b, max_patches, c)) + normalizer = contextual_h.new_zeros((b, max_patches, 1)) + index = patch_ids.unsqueeze(-1) + pooled.scatter_add_(1, index.expand(-1, -1, c), encoded) + normalizer.scatter_add_(1, index, weight) + return pooled / normalizer.clamp_min(1e-6) + + +class ByteLatentMambaCore(nn.Module): + """Causal hierarchical BLT with adaptive, bounded level transitions.""" + + def __init__( + self, + vocab_size: int, + dim: int = 256, + layers: int = 4, + position_bins: int = 8192, + use_mamba: bool = True, + mamba_version: int = 2, + mamba_d_state: int = 64, + mamba2_headdim: int = 0, + num_sections: int = 16, + min_patch_bytes: int = 4, + max_patch_bytes: int = 16, + patch_change_threshold: int = 48, + close_threshold: float = 0.90, + mid_close_bonus: float = 0.05, + nested_pool_factor: int = 16, + nested_min_pool_factor: int = 0, + nested_close_threshold: Optional[float] = None, + nested_layers: int = 0, + tertiary_pool_factor: int = 0, + tertiary_min_pool_factor: int = 0, + tertiary_close_threshold: Optional[float] = None, + tertiary_layers: int = 0, + decoder_dim: int = 0, + detach_inactive_coarse_gradients: bool = True, + legacy_pool_closure_training: bool = False, + decoder_pool_controller: bool = False, + pool_controller_alpha: float = 1.0, + pool_controller_beta: float = 0.5, + pool_controller_gamma: float = 0.5, + short_pool_budget: int = 0, + short_pool_window: int = 0, + secondary_min_patch_bytes: int = 16, + ): + super().__init__() + self.vocab_size = int(vocab_size) + self.mamba_version = int(mamba_version) + self.mamba_d_state = int(mamba_d_state) + inner_dim = int(dim) * 2 + requested_headdim = int(mamba2_headdim) + self.mamba2_headdim = requested_headdim + if self.mamba_version == 2 and self.mamba2_headdim <= 0: + self.mamba2_headdim = next( + (candidate for candidate in (64, 128, 100, 80, 50, 40, 32, 25, 20, 16, 10, 8, 5, 4, 2, 1) + if inner_dim % candidate == 0), + 1, + ) + self.min_patch_bytes = max(1, int(min_patch_bytes)) + self.max_patch_bytes = max(self.min_patch_bytes, int(max_patch_bytes)) + self.patch_change_threshold = max(0, int(patch_change_threshold)) + self.close_threshold = max(0.0, min(1.0, float(close_threshold))) + self.mid_close_bonus = max(0.0, float(mid_close_bonus)) + self.nested_pool_factor = max(2, int(nested_pool_factor)) + self.nested_min_pool_factor = ( + self.nested_pool_factor + if int(nested_min_pool_factor) <= 0 + else max(1, min(int(nested_min_pool_factor), self.nested_pool_factor)) + ) + self.nested_close_threshold = self.close_threshold if nested_close_threshold is None else max( + 0.0, min(1.0, float(nested_close_threshold)) + ) + self.tertiary_pool_factor = int(tertiary_pool_factor) + self.tertiary_min_pool_factor = ( + self.tertiary_pool_factor + if int(tertiary_min_pool_factor) <= 0 + else max(1, min(int(tertiary_min_pool_factor), self.tertiary_pool_factor)) + ) + self.tertiary_close_threshold = self.close_threshold if tertiary_close_threshold is None else max( + 0.0, min(1.0, float(tertiary_close_threshold)) + ) + self.tertiary_layers = int(tertiary_layers) + self.tertiary_enabled = self.tertiary_pool_factor >= 2 and self.tertiary_layers > 0 + if (self.tertiary_pool_factor > 0 or self.tertiary_layers > 0) and not self.tertiary_enabled: + raise ValueError("tertiary_pool_factor must be at least 2 and tertiary_layers must be positive together") + self.nested_layers = max(1, int(nested_layers) if int(nested_layers) > 0 else max(1, int(layers) // 2)) + self.fine_layers = max(0, int(layers) - self.nested_layers) + self.decoder_dim = int(decoder_dim) if int(decoder_dim) > 0 else int(dim) * 2 + if self.decoder_dim < int(dim) * 2: + raise ValueError("decoder_dim must be 0 (the 2*dim default) or at least 2*dim") + self.detach_inactive_coarse_gradients = bool(detach_inactive_coarse_gradients) + self.legacy_pool_closure_training = bool(legacy_pool_closure_training) + self.decoder_pool_controller = bool(decoder_pool_controller) + self.pool_controller_alpha = float(pool_controller_alpha) + self.pool_controller_beta = float(pool_controller_beta) + self.pool_controller_gamma = float(pool_controller_gamma) + self.short_pool_budget = max(0, int(short_pool_budget)) + self.short_pool_window = max(0, int(short_pool_window)) + self.secondary_min_patch_bytes = min( + self.max_patch_bytes, + max(self.min_patch_bytes, int(secondary_min_patch_bytes)), + ) + self.byte_offset = 4 + self.byte_emb = nn.Embedding(vocab_size, dim) + self.local_norm = nn.LayerNorm(dim) + self.local_conv = nn.Sequential( + CausalConv1d(dim, kernel_size=5), + nn.GELU(), + CausalConv1d(dim, kernel_size=3), + nn.GELU(), + ) + pool_encoder_kwargs = { + "detach_close_content_gradient": not self.legacy_pool_closure_training + } + self.pool_encoder = LearnedPoolEncoder(dim, **pool_encoder_kwargs) + # The second pooler operates on completed first-level pool states. + # Fixed-size grouping keeps this hierarchy vectorized and bounded; + # its representation/weights remain learned through this encoder. + self.nested_pool_encoder = LearnedPoolEncoder(dim, **pool_encoder_kwargs) + self.tertiary_pool_encoder = ( + LearnedPoolEncoder(dim, **pool_encoder_kwargs) + if self.tertiary_enabled else None + ) + self.boe_patch = nn.Parameter(torch.zeros(1, 1, dim)) + self.patch_pos_emb = nn.Embedding(position_bins, dim) + self.patch_len_emb = nn.Embedding(self.max_patch_bytes + 1, dim) + self.nested_len_emb = nn.Embedding(self.nested_pool_factor + 1, dim) + self.tertiary_len_emb = nn.Embedding(self.tertiary_pool_factor + 1, dim) if self.tertiary_enabled else None + block_kwargs = dict( + dim=dim, + use_mamba=use_mamba, + mamba_version=self.mamba_version, + mamba_d_state=self.mamba_d_state, + mamba2_headdim=self.mamba2_headdim, + ) + self.global_blocks = nn.ModuleList([ + BLTGlobalBlock(**block_kwargs, layer_idx=i) for i in range(self.fine_layers) + ]) + self.nested_global_blocks = nn.ModuleList([ + BLTGlobalBlock(**block_kwargs, layer_idx=self.fine_layers + i) + for i in range(self.nested_layers) + ]) + self.tertiary_global_blocks = nn.ModuleList([ + BLTGlobalBlock( + **block_kwargs, + layer_idx=self.fine_layers + self.nested_layers + i, + ) + for i in range(self.tertiary_layers) + ]) + self.decoder = nn.Sequential( + nn.LayerNorm(dim * (5 if self.tertiary_enabled else 4)), + nn.Linear(dim * (5 if self.tertiary_enabled else 4), self.decoder_dim), + nn.GELU(), + nn.Linear(self.decoder_dim, dim), + nn.GELU(), + ) + # Three causal readouts (hold, refresh, compress) for each hierarchy. + # The zero initialization produces 0.5 for every readout. With the + # default alpha=beta+gamma this is an exact no-op, which lets an older + # checkpoint acquire the controller without changing its first + # generated byte or its initial routing policy. + self.pool_controller = nn.Linear(dim, 9) if self.decoder_pool_controller else None + if self.pool_controller is not None: + nn.init.zeros_(self.pool_controller.weight) + nn.init.zeros_(self.pool_controller.bias) + self.lm_head = nn.Linear(dim, vocab_size, bias=False) + self.position_bins = int(position_bins) + # Placement is deliberately runtime-only: checkpoints contain the + # ordinary architecture state dict and can still be loaded on CPU, a + # single GPU, or a browser-export host. + self.model_parallel_enabled = False + # Standalone evaluation may request byte-aligned hierarchy details. + # Keep this disabled during training so large activations are not + # retained after forward/backward. + self.capture_evaluation_details = False + self._evaluation_details_by_thread: Dict[int, object] = {} + self.last_forward_evaluation_details = None + self.last_pool_controller_scores = None + self.last_pool_controller_close_signals = None + self.last_pool_routing_training = None + self.fine_device = None + self.nested_devices: List[torch.device] = [] + + @property + def last_forward_evaluation_details(self): + return self._evaluation_details_by_thread.get(threading.get_ident()) + + @last_forward_evaluation_details.setter + def last_forward_evaluation_details(self, value) -> None: + thread_id = threading.get_ident() + if value is None: + self._evaluation_details_by_thread.pop(thread_id, None) + else: + self._evaluation_details_by_thread[thread_id] = value + + def clear_forward_evaluation_details(self) -> None: + self._evaluation_details_by_thread.clear() + + def _pool_controller_scores(self, decoded: torch.Tensor) -> Optional[torch.Tensor]: + """Return [batch, time, level, (hold, refresh, compress)]. + + Routing is discrete, so controller supervision intentionally does not + backpropagate into the already-trained decoder. Recovery therefore + teaches only the newly introduced readout and cannot damage the + decoder while its policy is being calibrated. + """ + if self.pool_controller is None: + return None + return torch.sigmoid(self.pool_controller(decoded.detach())).view( + *decoded.shape[:-1], 3, 3 + ) + + def _controller_close_adjustment(self, scores, level: int) -> float: + if scores is None or not self.decoder_pool_controller: + return 0.0 + level_scores = scores[level] + hold, refresh, compress = ( + float(value.detach().item()) if isinstance(value, torch.Tensor) else float(value) + for value in level_scores + ) + # Centering makes the zero-initialized 0.5/0.5/0.5 controller neutral + # for every coefficient combination, not only the defaults. + return ( + self.pool_controller_alpha * (refresh - 0.5) + - self.pool_controller_beta * (hold - 0.5) + - self.pool_controller_gamma * (compress - 0.5) + ) + + def pool_controller_auxiliary_loss( + self, logits: torch.Tensor, targets: torch.Tensor + ) -> torch.Tensor: + """Supervise controller semantics without differentiating hard routing. + + The existing close head supplies the refresh target. Remaining mass + is split between hold and compress using detached next-byte surprise: + difficult bytes request more local evidence (hold), while predictable + bytes permit compression. Each level receives its own aligned close + signal from the encoder at that level. + """ + scores = self.last_pool_controller_scores + signals = self.last_pool_controller_close_signals + if scores is None or signals is None: + return logits.new_zeros(()) + valid = targets.ne(-100) + with torch.no_grad(): + nll = F.cross_entropy( + logits.detach().transpose(1, 2), targets, ignore_index=-100, + reduction="none", + ) + difficulty = nll / (nll + 2.0) + refresh = signals.detach().clamp(0.0, 1.0) + unresolved = 1.0 - refresh + target_scores = torch.stack( + [ + unresolved * difficulty.unsqueeze(-1), + refresh, + unresolved * (1.0 - difficulty.unsqueeze(-1)), + ], + dim=-1, + ) + # C[t] controls the boundary evaluated at the beginning of step t+1. + # Train it against that future step rather than the byte that emitted + # it. The final controller output has no within-window target. + future_scores = scores[:, :-1] + target_scores = target_scores[:, 1:] + future_valid = valid[:, 1:] + expanded_valid = future_valid.unsqueeze(-1).unsqueeze(-1).expand_as( + future_scores + ) + if not bool(expanded_valid.any()): + return logits.new_zeros(()) + # PyTorch deliberately rejects probability-space BCE inside CUDA + # autocast because its sigmoid gradient can underflow in FP16/BF16. + # Keep the checkpoint-compatible sigmoid controller, but evaluate this + # small auxiliary objective in FP32 outside the surrounding training + # autocast region. Gradients still flow through the float conversion + # into pool_controller; decoded was intentionally detached above. + with torch.autocast(device_type=future_scores.device.type, enabled=False): + return F.binary_cross_entropy( + future_scores[expanded_valid].float(), + target_scores[expanded_valid].float(), + ) + + def configure_model_parallel( + self, + fine_device: torch.device, + nested_devices: List[torch.device], + tertiary_device: Optional[torch.device] = None, + ) -> None: + """Place fine BLT work on one device and coarse levels on others. + + The byte embedding, local convolution, first-level pooling, fine + blocks, and decoder stay together because they carry the long sequence. + Only completed fine-pool states cross to the coarse hierarchy. Coarse + blocks are assigned in contiguous chunks to avoid an interconnect hop + after every layer. + """ + if not nested_devices: + raise ValueError("nested model parallelism requires at least one coarse device") + fine_device = torch.device(fine_device) + nested_devices = [torch.device(item) for item in nested_devices] + if fine_device.type != "cuda" or any(item.type != "cuda" for item in nested_devices): + raise ValueError("nested model parallelism requires CUDA devices") + + # Start with a coherent root copy, then move just the coarse stage. + self.to(fine_device) + coarse_root = nested_devices[0] + self.nested_pool_encoder.to(coarse_root) + self.nested_len_emb.to(coarse_root) + if self.tertiary_enabled: + if tertiary_device is None: + raise ValueError("three-level nested model parallelism requires a tertiary device") + tertiary_device = torch.device(tertiary_device) + if tertiary_device.type != "cuda": + raise ValueError("tertiary model-parallel device must be CUDA") + assert self.tertiary_pool_encoder is not None and self.tertiary_len_emb is not None + self.tertiary_pool_encoder.to(tertiary_device) + self.tertiary_len_emb.to(tertiary_device) + for block in self.tertiary_global_blocks: + block.to(tertiary_device) + block_count = len(self.nested_global_blocks) + for index, block in enumerate(self.nested_global_blocks): + # Contiguous placement gives one transfer per participating GPU, + # rather than alternating CUDA devices every coarse layer. + device_index = min(len(nested_devices) - 1, index * len(nested_devices) // max(1, block_count)) + block.to(nested_devices[device_index]) + self.model_parallel_enabled = True + self.fine_device = fine_device + self.nested_devices = nested_devices + + @staticmethod + def _module_device(module: nn.Module) -> torch.device: + return next(module.parameters()).device + + @classmethod + def _run_block_on_own_device( + cls, + block: nn.Module, + h: torch.Tensor, + inference_params=None, + ) -> torch.Tensor: + """Launch Triton-backed mixers under their parameter device context. + + Ordinary PyTorch operators dispatch from tensor placement, but + Mamba-2's Triton SSD kernels use CUDA's current device when launching. + In a model-parallel process that current device can remain cuda:0 even + while a level and its activations live on cuda:1 or cuda:2, producing + Triton's misleading "cpu tensor?" pointer error. + """ + block_device = cls._module_device(block) + if h.device != block_device: + h = h.to(block_device, non_blocking=True) + if block_device.type == "cuda": + with torch.cuda.device(block_device): + return block(h, inference_params=inference_params) + return block(h, inference_params=inference_params) + + def _run_nested_blocks(self, nested_global_h: torch.Tensor, nested_global_mask: torch.Tensor) -> torch.Tensor: + """Run coarse blocks, transferring only at device-stage boundaries.""" + h = nested_global_h + mask = nested_global_mask + for block in self.nested_global_blocks: + block_device = self._module_device(block) + if h.device != block_device: + h = h.to(block_device, non_blocking=True) + mask = mask.to(block_device, non_blocking=True) + h = self._run_block_on_own_device(block, h, inference_params=None) + h = h.masked_fill(~mask.unsqueeze(-1), 0.0) + return h + + def _run_tertiary_blocks(self, h: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: + for block in self.tertiary_global_blocks: + block_device = self._module_device(block) + if h.device != block_device: + h, mask = h.to(block_device, non_blocking=True), mask.to(block_device, non_blocking=True) + h = self._run_block_on_own_device( + block, h, inference_params=None + ).masked_fill(~mask.unsqueeze(-1), 0.0) + return h + + def _local_features(self, x: torch.Tensor) -> torch.Tensor: + h = self.byte_emb(x.clamp(0, self.vocab_size - 1)) + conv = self.local_conv(self.local_norm(h).transpose(1, 2)).transpose(1, 2) + return h + conv + + @staticmethod + def _stream_conv_last( + conv: nn.Conv1d, + current: torch.Tensor, + state: Dict[str, object], + cache_key: str, + ) -> torch.Tensor: + """Evaluate one causal Conv1d position and retain only its left context.""" + kernel_size = int(conv.kernel_size[0]) + history = state.get(cache_key) + values = current if history is None else torch.cat([history, current], dim=1) + conv_input = values + if values.shape[1] < kernel_size: + conv_input = F.pad(values.transpose(1, 2), (kernel_size - values.shape[1], 0)) + else: + conv_input = values[:, -kernel_size:].transpose(1, 2) + output = conv(conv_input).transpose(1, 2)[:, -1:, :] + keep = max(0, kernel_size - 1) + state[cache_key] = values[:, -keep:].detach() if keep else None + return output + + def _local_stream_context(self, token: torch.Tensor, state: Dict[str, object]) -> torch.Tensor: + """Incrementally compute the exact causal byte and pool-context feature.""" + embedded = self.byte_emb(token.clamp(0, self.vocab_size - 1)) + conv1 = self.local_conv[0] + conv2 = self.local_conv[2] + first = self._stream_conv_last( + conv1.conv, self.local_norm(embedded), state, "local_conv1_tail" + ) + first = F.gelu(first) + second = self._stream_conv_last( + conv2.conv, first, state, "local_conv2_tail" + ) + local_h = embedded + F.gelu(second) + contextual_delta = self._stream_conv_last( + self.pool_encoder.context, + self.pool_encoder.context_norm(local_h), + state, + "pool_context_tail", + ) + return local_h + F.gelu(contextual_delta) + + def pool_close_probabilities(self, x: torch.Tensor) -> torch.Tensor: + """Run only the learned encoder, for greedy-policy bootstrap training.""" + local_h = self._local_features(x) + return self.pool_encoder.close_probability(self.pool_encoder.contextualize(local_h)) + + def _close_score( + self, probability: float, patch_len: int, minimum: Optional[int] = None + ) -> float: + """Positive-only midpoint incentive; it never lowers a close score.""" + minimum = self.min_patch_bytes if minimum is None else int(minimum) + span = max(1, self.max_patch_bytes - minimum) + progress = max(0.0, min(1.0, (float(patch_len) - minimum) / span)) + midpoint = max(0.0, 1.0 - abs(2.0 * progress - 1.0)) + return float(probability) + self.mid_close_bonus * midpoint + + def _runtime_patch_ids( + self, + contextual_h: torch.Tensor, + token_mask: torch.Tensor, + close_probabilities: Optional[torch.Tensor] = None, + minimum: Optional[int] = None, + minima_by_token: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """Causally route tokens into learned pools, closing the final pool too. + + Boundary decisions are only allowed after ``min_patch_bytes`` and are + always forced at ``max_patch_bytes``. Padding never opens a pool. + """ + minimum = self.min_patch_bytes if minimum is None else max( + 1, min(int(minimum), self.max_patch_bytes) + ) + probabilities = ( + self.pool_encoder.close_probability(contextual_h) + if close_probabilities is None + else close_probabilities + ).detach() + positions = torch.arange(contextual_h.shape[1], device=contextual_h.device).view(1, -1) + candidate_close = token_mask & (positions >= minimum) & ( + probabilities + self.mid_close_bonus >= self.close_threshold + ) + if minima_by_token is None and not bool(candidate_close.any()): + # The initialized close bias intentionally takes this fast path: + # bounded, vectorized max-sized pools with no Python per-token + # loop or CUDA synchronization for long training windows. + self.last_short_pool_budget_metrics = { + "short_pool_count": 0.0, + "budget_activated_rows": 0.0, + "batch_rows": float(contextual_h.shape[0]), + } + return ( + (positions // self.max_patch_bytes) + .expand_as(token_mask) + .long() + .masked_fill(~token_mask, 0) + ) + + # Learned closes are sparse. Move the small routing decision to CPU + # once and scan ordinary Python values, avoiding one CUDA sync per byte. + probability_rows = probabilities.float().cpu().tolist() + mask_rows = token_mask.cpu().tolist() + minima_rows = ( + minima_by_token.detach().to(device="cpu", dtype=torch.long).tolist() + if minima_by_token is not None else None + ) + rows: List[torch.Tensor] = [] + short_pool_counts: List[int] = [] + budget_activated_rows = 0 + for row_index, (row_prob, row_mask) in enumerate(zip(probability_rows, mask_rows)): + patch_id = 0 + patch_len = 0 + short_pools = 0 + secondary_active = False + recent_short_pools = deque() + recent_short_count = 0 + budget_was_activated = False + ids: List[int] = [] + previous_policy = None + row_minima = minima_rows[row_index] if minima_rows is not None else None + for token_index, (prob, valid) in enumerate(zip(row_prob, row_mask)): + if not valid: + ids.append(patch_id) + continue + policy = ( + tuple(int(value) for value in row_minima[token_index]) + if row_minima is not None else (minimum, 0, 0) + ) + policy_minimum = max(1, min(int(policy[0]), self.max_patch_bytes)) + runtime_minimum = max( + policy_minimum, + self.secondary_min_patch_bytes if secondary_active else 1, + ) + policy_transition = ( + patch_len > 0 and previous_policy is not None and policy != previous_policy + ) + learned_close = patch_len >= self.max_patch_bytes or ( + patch_len >= runtime_minimum + and self._close_score(prob, patch_len, policy_minimum) >= self.close_threshold + ) + if policy_transition or learned_close: + # A text<->codec transition is a structural boundary, not + # evidence that the learned router emitted a short pool. + completed_short = ( + not policy_transition and patch_len < self.secondary_min_patch_bytes + ) + if completed_short: + short_pools += 1 + if self.short_pool_budget > 0 and self.short_pool_window > 0: + recent_short_pools.append(completed_short) + recent_short_count += int(completed_short) + if len(recent_short_pools) > self.short_pool_window: + recent_short_count -= int(recent_short_pools.popleft()) + secondary_active = recent_short_count >= self.short_pool_budget + elif self.short_pool_budget > 0 and short_pools >= self.short_pool_budget: + secondary_active = True + budget_was_activated = budget_was_activated or secondary_active + patch_id += 1 + patch_len = 0 + ids.append(patch_id) + patch_len += 1 + previous_policy = policy + rows.append(torch.tensor(ids, dtype=torch.long, device=contextual_h.device)) + short_pool_counts.append(short_pools) + budget_activated_rows += int(budget_was_activated) + self.last_short_pool_budget_metrics = { + "short_pool_count": float(sum(short_pool_counts)), + "budget_activated_rows": float(budget_activated_rows), + "batch_rows": float(len(rows)), + } + return torch.stack(rows, dim=0).masked_fill(~token_mask, 0) + + def _runtime_hierarchy_ids( + self, + contextual_h: torch.Tensor, + token_mask: torch.Tensor, + encoder: LearnedPoolEncoder, + minimum: int, + maximum: int, + threshold: float, + minima_by_token: Optional[torch.Tensor] = None, + policy_by_token: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """Causally group completed lower-level states with learned closures. + + A lower-level state first joins its current group. Its close score may + then close that completed group for the *next* lower-level state. This + is important: closing the previous group from the current state would + let full-window inference expose a coarse result before cached + inference had actually completed the deciding state. + """ + positions = torch.arange(contextual_h.shape[1], device=contextual_h.device).view(1, -1) + if minimum >= maximum and minima_by_token is None and policy_by_token is None: + return ( + (positions // maximum) + .expand_as(token_mask) + .long() + .masked_fill(~token_mask, 0) + ) + probabilities = encoder.close_probability(contextual_h).detach() + probability_rows = probabilities.float().cpu().tolist() + mask_rows = token_mask.cpu().tolist() + minima_rows = ( + minima_by_token.detach().to(device="cpu", dtype=torch.long).tolist() + if minima_by_token is not None else None + ) + policy_rows = ( + policy_by_token.detach().to(device="cpu", dtype=torch.long).tolist() + if policy_by_token is not None else None + ) + rows: List[torch.Tensor] = [] + span = max(1, maximum - minimum) + for row_index, (row_probabilities, row_mask) in enumerate(zip(probability_rows, mask_rows)): + pool_id = 0 + pool_len = 0 + ids: List[int] = [] + previous_policy = None + for token_index, (probability, valid) in enumerate(zip(row_probabilities, row_mask)): + if not valid: + ids.append(pool_id) + continue + runtime_minimum = ( + max(1, min(int(minima_rows[row_index][token_index]), maximum)) + if minima_rows is not None else minimum + ) + policy = ( + tuple(int(value) for value in policy_rows[row_index][token_index]) + if policy_rows is not None else (runtime_minimum,) + ) + if pool_len > 0 and previous_policy is not None and policy != previous_policy: + pool_id += 1 + pool_len = 0 + ids.append(pool_id) + pool_len += 1 + runtime_span = max(1, maximum - runtime_minimum) + progress = max(0.0, min(1.0, (pool_len - runtime_minimum) / runtime_span)) + midpoint = max(0.0, 1.0 - abs(2.0 * progress - 1.0)) + close_score = float(probability) + self.mid_close_bonus * midpoint + if pool_len >= maximum or ( + pool_len >= runtime_minimum and close_score >= threshold + ): + pool_id += 1 + pool_len = 0 + previous_policy = policy + rows.append(torch.tensor(ids, dtype=torch.long, device=contextual_h.device)) + # scatter_add validates every index, including entries whose weight is + # zero. A group that closes on a row's final valid state increments + # pool_id for the following padding positions; force those invalid + # positions to the always-valid zero bucket before learned pooling. + return torch.stack(rows, dim=0).masked_fill(~token_mask, 0) + + @staticmethod + def _pool_runtime_minima( + minima: torch.Tensor, + ids: torch.Tensor, + mask: torch.Tensor, + group_count: int, + ) -> torch.Tensor: + """Carry a causal span policy into its forced-aligned pooled states.""" + output = torch.zeros( + (minima.shape[0], int(group_count), 3), + dtype=torch.long, + device=minima.device, + ) + safe_ids = ids.masked_fill(~mask, 0).unsqueeze(-1).expand(-1, -1, 3) + output.scatter_reduce_( + 1, + safe_ids, + minima.masked_fill(~mask.unsqueeze(-1), 0), + reduce="amax", + include_self=True, + ) + return output + + def _pool_patches( + self, + local_h: torch.Tensor, + patch_ids: torch.Tensor, + token_mask: torch.Tensor, + patch_counts: Optional[torch.Tensor] = None, + max_patches: Optional[int] = None, + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + b, t, c = local_h.shape + if patch_counts is None: + patch_counts = patch_ids.amax(dim=1) + 1 + else: + patch_counts = patch_counts.to(device=local_h.device, dtype=torch.long).reshape(-1) + if int(patch_counts.numel()) != b: + raise ValueError(f"patch_counts has {patch_counts.numel()} values for batch size {b}") + if max_patches is None: + max_patches = int(patch_counts.max().detach().cpu().item()) + else: + max_patches = int(max_patches) + if max_patches < 1: + raise ValueError("max_patches must be positive") + pooled = self.pool_encoder.pool(local_h, patch_ids, patch_counts, token_mask, max_patches) + counts = local_h.new_zeros((b, max_patches, 1)) + counts.scatter_add_(1, patch_ids.unsqueeze(-1), token_mask.unsqueeze(-1).to(dtype=local_h.dtype)) + patch_mask = torch.arange(max_patches, device=local_h.device).view(1, -1) < patch_counts.view(-1, 1) + patch_lens = counts.squeeze(-1).round().long().clamp(0, self.max_patch_bytes) + return pooled, patch_lens, patch_mask + + @staticmethod + def _first_valid_token_mask(token_mask: torch.Tensor) -> torch.Tensor: + """Select exactly the first non-padding token in each batch row.""" + return token_mask & token_mask.long().cumsum(dim=1).eq(1) + + def _gate_inactive_coarse_gradient( + self, + latent: torch.Tensor, + active_mask: torch.Tensor, + token_mask: torch.Tensor, + ) -> torch.Tensor: + """Keep forward values while limiting inactive coarse BOE gradients. + + The bootstrap state remains trainable on the first valid token. After + that, an inactive state is detached until a completed causal pool is + available. ``torch.where`` changes only autograd routing; inference + values and checkpoint shapes are unchanged. + """ + if not self.detach_inactive_coarse_gradients: + return latent + allow_gradient = active_mask | self._first_valid_token_mask(token_mask) + return torch.where(allow_gradient.unsqueeze(-1), latent, latent.detach()) + + def _global_stream_step(self, patch_embed: torch.Tensor, inference_params, offset: int) -> torch.Tensor: + _set_inference_seqlen_offset(inference_params, int(offset)) + if patch_embed.device.type == "cuda": + # Mamba's inference conv/SSM cache is dtype-strict. Keep recurrent + # state updates in fp32 even when streaming validation runs under AMP. + with torch.amp.autocast("cuda", enabled=False): + stream_dtype = next(self.global_blocks[0].parameters()).dtype + h = patch_embed.to(dtype=stream_dtype) + for block in self.global_blocks: + h = self._run_block_on_own_device( + block, h, inference_params=inference_params + ) + else: + h = patch_embed + for block in self.global_blocks: + h = self._run_block_on_own_device( + block, h, inference_params=inference_params + ) + return h + + def _nested_global_stream_step(self, pool_embed: torch.Tensor, inference_params, offset: int) -> torch.Tensor: + """Advance the coarse Mamba cache by one completed coarse pool.""" + _set_inference_seqlen_offset(inference_params, int(offset)) + if pool_embed.device.type == "cuda": + with torch.amp.autocast("cuda", enabled=False): + stream_dtype = next(self.nested_global_blocks[0].parameters()).dtype + h = pool_embed.to(dtype=stream_dtype) + for block in self.nested_global_blocks: + block_device = self._module_device(block) + if h.device != block_device: + h = h.to(block_device, non_blocking=True) + h = self._run_block_on_own_device( + block, h, inference_params=inference_params + ) + else: + h = pool_embed + for block in self.nested_global_blocks: + h = self._run_block_on_own_device( + block, h, inference_params=inference_params + ) + return h + + def _tertiary_global_stream_step(self, pool_embed: torch.Tensor, inference_params, offset: int) -> torch.Tensor: + """Advance level 3 by one completed level-2 group.""" + _set_inference_seqlen_offset(inference_params, int(offset)) + h = pool_embed + if h.device.type == "cuda": + with torch.amp.autocast("cuda", enabled=False): + stream_dtype = next(self.tertiary_global_blocks[0].parameters()).dtype + h = h.to(dtype=stream_dtype) + for block in self.tertiary_global_blocks: + block_device = self._module_device(block) + if h.device != block_device: + h = h.to(block_device, non_blocking=True) + h = self._run_block_on_own_device( + block, h, inference_params=inference_params + ) + else: + for block in self.tertiary_global_blocks: + h = self._run_block_on_own_device( + block, h, inference_params=inference_params + ) + return h + + def _initialize_stream_mamba_caches( + self, + blocks: nn.ModuleList, + inference_params, + max_seqlen: int, + ) -> None: + """Preallocate numerically stable Mamba-2 recurrent states. + + Casting a checkpoint to FP16 should not also reduce the accumulator + range of its long-lived SSD state. Mamba-2's public cache allocator + uses one dtype for both caches, so retain the short convolution cache + in the model dtype and promote only the SSM accumulator to FP32. + The Triton selective-state-update kernel supports this mixed layout. + """ + if inference_params is None: + return + for block in blocks: + if ( + getattr(block, "kind", None) != "mamba" + or int(getattr(block, "mamba_version", 1)) != 2 + ): + continue + mixer = block.mixer + layer_idx = getattr(mixer, "layer_idx", None) + if layer_idx is None: + raise RuntimeError("cached Mamba-2 inference requires a unique layer_idx") + block_device = self._module_device(block) + if block_device.type == "cuda": + with torch.cuda.device(block_device): + conv_state, ssm_state = mixer.allocate_inference_cache( + 1, int(max_seqlen) + ) + else: + conv_state, ssm_state = mixer.allocate_inference_cache( + 1, int(max_seqlen) + ) + inference_params.key_value_memory_dict[int(layer_idx)] = ( + conv_state, + ssm_state.float(), + ) + + def _patch_conditioning( + self, + patch_h: torch.Tensor, + patch_len: int, + patch_index: int, + ) -> torch.Tensor: + b, _, _ = patch_h.shape + device = patch_h.device + patch_pos = torch.full((b, 1), int(patch_index), dtype=torch.long, device=device) + patch_lens = torch.full((b, 1), max(1, min(int(patch_len), self.max_patch_bytes)), dtype=torch.long, device=device) + return patch_h + self.patch_pos_emb(patch_pos.remainder(self.position_bins)) + self.patch_len_emb(patch_lens) + + def _concat_patch_parts(self, parts: List[Dict[str, torch.Tensor]]) -> Dict[str, torch.Tensor]: + keys = [k for k, v in parts[0].items() if isinstance(v, torch.Tensor) and v.dim() >= 2 and int(v.shape[1]) == 1] + return {k: torch.cat([p[k] for p in parts], dim=1) for k in keys} + + def _hierarchy_should_close( + self, + probability: float, + pool_len: int, + minimum: int, + maximum: int, + threshold: float, + controller_adjustment: float = 0.0, + ) -> bool: + if pool_len >= maximum: + return True + if pool_len < minimum: + return False + span = max(1, maximum - minimum) + progress = max(0.0, min(1.0, (pool_len - minimum) / span)) + midpoint = max(0.0, 1.0 - abs(2.0 * progress - 1.0)) + return ( + float(probability) + + self.mid_close_bonus * midpoint + + float(controller_adjustment) + >= threshold + ) + + def _append_nested_stream_part( + self, state: Dict[str, object], fine_completed_h: torch.Tensor + ) -> None: + """Route one completed fine state into the adaptive level-2 pool.""" + nested_device = self._module_device(self.nested_pool_encoder) + current = fine_completed_h + if current.device != nested_device: + current = current.to(nested_device, non_blocking=True) + routing_tail = state.get("nested_routing_tail") + routing_input = current if routing_tail is None else torch.cat([routing_tail, current], dim=1) + contextual_current = self.nested_pool_encoder.contextualize(routing_input)[:, -1:, :] + probability = float( + self.nested_pool_encoder.close_probability(contextual_current)[0, 0] + .detach() + .cpu() + .item() + ) + state.setdefault("nested_fine_parts", []).append(current) + state["nested_routing_tail"] = routing_input[:, -2:].detach() + state["pending_nested_close_probability"] = probability + controller_adjustment = self._controller_close_adjustment( + state.get("pending_pool_controller_scores"), 1 + ) + if self._hierarchy_should_close( + probability, + len(state["nested_fine_parts"]), + self.nested_min_pool_factor, + self.nested_pool_factor, + self.nested_close_threshold, + controller_adjustment, + ): + self._close_nested_stream_pool(state, force=True) + + def _close_nested_stream_pool(self, state: Dict[str, object], force: bool = False) -> None: + """Commit a group of completed fine pools to the coarse cached stream.""" + fine_parts = state.get("nested_fine_parts", []) + if not fine_parts or (not force and len(fine_parts) < self.nested_pool_factor): + return + fine_h = torch.cat(fine_parts, dim=1) + nested_device = self._module_device(self.nested_pool_encoder) + if fine_h.device != nested_device: + fine_h = fine_h.to(nested_device, non_blocking=True) + tail = state.get("nested_context_tail") + contextual_input = fine_h if tail is None else torch.cat([tail, fine_h], dim=1) + contextual_h = self.nested_pool_encoder.contextualize(contextual_input)[:, -fine_h.shape[1]:] + token_mask = torch.ones(fine_h.shape[:2], dtype=torch.bool, device=fine_h.device) + nested_ids = torch.zeros(fine_h.shape[:2], dtype=torch.long, device=fine_h.device) + nested_h = self.nested_pool_encoder.pool( + contextual_h, nested_ids, torch.ones(1, dtype=torch.long, device=fine_h.device), token_mask, 1 + ) + nested_index = int(state["completed_nested_patches"]) + position_device = self.patch_pos_emb.weight.device + nested_pos = torch.full((1, 1), nested_index, dtype=torch.long, device=position_device) + nested_len = torch.full((1, 1), int(fine_h.shape[1]), dtype=torch.long, device=fine_h.device) + nested_pos_h = self.patch_pos_emb(nested_pos.remainder(self.position_bins)) + if nested_pos_h.device != nested_device: + nested_pos_h = nested_pos_h.to(nested_device, non_blocking=True) + nested_h = nested_h + nested_pos_h + self.nested_len_emb(nested_len) + state["current_nested_global"] = self._nested_global_stream_step( + nested_h, state["nested_inference_params"], nested_index + 1 + ) + state["completed_nested_patches"] = nested_index + 1 + state["nested_context_tail"] = contextual_input[:, -2:].detach() + state.setdefault("closed_nested_pools", []).append({ + "length": int(fine_h.shape[1]), + "close_probability": float( + state.get("pending_nested_close_probability", 0.0) + ), + }) + state["nested_fine_parts"] = [] + if self.tertiary_enabled: + self._append_tertiary_stream_part(state, state["current_nested_global"]) + + def _append_tertiary_stream_part( + self, state: Dict[str, object], nested_completed_h: torch.Tensor + ) -> None: + """Route one completed level-2 state into the adaptive level-3 pool.""" + if not self.tertiary_enabled: + return + assert self.tertiary_pool_encoder is not None + tertiary_device = self._module_device(self.tertiary_pool_encoder) + current = nested_completed_h + if current.device != tertiary_device: + current = current.to(tertiary_device, non_blocking=True) + routing_tail = state.get("tertiary_routing_tail") + routing_input = current if routing_tail is None else torch.cat([routing_tail, current], dim=1) + contextual_current = self.tertiary_pool_encoder.contextualize(routing_input)[:, -1:, :] + probability = float( + self.tertiary_pool_encoder.close_probability(contextual_current)[0, 0] + .detach() + .cpu() + .item() + ) + state.setdefault("tertiary_nested_parts", []).append(current) + state["tertiary_routing_tail"] = routing_input[:, -2:].detach() + state["pending_tertiary_close_probability"] = probability + controller_adjustment = self._controller_close_adjustment( + state.get("pending_pool_controller_scores"), 2 + ) + if self._hierarchy_should_close( + probability, + len(state["tertiary_nested_parts"]), + self.tertiary_min_pool_factor, + self.tertiary_pool_factor, + self.tertiary_close_threshold, + controller_adjustment, + ): + self._close_tertiary_stream_pool(state, force=True) + + def _close_tertiary_stream_pool(self, state: Dict[str, object], force: bool = False) -> None: + """Commit completed level-2 states to the cached level-3 stream.""" + if not self.tertiary_enabled: + return + nested_parts = state.get("tertiary_nested_parts", []) + if not nested_parts or (not force and len(nested_parts) < self.tertiary_pool_factor): + return + assert self.tertiary_pool_encoder is not None and self.tertiary_len_emb is not None + tertiary_device = self._module_device(self.tertiary_pool_encoder) + nested_h = torch.cat(nested_parts, dim=1) + if nested_h.device != tertiary_device: + nested_h = nested_h.to(tertiary_device, non_blocking=True) + tail = state.get("tertiary_context_tail") + contextual_input = nested_h if tail is None else torch.cat([tail, nested_h], dim=1) + contextual_h = self.tertiary_pool_encoder.contextualize(contextual_input)[:, -nested_h.shape[1]:] + token_mask = torch.ones(nested_h.shape[:2], dtype=torch.bool, device=tertiary_device) + tertiary_ids = torch.zeros(nested_h.shape[:2], dtype=torch.long, device=tertiary_device) + tertiary_h = self.tertiary_pool_encoder.pool( + contextual_h, + tertiary_ids, + torch.ones(1, dtype=torch.long, device=tertiary_device), + token_mask, + 1, + ) + tertiary_index = int(state["completed_tertiary_patches"]) + position_device = self.patch_pos_emb.weight.device + tertiary_pos = torch.full((1, 1), tertiary_index, dtype=torch.long, device=position_device) + tertiary_pos_h = self.patch_pos_emb(tertiary_pos.remainder(self.position_bins)) + if tertiary_pos_h.device != tertiary_device: + tertiary_pos_h = tertiary_pos_h.to(tertiary_device, non_blocking=True) + tertiary_len = torch.full( + (1, 1), int(nested_h.shape[1]), dtype=torch.long, device=tertiary_device + ) + tertiary_h = tertiary_h + tertiary_pos_h + self.tertiary_len_emb(tertiary_len) + state["current_tertiary_global"] = self._tertiary_global_stream_step( + tertiary_h, state["tertiary_inference_params"], tertiary_index + 1 + ) + state["completed_tertiary_patches"] = tertiary_index + 1 + state["tertiary_context_tail"] = contextual_input[:, -2:].detach() + state.setdefault("closed_tertiary_pools", []).append({ + "length": int(nested_h.shape[1]), + "close_probability": float( + state.get("pending_tertiary_close_probability", 0.0) + ), + }) + state["tertiary_nested_parts"] = [] + + def _close_stream_patch(self, state: Dict[str, object]) -> None: + patch_tokens = state.get("patch_tokens", []) + if not patch_tokens: + return + contextual_h = torch.cat(state["patch_contextual_parts"], dim=1) + patch_ids = torch.zeros(contextual_h.shape[:2], dtype=torch.long, device=contextual_h.device) + token_mask = torch.ones(contextual_h.shape[:2], dtype=torch.bool, device=contextual_h.device) + patch_h = self.pool_encoder.pool( + contextual_h, + patch_ids, + torch.ones(1, dtype=torch.long, device=contextual_h.device), + token_mask, + 1, + ) + patch_index = int(state["completed_patches"]) + patch_h = self._patch_conditioning( + patch_h, + int(contextual_h.shape[1]), + patch_index, + ) + fine_completed_h = self._global_stream_step( + patch_h, + state["inference_params"], + int(state["completed_patches"]) + 1, + ) + state["current_global"] = fine_completed_h + state["completed_patches"] = int(state["completed_patches"]) + 1 + self._append_nested_stream_part(state, fine_completed_h) + state.setdefault("closed_pools", []).append({ + "length": int(contextual_h.shape[1]), + "close_probability": float(state.get("pending_close_probability", 0.0)), + "reason": str(state.get("pending_close_reason", "end")), + "byte_pos": int(state.get("patch_end_byte_pos", 0)), + }) + if int(contextual_h.shape[1]) < self.secondary_min_patch_bytes: + state["short_pool_count"] = int(state.get("short_pool_count", 0)) + 1 + if self.short_pool_budget > 0 and self.short_pool_window > 0: + history = state.setdefault("short_pool_history", deque()) + completed_short = int(contextual_h.shape[1]) < self.secondary_min_patch_bytes + history.append(completed_short) + state["short_pool_window_count"] = int( + state.get("short_pool_window_count", 0) + ) + int(completed_short) + if len(history) > self.short_pool_window: + state["short_pool_window_count"] -= int(history.popleft()) + state["secondary_min_active"] = ( + int(state["short_pool_window_count"]) >= self.short_pool_budget + ) + elif ( + self.short_pool_budget > 0 + and int(state.get("short_pool_count", 0)) >= self.short_pool_budget + ): + state["secondary_min_active"] = True + state["patch_tokens"] = [] + state["patch_contextual_parts"] = [] + state["patch_prev_raw"] = None + state["pending_close_probability"] = 0.0 + state["pending_close_reason"] = "end" + + def new_stream_state(self, batch: Dict[str, torch.Tensor], max_patches: int) -> Dict[str, object]: + inference_params = _new_mamba_inference_params(max_batch_size=1, max_seqlen=max(2, int(max_patches) + 1)) + # Adaptive groups can close at their minimum, so cache capacity must be + # based on the shortest possible group rather than the hard maximum. + nested_max_patches = max( + 2, int(math.ceil(int(max_patches) / self.nested_min_pool_factor)) + 1 + ) + nested_inference_params = _new_mamba_inference_params(max_batch_size=1, max_seqlen=nested_max_patches) + tertiary_max_patches = max( + 2, + int( + math.ceil( + nested_max_patches / max(1, self.tertiary_min_pool_factor) + ) + ) + + 1, + ) + tertiary_inference_params = ( + _new_mamba_inference_params(max_batch_size=1, max_seqlen=tertiary_max_patches) + if self.tertiary_enabled + else None + ) + if inference_params is None or nested_inference_params is None or ( + self.tertiary_enabled and tertiary_inference_params is None + ): + raise RuntimeError("BLT streaming evaluation requires mamba_ssm InferenceParams.") + self._initialize_stream_mamba_caches( + self.global_blocks, inference_params, max(2, int(max_patches) + 1) + ) + self._initialize_stream_mamba_caches( + self.nested_global_blocks, + nested_inference_params, + nested_max_patches, + ) + if self.tertiary_enabled: + self._initialize_stream_mamba_caches( + self.tertiary_global_blocks, + tertiary_inference_params, + tertiary_max_patches, + ) + boe = self.boe_patch.expand(1, 1, self.boe_patch.shape[-1]) + current_global = self._global_stream_step(boe, inference_params, 0) + nested_device = self._module_device(self.nested_global_blocks[0]) + current_nested_global = self._nested_global_stream_step( + boe.to(nested_device, non_blocking=True), nested_inference_params, 0 + ) + current_tertiary_global = None + if self.tertiary_enabled: + tertiary_device = self._module_device(self.tertiary_global_blocks[0]) + current_tertiary_global = self._tertiary_global_stream_step( + boe.to(tertiary_device, non_blocking=True), tertiary_inference_params, 0 + ) + return { + "inference_params": inference_params, + "nested_inference_params": nested_inference_params, + "tertiary_inference_params": tertiary_inference_params, + "current_global": current_global, + "current_nested_global": current_nested_global, + "current_tertiary_global": current_tertiary_global, + "initial_nested_global": current_nested_global.detach().clone(), + "initial_tertiary_global": ( + current_tertiary_global.detach().clone() + if current_tertiary_global is not None + else None + ), + "prev_byte_latent": current_global, + "completed_patches": 0, + "completed_nested_patches": 0, + "completed_tertiary_patches": 0, + "patch_tokens": [], + "patch_contextual_parts": [], + "patch_prev_raw": None, + "pending_close_probability": 0.0, + "pending_close_reason": "end", + "closed_pools": [], + "nested_fine_parts": [], + "closed_nested_pools": [], + "nested_context_tail": None, + "nested_routing_tail": None, + "pending_nested_close_probability": 0.0, + "tertiary_nested_parts": [], + "closed_tertiary_pools": [], + "tertiary_context_tail": None, + "tertiary_routing_tail": None, + "pending_tertiary_close_probability": 0.0, + "pending_pool_controller_scores": None, + "short_pool_count": 0, + "short_pool_history": deque(), + "short_pool_window_count": 0, + "secondary_min_active": False, + "local_conv1_tail": None, + "local_conv2_tail": None, + "pool_context_tail": None, + } + + def stream_step(self, state: Dict[str, object], one: Dict[str, torch.Tensor]) -> torch.Tensor: + # A controller emitted at t is consumed only when step t+1 begins. + # Closing here makes any newly completed fine/L2/L3 state visible to + # the next decode, while never allowing that decode to alter its own + # routing decision. + token = one["x"] + contextual_current = self._local_stream_context(token, state) + close_prob = float( + self.pool_encoder.close_probability(contextual_current)[0, 0].detach().cpu().item() + ) + patch_len = len(state.get("patch_tokens", [])) + runtime_minimum = ( + self.secondary_min_patch_bytes + if bool(state.get("secondary_min_active", False)) + else self.min_patch_bytes + ) + if patch_len: + controller_adjustment = self._controller_close_adjustment( + state.get("pending_pool_controller_scores"), 0 + ) + close_reason = None + if patch_len >= self.max_patch_bytes: + close_reason = "max" + elif ( + patch_len >= runtime_minimum + and self._close_score(close_prob, patch_len) + + controller_adjustment + >= self.close_threshold + ): + close_reason = ( + "decoder_controller" + if self.decoder_pool_controller + else "learned" + ) + if close_reason is not None: + state["pending_close_probability"] = close_prob + state["pending_close_reason"] = close_reason + self._close_stream_patch(state) + + state["patch_tokens"].append(token) + state["patch_contextual_parts"].append(contextual_current) + + local_last = contextual_current + byte_latent = state["current_global"] + prev_latent = state.get("prev_byte_latent", byte_latent) + nested_latent = state["current_nested_global"] + decoder_device = local_last.device + if nested_latent.device != decoder_device: + nested_latent = nested_latent.to(decoder_device, non_blocking=True) + decode_parts = [local_last, byte_latent, prev_latent, nested_latent] + if self.tertiary_enabled: + tertiary_latent = state["current_tertiary_global"] + if tertiary_latent.device != decoder_device: + tertiary_latent = tertiary_latent.to(decoder_device, non_blocking=True) + decode_parts.append(tertiary_latent) + self.last_stream_decode_parts = tuple(part.detach() for part in decode_parts) + decoded = self.decoder(torch.cat(decode_parts, dim=-1)) + logits = self.lm_head(decoded) + controller_scores = self._pool_controller_scores(decoded) + if controller_scores is not None: + current_scores = controller_scores[0, 0].detach().float().cpu().tolist() + state["pending_pool_controller_scores"] = current_scores + else: + state["pending_pool_controller_scores"] = None + state["prev_byte_latent"] = byte_latent + return logits + + def finish_stream(self, state: Dict[str, object]) -> None: + """Flush an under-length pool at an explicit end-of-stream boundary.""" + self._close_stream_patch(state) + self._close_nested_stream_pool(state, force=True) + self._close_tertiary_stream_pool(state, force=True) + + def forward( + self, + x, + patch_ids=None, + patch_counts=None, + max_patches=None, + inference_params=None, + runtime_pool_minima=None, + **_unused, + ): + if self.model_parallel_enabled and inference_params is not None: + raise RuntimeError("stateful nested streaming is not supported with --nested-devices; use full-window training/validation or a single device.") + fine_minimum = self.min_patch_bytes + nested_minimum = self.nested_min_pool_factor + tertiary_minimum = self.tertiary_min_pool_factor + per_token_minima = None + if runtime_pool_minima is not None: + if runtime_pool_minima.numel() == 3: + minima = runtime_pool_minima.detach().to(device="cpu").reshape(-1).tolist() + fine_minimum = max(1, min(int(minima[0]), self.max_patch_bytes)) + nested_minimum = max(1, min(int(minima[1]), self.nested_pool_factor)) + tertiary_minimum = max(1, min(int(minima[2]), self.tertiary_pool_factor)) + else: + if runtime_pool_minima.ndim != 3 or runtime_pool_minima.shape[-1] != 3: + raise ValueError( + "runtime_pool_minima must be [3] or [batch, bytes, 3]" + ) + if tuple(runtime_pool_minima.shape[:2]) != tuple(x.shape): + raise ValueError( + "per-byte runtime_pool_minima must match x batch and length" + ) + per_token_minima = runtime_pool_minima.to( + device=x.device, dtype=torch.long, non_blocking=True + ).clone() + defaults = per_token_minima.new_tensor([ + fine_minimum, nested_minimum, tertiary_minimum, + ]).view(1, 1, 3) + per_token_minima = torch.where( + per_token_minima.gt(0), per_token_minima, defaults + ) + per_token_minima[..., 0].clamp_(1, self.max_patch_bytes) + per_token_minima[..., 1].clamp_(1, self.nested_pool_factor) + per_token_minima[..., 2].clamp_(1, self.tertiary_pool_factor) + local_h = self._local_features(x) + contextual_h = self.pool_encoder.contextualize(local_h) + # Pool boundaries are model predictions. Ignore legacy precomputed + # greedy IDs so cached batches/checkpoints remain usable during the + # transition to learned pooling. + token_mask = x.ne(0) + close_logits = self.pool_encoder.close_logits(contextual_h) + close_probabilities = torch.sigmoid(close_logits) + self.last_pool_close_probabilities = close_probabilities + patch_ids = self._runtime_patch_ids( + contextual_h, token_mask, close_probabilities, minimum=fine_minimum, + minima_by_token=per_token_minima, + ) + patch_counts = patch_ids.masked_fill(~token_mask, 0).amax(dim=1) + 1 + max_patches = int(patch_counts.max().detach().cpu().item()) + patch_h, patch_lens, patch_mask = self._pool_patches( + contextual_h, + patch_ids, + token_mask, + patch_counts=patch_counts, + max_patches=max_patches, + ) + patch_minima = ( + self._pool_runtime_minima( + per_token_minima, patch_ids, token_mask, max_patches + ) + if per_token_minima is not None else None + ) + valid_pool_lens = patch_lens[patch_mask] + self.last_pool_metrics = { + "pool_length_sum": valid_pool_lens.detach().float().sum(), + "pool_count": valid_pool_lens.detach().new_tensor(float(valid_pool_lens.numel())), + "pool_length_mean": valid_pool_lens.detach().float().mean() if valid_pool_lens.numel() else patch_h.new_zeros(()), + "pools_per_window": patch_counts.detach().float().mean(), + } + budget_metrics = getattr(self, "last_short_pool_budget_metrics", {}) + self.last_pool_metrics.update({ + "short_pool_count": patch_h.new_tensor( + float(budget_metrics.get("short_pool_count", 0.0)) + ), + "short_pool_budget_activated_rows": patch_h.new_tensor( + float(budget_metrics.get("budget_activated_rows", 0.0)) + ), + "short_pool_budget_batch_rows": patch_h.new_tensor( + float(budget_metrics.get("batch_rows", contextual_h.shape[0])) + ), + }) + b, p, c = patch_h.shape + patch_pos = torch.arange(p, device=x.device, dtype=torch.long).view(1, p).expand(b, p) + patch_h = patch_h + self.patch_pos_emb(patch_pos.remainder(self.position_bins)) + self.patch_len_emb(patch_lens) + + patch_h = patch_h.masked_fill(~patch_mask.unsqueeze(-1), 0.0) + boe_patch = self.boe_patch.expand(b, 1, c) + global_h = torch.cat([boe_patch, patch_h], dim=1) + global_mask = torch.cat( + [torch.ones((b, 1), dtype=torch.bool, device=x.device), patch_mask], + dim=1, + ) + for block in self.global_blocks: + global_h = self._run_block_on_own_device( + block, global_h, inference_params=inference_params + ) + global_h = global_h.masked_fill(~global_mask.unsqueeze(-1), 0.0) + + # Causal BLT alignment: bytes in patch p decode from global state p, + # which is BOE for patch 0 and completed patch p-1 thereafter. + byte_latents = global_h.gather(1, patch_ids.unsqueeze(-1).expand(-1, -1, c)) + prev_latents = torch.cat([byte_latents[:, :1], byte_latents[:, :-1]], dim=1) + + # Nest completed fine-pool states into a coarser causal stream. Each + # coarse state is only read by a *later* group of fine pools, so the + # multi-scale decoder never receives information from its current + # target pool. This lets later layers operate on roughly + # ``nested_pool_factor`` fewer positions without losing fine readouts. + fine_h = global_h[:, 1:] + # The nested hierarchy is substantially shorter than the byte path. + # In model-parallel mode, make that single fine->coarse transfer here. + nested_device = self._module_device(self.nested_pool_encoder) + fine_h_nested = fine_h if fine_h.device == nested_device else fine_h.to(nested_device, non_blocking=True) + patch_mask_nested = patch_mask if patch_mask.device == nested_device else patch_mask.to(nested_device, non_blocking=True) + patch_counts_nested = patch_counts if patch_counts.device == nested_device else patch_counts.to(nested_device, non_blocking=True) + # All IDs used by coarse pooling must live with the coarse activations. + fine_contextual_h = self.nested_pool_encoder.contextualize(fine_h_nested) + nested_close_logits = self.nested_pool_encoder.close_logits(fine_contextual_h) + nested_close_probabilities = torch.sigmoid(nested_close_logits) + nested_ids = self._runtime_hierarchy_ids( + fine_contextual_h, + patch_mask_nested, + self.nested_pool_encoder, + nested_minimum, + self.nested_pool_factor, + self.nested_close_threshold, + minima_by_token=( + patch_minima[..., 1].to(nested_device, non_blocking=True) + if patch_minima is not None else None + ), + policy_by_token=( + patch_minima.to(nested_device, non_blocking=True) + if patch_minima is not None else None + ), + ) + nested_counts = nested_ids.masked_fill(~patch_mask_nested, 0).amax(dim=1) + 1 + max_nested_patches = int(nested_counts.max().detach().cpu().item()) + nested_h = self.nested_pool_encoder.pool( + fine_contextual_h, + nested_ids, + nested_counts, + patch_mask_nested, + max_nested_patches, + ) + nested_minima = ( + self._pool_runtime_minima( + patch_minima.to(nested_device, non_blocking=True), + nested_ids, + patch_mask_nested, + max_nested_patches, + ) + if patch_minima is not None else None + ) + nested_token_counts = fine_h_nested.new_zeros((b, max_nested_patches, 1)) + nested_token_counts.scatter_add_( + 1, + nested_ids.unsqueeze(-1), + patch_mask_nested.unsqueeze(-1).to(dtype=fine_h_nested.dtype), + ) + nested_lens = nested_token_counts.squeeze(-1).round().long().clamp(0, self.nested_pool_factor) + nested_mask = torch.arange(max_nested_patches, device=nested_device).view(1, -1) < nested_counts.view(-1, 1) + valid_nested_lens = nested_lens[nested_mask] + self.last_pool_metrics.update({ + "nested_pool_length_sum": valid_nested_lens.detach().float().sum(), + "nested_pool_count": valid_nested_lens.detach().new_tensor(float(valid_nested_lens.numel())), + "nested_pool_length_mean": ( + valid_nested_lens.detach().float().mean() + if valid_nested_lens.numel() + else nested_h.new_zeros(()) + ), + "nested_pools_per_window": nested_counts.detach().float().mean(), + }) + nested_pos = torch.arange(max_nested_patches, device=x.device, dtype=torch.long).view(1, -1).expand(b, -1) + # patch_pos_emb is intentionally shared with the fine stage, so look + # it up on the fine device and transfer its small coarse sequence. + nested_pos_emb = self.patch_pos_emb(nested_pos.remainder(self.position_bins)).to(nested_device, non_blocking=True) + nested_h = ( + nested_h + + nested_pos_emb + + self.nested_len_emb(nested_lens) + ).masked_fill(~nested_mask.unsqueeze(-1), 0.0) + nested_boe = boe_patch if boe_patch.device == nested_device else boe_patch.to(nested_device, non_blocking=True) + nested_global_h = torch.cat([nested_boe, nested_h], dim=1) + nested_global_mask = torch.cat( + [torch.ones((b, 1), dtype=torch.bool, device=nested_device), nested_mask], + dim=1, + ) + nested_global_h = self._run_nested_blocks(nested_global_h, nested_global_mask) + tertiary_byte_latents = None + tertiary_initial_byte_latents = None + tertiary_ids_fine = None + tertiary_lens = None + tertiary_mask = None + tertiary_close_probabilities = None + tertiary_close_logits = None + tertiary_ids = None + level2_mask = None + if self.tertiary_enabled: + assert self.tertiary_pool_encoder is not None and self.tertiary_len_emb is not None + tertiary_device = self._module_device(self.tertiary_pool_encoder) + level2_h = nested_global_h[:, 1:] + if level2_h.device != tertiary_device: + level2_h = level2_h.to(tertiary_device, non_blocking=True) + level2_mask = nested_mask if nested_mask.device == tertiary_device else nested_mask.to(tertiary_device, non_blocking=True) + level2_counts = nested_counts if nested_counts.device == tertiary_device else nested_counts.to(tertiary_device, non_blocking=True) + level2_contextual = self.tertiary_pool_encoder.contextualize(level2_h) + tertiary_close_logits = self.tertiary_pool_encoder.close_logits(level2_contextual) + tertiary_close_probabilities = torch.sigmoid(tertiary_close_logits) + tertiary_ids = self._runtime_hierarchy_ids( + level2_contextual, + level2_mask, + self.tertiary_pool_encoder, + tertiary_minimum, + self.tertiary_pool_factor, + self.tertiary_close_threshold, + minima_by_token=( + nested_minima[..., 2].to(tertiary_device, non_blocking=True) + if nested_minima is not None else None + ), + policy_by_token=( + nested_minima.to(tertiary_device, non_blocking=True) + if nested_minima is not None else None + ), + ) + tertiary_counts = tertiary_ids.masked_fill(~level2_mask, 0).amax(dim=1) + 1 + max_tertiary = int(tertiary_counts.max().detach().cpu().item()) + tertiary_h = self.tertiary_pool_encoder.pool(level2_contextual, tertiary_ids, tertiary_counts, level2_mask, max_tertiary) + tertiary_token_counts = level2_h.new_zeros((b, max_tertiary, 1)) + tertiary_token_counts.scatter_add_(1, tertiary_ids.unsqueeze(-1), level2_mask.unsqueeze(-1).to(level2_h.dtype)) + tertiary_lens = tertiary_token_counts.squeeze(-1).round().long().clamp(0, self.tertiary_pool_factor) + tertiary_mask = torch.arange(max_tertiary, device=tertiary_device).view(1, -1) < tertiary_counts.view(-1, 1) + valid_tertiary_lens = tertiary_lens[tertiary_mask] + self.last_pool_metrics.update({ + "tertiary_pool_length_sum": valid_tertiary_lens.detach().float().sum(), + "tertiary_pool_count": valid_tertiary_lens.detach().new_tensor(float(valid_tertiary_lens.numel())), + "tertiary_pool_length_mean": ( + valid_tertiary_lens.detach().float().mean() + if valid_tertiary_lens.numel() + else tertiary_h.new_zeros(()) + ), + "tertiary_pools_per_window": tertiary_counts.detach().float().mean(), + }) + tertiary_pos = torch.arange(max_tertiary, device=x.device, dtype=torch.long).view(1, -1).expand(b, -1) + tertiary_pos_emb = self.patch_pos_emb(tertiary_pos.remainder(self.position_bins)).to(tertiary_device, non_blocking=True) + tertiary_h = (tertiary_h + tertiary_pos_emb + self.tertiary_len_emb(tertiary_lens)).masked_fill(~tertiary_mask.unsqueeze(-1), 0.0) + tertiary_boe = boe_patch.to(tertiary_device, non_blocking=True) + tertiary_global_h = self._run_tertiary_blocks( + torch.cat([tertiary_boe, tertiary_h], dim=1), + torch.cat([torch.ones((b, 1), dtype=torch.bool, device=tertiary_device), tertiary_mask], dim=1), + ) + # Map level-3 causal readouts back to level 2, then fine patches. + tertiary_global_h = tertiary_global_h.to(x.device, non_blocking=True) + tertiary_initial_byte_latents = tertiary_global_h[:, :1].expand( + -1, x.shape[1], -1 + ) + tertiary_ids_fine = tertiary_ids.to(x.device, non_blocking=True) + level2_tertiary = tertiary_global_h.gather(1, tertiary_ids_fine.unsqueeze(-1).expand(-1, -1, c)) + nested_ids_fine_for_tertiary = nested_ids.to(x.device, non_blocking=True) + tertiary_byte_latents = level2_tertiary.gather(1, nested_ids_fine_for_tertiary.unsqueeze(-1).expand(-1, -1, c)).gather(1, patch_ids.unsqueeze(-1).expand(-1, -1, c)) + # Return only the coarse result; decoder and all byte-level gathers + # remain on the fine device. + if nested_global_h.device != x.device: + nested_global_h = nested_global_h.to(x.device, non_blocking=True) + nested_ids_fine = nested_ids if nested_ids.device == x.device else nested_ids.to(x.device, non_blocking=True) + nested_readouts = nested_global_h.gather(1, nested_ids_fine.unsqueeze(-1).expand(-1, -1, c)) + nested_byte_latents = nested_readouts.gather(1, patch_ids.unsqueeze(-1).expand(-1, -1, c)) + nested_byte_pool_ids = nested_ids_fine.gather(1, patch_ids) + nested_byte_latents = self._gate_inactive_coarse_gradient( + nested_byte_latents, + nested_byte_pool_ids.gt(0), + token_mask, + ) + nested_initial_byte_latents = nested_global_h[:, :1].expand(-1, x.shape[1], -1) + if tertiary_byte_latents is not None and tertiary_ids_fine is not None: + tertiary_fine_pool_ids = tertiary_ids_fine.gather(1, nested_ids_fine) + tertiary_byte_pool_ids = tertiary_fine_pool_ids.gather(1, patch_ids) + tertiary_byte_latents = self._gate_inactive_coarse_gradient( + tertiary_byte_latents, + tertiary_byte_pool_ids.gt(0), + token_mask, + ) + decode_parts = [contextual_h, byte_latents, prev_latents, nested_byte_latents] + if tertiary_byte_latents is not None: + decode_parts.append(tertiary_byte_latents) + if self.capture_evaluation_details: + self.last_forward_evaluation_details = { + "token_mask": token_mask.detach(), + "patch_ids": patch_ids.detach(), + "patch_lens": patch_lens.detach(), + "patch_mask": patch_mask.detach(), + "nested_ids": nested_ids_fine.detach(), + "nested_lens": nested_lens.detach(), + "nested_mask": nested_mask.detach(), + "tertiary_ids": ( + tertiary_ids_fine.detach() if tertiary_ids_fine is not None else None + ), + "tertiary_lens": ( + tertiary_lens.detach() if tertiary_lens is not None else None + ), + "tertiary_mask": ( + tertiary_mask.detach() if tertiary_mask is not None else None + ), + "decode_parts": tuple(part.detach() for part in decode_parts), + "initial_nested_byte_latents": nested_initial_byte_latents.detach(), + "initial_tertiary_byte_latents": ( + tertiary_initial_byte_latents.detach() + if tertiary_initial_byte_latents is not None + else None + ), + } + self.last_pool_routing_training = { + "fine_logits": close_logits, + "fine_ids": patch_ids.detach(), + "fine_mask": token_mask.detach(), + "fine_minimum": fine_minimum, + "fine_maximum": self.max_patch_bytes, + "fine_threshold": self.close_threshold, + "nested_logits": nested_close_logits, + "nested_ids": nested_ids.detach(), + "nested_mask": patch_mask_nested.detach(), + "nested_minimum": nested_minimum, + "nested_maximum": self.nested_pool_factor, + "nested_threshold": self.nested_close_threshold, + "tertiary_logits": tertiary_close_logits, + "tertiary_ids": tertiary_ids.detach() if tertiary_ids is not None else None, + "tertiary_mask": level2_mask.detach() if level2_mask is not None else None, + "tertiary_minimum": tertiary_minimum, + "tertiary_maximum": self.tertiary_pool_factor, + "tertiary_threshold": self.tertiary_close_threshold, + } + decoded = self.decoder(torch.cat(decode_parts, dim=-1)) + controller_scores = self._pool_controller_scores(decoded) + self.last_pool_controller_scores = controller_scores + if controller_scores is not None: + nested_signal_fine = nested_close_probabilities.to( + x.device, non_blocking=True + ).gather(1, patch_ids) + if tertiary_close_probabilities is not None and tertiary_ids_fine is not None: + tertiary_signal_level2 = tertiary_close_probabilities.to( + x.device, non_blocking=True + ).gather(1, nested_ids_fine) + tertiary_signal_fine = tertiary_signal_level2.gather(1, patch_ids) + else: + tertiary_signal_fine = nested_signal_fine + self.last_pool_controller_close_signals = torch.stack( + [close_probabilities, nested_signal_fine, tertiary_signal_fine], + dim=-1, + ) + else: + self.last_pool_controller_close_signals = None + return self.lm_head(decoded) + + @staticmethod + def _routing_decision_ages( + ids: torch.Tensor, mask: torch.Tensor, *, closes_before_token: bool + ) -> torch.Tensor: + """Age seen by a boundary decision, stable across the resulting close.""" + batch, length = ids.shape + positions = torch.arange(length, device=ids.device).view(1, -1).expand(batch, -1) + group_count = max( + 1, int(ids.masked_fill(~mask, 0).amax().detach().cpu().item()) + 1 + ) + safe_ids = ids.masked_fill(~mask, 0).clamp(0, group_count - 1) + starts = torch.full( + (batch, group_count), length, dtype=torch.long, device=ids.device + ) + starts.scatter_reduce_( + 1, safe_ids, positions.masked_fill(~mask, length), + reduce="amin", include_self=True, + ) + ages = (positions - starts.gather(1, safe_ids) + 1).masked_fill(~mask, 0) + if not closes_before_token: + return ages + # Fine routing closes the previous patch immediately before the current + # byte. On a boundary byte, supervise from the completed previous + # length so the positive label remains stable after the close occurs. + counts = torch.zeros( + (batch, group_count), dtype=torch.long, device=ids.device + ) + counts.scatter_add_(1, safe_ids, mask.long()) + previous_lengths = counts.gather(1, (safe_ids - 1).clamp_min(0)) + starts_group = mask & positions.gt(0) & ages.eq(1) + return torch.where(starts_group, previous_lengths, (ages - 1).clamp_min(0)) + + @staticmethod + def _routing_aggregate( + values: torch.Tensor, ids: torch.Tensor, mask: torch.Tensor + ) -> torch.Tensor: + values = values.to(ids.device) + mask = mask.to(ids.device) + groups = max( + 1, int(ids.masked_fill(~mask, 0).amax().detach().cpu().item()) + 1 + ) + safe_ids = ids.masked_fill(~mask, 0).clamp(0, groups - 1) + sums = values.new_zeros((values.shape[0], groups)) + counts = values.new_zeros((values.shape[0], groups)) + sums.scatter_add_(1, safe_ids, values.masked_fill(~mask, 0.0)) + counts.scatter_add_(1, safe_ids, mask.to(values.dtype)) + return sums / counts.clamp_min(1.0) + + @staticmethod + def _routing_future_surprise( + values: torch.Tensor, valid: torch.Tensor + ) -> torch.Tensor: + future = torch.cat([values[:, 1:], values[:, -1:]], dim=1).detach().float() + valid = valid.to(future.device) + selected = future[valid] + if not selected.numel(): + return future.new_zeros(future.shape) + return torch.tanh( + (future - selected.mean()) / selected.std(unbiased=False).clamp_min(1e-4) + ).masked_fill(~valid, 0.0) + + @staticmethod + def _routing_level_objective( + logits: torch.Tensor, + ids: torch.Tensor, + mask: torch.Tensor, + surprise: torch.Tensor, + *, + minimum: int, + maximum: int, + target_length: int, + semantic_span: float, + closes_before_token: bool, + threshold: float, + collect_metrics: bool = True, + ) -> Tuple[torch.Tensor, Dict[str, float]]: + mask = mask.to(logits.device) + ids = ids.to(logits.device) + surprise = surprise.to(logits.device) + minimum = max(1, min(int(minimum), int(maximum))) + target_length = max(minimum, min(int(target_length), int(maximum))) + valid_logits = logits.float()[mask] + if not valid_logits.numel(): + return logits.float().sum() * 0.0, {} + valid_surprise = surprise.float()[mask] + # Select a fixed causal refresh budget, assigning it to positions whose + # following region was hardest to predict. This makes the close head a + # semantic selector rather than asking a content-only network to infer + # an invisible current-pool age. + target_rate = min(1.0, 1.0 / max(1.0, float(target_length))) + positive_count = max(1, int(round(valid_logits.numel() * target_rate))) + positive_count = min(positive_count, valid_logits.numel()) + selected = torch.topk(valid_surprise, positive_count, sorted=False).indices + labels = torch.zeros_like(valid_logits) + labels[selected] = 1.0 + threshold = max(1e-4, min(1.0 - 1e-4, float(threshold))) + threshold_logit = math.log(threshold / (1.0 - threshold)) + decision_logits = (valid_logits - threshold_logit) / 0.5 + semantic_loss = F.binary_cross_entropy_with_logits(decision_logits, labels) + soft_rate = torch.sigmoid(decision_logits).mean() + rate_loss = (soft_rate - target_rate) ** 2 + objective = max(0.0, float(semantic_span)) * semantic_loss + 4.0 * rate_loss + if not collect_metrics: + return objective, {} + probabilities = torch.sigmoid(valid_logits.detach()) + return objective, { + "prob_mean": float(probabilities.mean().cpu()), + "prob_std": float(probabilities.std(unbiased=False).cpu()), + "above_threshold": float(probabilities.ge(float(threshold)).float().mean().cpu()), + "target_length": float(target_length), + "target_close_rate": float(target_rate), + "soft_close_rate": float(soft_rate.detach().cpu()), + } + + def pool_routing_auxiliary_loss( + self, + logits: torch.Tensor, + targets: torch.Tensor, + *, + fine_target: int = 0, + nested_target: int = 0, + tertiary_target: int = 0, + semantic_span: float = 2.0, + collect_metrics: bool = True, + ) -> Tuple[torch.Tensor, Dict[str, float]]: + """Explicitly train causal close heads against rate and surprise labels.""" + state = self.last_pool_routing_training + if not state: + return logits.float().sum() * 0.0, {} + with torch.no_grad(): + byte_loss = F.cross_entropy( + logits.detach().float().transpose(1, 2), targets, + ignore_index=-100, reduction="none", + ) + fine_mask = state["fine_mask"] + scored_bytes = targets.ne(-100).to(fine_mask.device) & fine_mask + fine_surprise = self._routing_future_surprise( + byte_loss.to(fine_mask.device), scored_bytes + ) + fine_group_loss = self._routing_aggregate( + byte_loss, state["fine_ids"], scored_bytes + ) + nested_mask = state["nested_mask"] + nested_surprise = self._routing_future_surprise( + fine_group_loss.to(nested_mask.device), nested_mask + ) + nested_group_loss = self._routing_aggregate( + fine_group_loss, state["nested_ids"], nested_mask + ) + tertiary_mask = state["tertiary_mask"] + tertiary_surprise = ( + self._routing_future_surprise( + nested_group_loss.to(tertiary_mask.device), tertiary_mask + ) + if tertiary_mask is not None else None + ) + + def target(requested: int, prefix: str) -> int: + minimum = int(state[f"{prefix}_minimum"]) + maximum = int(state[f"{prefix}_maximum"]) + return ( + max(minimum, min(int(requested), maximum)) + if int(requested) > 0 + else max(minimum, min(maximum, minimum * 4)) + ) + + levels = [ + ("fine", "fine", fine_target, fine_surprise, True), + ("level2", "nested", nested_target, nested_surprise, False), + ] + if state["tertiary_logits"] is not None: + levels.append(( + "level3", "tertiary", tertiary_target, tertiary_surprise, False + )) + objectives: List[torch.Tensor] = [] + metrics: Dict[str, float] = {} + destination = logits.device + for name, prefix, requested, surprise, before in levels: + level_objective, level_metrics = self._routing_level_objective( + state[f"{prefix}_logits"], state[f"{prefix}_ids"], + state[f"{prefix}_mask"], surprise, + minimum=int(state[f"{prefix}_minimum"]), + maximum=int(state[f"{prefix}_maximum"]), + target_length=target(requested, prefix), + semantic_span=semantic_span, + closes_before_token=before, + threshold=float(state[f"{prefix}_threshold"]), + collect_metrics=collect_metrics, + ) + objectives.append(level_objective.to(destination, non_blocking=True)) + metrics.update({f"{name}_{key}": value for key, value in level_metrics.items()}) + return torch.stack(objectives).mean(), metrics + + +class ForwardBackwardRepairModel(nn.Module): + def __init__( + self, + vocab_size: int, + dim: int = 256, + layers: int = 4, + position_bins: int = 8192, + use_mamba: bool = True, + mamba_version: int = 2, + mamba_d_state: int = 64, + mamba2_headdim: int = 0, + bidirectional: bool = False, + num_sections: int = 16, + min_patch_bytes: int = 4, + max_patch_bytes: int = 16, + patch_change_threshold: int = 48, + close_threshold: float = 0.90, + mid_close_bonus: float = 0.05, + nested_pool_factor: int = 16, + nested_min_pool_factor: int = 0, + nested_close_threshold: Optional[float] = None, + nested_layers: int = 0, + tertiary_pool_factor: int = 0, + tertiary_min_pool_factor: int = 0, + tertiary_close_threshold: Optional[float] = None, + tertiary_layers: int = 0, + decoder_dim: int = 0, + detach_inactive_coarse_gradients: bool = True, + legacy_pool_closure_training: bool = False, + decoder_pool_controller: bool = False, + pool_controller_alpha: float = 1.0, + pool_controller_beta: float = 0.5, + pool_controller_gamma: float = 0.5, + short_pool_budget: int = 0, + short_pool_window: int = 0, + secondary_min_patch_bytes: int = 16, + ): + super().__init__() + if bidirectional: + raise ValueError("BLT Mamba trainer is forward-only; disable --bidirectional-repair.") + self.bidirectional = False + self.forward_model = ByteLatentMambaCore( + vocab_size, + dim=dim, + layers=layers, + position_bins=position_bins, + use_mamba=use_mamba, + mamba_version=mamba_version, + mamba_d_state=mamba_d_state, + mamba2_headdim=mamba2_headdim, + num_sections=num_sections, + min_patch_bytes=min_patch_bytes, + max_patch_bytes=max_patch_bytes, + patch_change_threshold=patch_change_threshold, + close_threshold=close_threshold, + mid_close_bonus=mid_close_bonus, + nested_pool_factor=nested_pool_factor, + nested_min_pool_factor=nested_min_pool_factor, + nested_close_threshold=nested_close_threshold, + nested_layers=nested_layers, + tertiary_pool_factor=tertiary_pool_factor, + tertiary_min_pool_factor=tertiary_min_pool_factor, + tertiary_close_threshold=tertiary_close_threshold, + tertiary_layers=tertiary_layers, + decoder_dim=decoder_dim, + detach_inactive_coarse_gradients=detach_inactive_coarse_gradients, + legacy_pool_closure_training=legacy_pool_closure_training, + decoder_pool_controller=decoder_pool_controller, + pool_controller_alpha=pool_controller_alpha, + pool_controller_beta=pool_controller_beta, + pool_controller_gamma=pool_controller_gamma, + short_pool_budget=short_pool_budget, + short_pool_window=short_pool_window, + secondary_min_patch_bytes=secondary_min_patch_bytes, + ) + + def configure_model_parallel(self, fine_device: torch.device, nested_devices: List[torch.device], tertiary_device: Optional[torch.device] = None) -> None: + """Enable nested-level model parallelism without changing checkpoint layout.""" + self.forward_model.configure_model_parallel(fine_device, nested_devices, tertiary_device=tertiary_device) + + def forward(self, batch, inference_params=None): + return self.forward_model(**batch, inference_params=inference_params) +