| """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 |
|
|
|
|
| def _new_mamba_inference_params(max_batch_size: int, max_seqlen: int): |
| if InferenceParams is None: |
| return None |
| try: |
| return InferenceParams( |
| max_seqlen=max_seqlen, |
| max_batch_size=max_batch_size, |
| ) |
| except TypeError: |
| try: |
| return InferenceParams( |
| max_batch_size=max_batch_size, |
| max_seqlen=max_seqlen, |
| ) |
| except TypeError: |
| return InferenceParams(max_seqlen, max_batch_size) |
|
|
|
|
| def _set_inference_seqlen_offset(inference_params, offset: int) -> None: |
| if inference_params is not None and hasattr(inference_params, "seqlen_offset"): |
| inference_params.seqlen_offset = int(offset) |
|
|
| 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) |
| |
| |
| |
| |
| 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) |
| |
| |
| 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) |
| |
| |
| |
| 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) |
| |
| |
| |
| 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(), |
| ) |
| |
| |
| |
| |
| |
| 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) |
| |
| |
| |
| self.model_parallel_enabled = False |
| |
| |
| |
| 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 |
| ) |
| |
| |
| 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, |
| ) |
| |
| |
| |
| 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(()) |
| |
| |
| |
| |
| |
| |
| 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") |
|
|
| |
| 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): |
| |
| |
| 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()): |
| |
| |
| |
| 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) |
| ) |
|
|
| |
| |
| 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: |
| |
| |
| 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)) |
| |
| |
| |
| |
| 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": |
| |
| |
| 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)) |
| |
| |
| 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: |
| |
| |
| |
| |
| 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) |
| |
| |
| |
| 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) |
|
|
| |
| |
| 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) |
|
|
| |
| |
| |
| |
| |
| fine_h = global_h[:, 1:] |
| |
| |
| 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) |
| |
| 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) |
| |
| |
| 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), |
| ) |
| |
| 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)) |
| |
| |
| 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 |
| |
| |
| |
| 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] |
| |
| |
| |
| |
| 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) |
|
|