| |
| """Shared loading, streaming, and sampling helpers for nested byte Mamba tools.""" |
| from __future__ import annotations |
|
|
| import gc |
| import json |
| import math |
| import random |
| import warnings |
| from dataclasses import dataclass |
| from pathlib import Path |
| from typing import Dict, Iterable, List, Optional, Sequence, Tuple |
|
|
| import torch |
| import torch.nn.functional as F |
|
|
| from modeling_nested_mamba import ( |
| ForwardBackwardRepairModel, |
| ) |
|
|
|
|
| PAD = 0 |
| BOS = 1 |
| EOS = 2 |
| UNK = 3 |
| BYTE_OFFSET = 4 |
| VOCAB_SIZE = 260 |
|
|
|
|
| @dataclass |
| class LoadedNestedModel: |
| model: ForwardBackwardRepairModel |
| config: Dict[str, object] |
| checkpoint_step: int |
| trained_tokens: int |
| checkpoint_path: Path |
| precision: str |
| fine_device: torch.device |
| model_parallel: bool |
|
|
| @property |
| def core(self): |
| return self.model.forward_model |
|
|
|
|
| def _dtype_for_precision(precision: str) -> torch.dtype: |
| normalized = precision.lower() |
| if normalized == "fp16": |
| return torch.float16 |
| if normalized == "bf16": |
| return torch.bfloat16 |
| if normalized == "fp32": |
| return torch.float32 |
| raise ValueError(f"unsupported precision {precision!r}; choose fp16, bf16, or fp32") |
|
|
|
|
| def _mamba2_fused_causal_conv_available() -> bool: |
| """Whether Mamba-2's combined full-sequence kernel can call causal-conv1d.""" |
| try: |
| from mamba_ssm.ops.triton import ssd_combined |
|
|
| return getattr(ssd_combined, "causal_conv1d_fwd_function", None) is not None |
| except Exception: |
| return False |
|
|
|
|
| def _configure_mamba2_inference_kernels(model: ForwardBackwardRepairModel) -> bool: |
| """Select the portable full-sequence path when fused causal-conv1d is absent.""" |
| fused_available = _mamba2_fused_causal_conv_available() |
| changed = False |
| for block_group in ( |
| model.forward_model.global_blocks, |
| model.forward_model.nested_global_blocks, |
| model.forward_model.tertiary_global_blocks, |
| ): |
| for block in block_group: |
| if int(getattr(block, "mamba_version", 1)) != 2: |
| continue |
| if not fused_available and bool(getattr(block.mixer, "use_mem_eff_path", False)): |
| block.mixer.use_mem_eff_path = False |
| changed = True |
| return changed |
|
|
|
|
| def _build_model(config: Dict[str, object]) -> ForwardBackwardRepairModel: |
| if config.get("training_method") == "diffusionblocks": |
| raise ValueError( |
| "this checkpoint declares an architecture that is incompatible with " |
| "the nested Mamba inference implementation" |
| ) |
| model = ForwardBackwardRepairModel( |
| vocab_size=int(config.get("vocab_size", VOCAB_SIZE)), |
| dim=int(config["dim"]), |
| layers=int(config["layers"]), |
| position_bins=int(config.get("position_bins", 8192)), |
| use_mamba=not bool(config.get("no_mamba", False)), |
| |
| mamba_version=int(config.get("mamba_version", 1)), |
| mamba_d_state=int( |
| config.get( |
| "mamba_d_state", |
| 64 if int(config.get("mamba_version", 1)) == 2 else 16, |
| ) |
| ), |
| mamba2_headdim=int(config.get("mamba2_headdim", 0)), |
| num_sections=1, |
| min_patch_bytes=int(config.get("blt_min_patch_bytes", 16)), |
| max_patch_bytes=int(config.get("blt_max_patch_bytes", 96)), |
| patch_change_threshold=int(config.get("blt_patch_change_threshold", 48)), |
| close_threshold=float(config.get("blt_close_threshold", 0.98)), |
| mid_close_bonus=float(config.get("blt_mid_close_bonus", 0.05)), |
| nested_pool_factor=int(config.get("nested_pool_factor", 16)), |
| nested_min_pool_factor=int(config.get("nested_min_pool_factor", 0)), |
| nested_close_threshold=( |
| None if config.get("nested_close_threshold") is None |
| else float(config["nested_close_threshold"]) |
| ), |
| nested_layers=int(config.get("nested_layers", 0)), |
| tertiary_pool_factor=int(config.get("tertiary_pool_factor", 0)), |
| tertiary_min_pool_factor=int(config.get("tertiary_min_pool_factor", 0)), |
| tertiary_close_threshold=( |
| None if config.get("tertiary_close_threshold") is None |
| else float(config["tertiary_close_threshold"]) |
| ), |
| tertiary_layers=int(config.get("tertiary_layers", 0)), |
| decoder_dim=int(config.get("decoder_dim", 0)), |
| detach_inactive_coarse_gradients=bool( |
| config.get("detach_inactive_coarse_gradients", True) |
| ), |
| decoder_pool_controller=bool(config.get("decoder_pool_controller", False)), |
| pool_controller_alpha=float(config.get("pool_controller_alpha", 1.0)), |
| pool_controller_beta=float(config.get("pool_controller_beta", 0.5)), |
| pool_controller_gamma=float(config.get("pool_controller_gamma", 0.5)), |
| |
| |
| short_pool_budget=int(config.get("blt_short_pool_budget", 0)), |
| short_pool_window=int(config.get("blt_short_pool_window", 0)), |
| secondary_min_patch_bytes=int(config.get("blt_secondary_min_patch_bytes", 16)), |
| ) |
| model.forward_model.mamba2_unfused_inference = _configure_mamba2_inference_kernels(model) |
| return model |
|
|
|
|
| def _parse_device_list(value: Optional[str]) -> List[torch.device]: |
| if not value: |
| return [] |
| return [torch.device(item.strip()) for item in value.split(",") if item.strip()] |
|
|
|
|
| def _validate_cuda_devices(devices: Sequence[torch.device]) -> None: |
| if not torch.cuda.is_available(): |
| raise RuntimeError( |
| "CUDA is unavailable. mamba_ssm selective-scan inference requires CUDA in this environment." |
| ) |
| count = torch.cuda.device_count() |
| for device in devices: |
| if device.type != "cuda" or device.index is None or device.index >= count: |
| raise ValueError(f"requested device {device} is unavailable; CUDA device count is {count}") |
|
|
|
|
| def load_nested_checkpoint( |
| checkpoint_path: str, |
| *, |
| precision: str = "fp16", |
| device: str = "cuda:0", |
| fine_device: Optional[str] = None, |
| nested_devices: Optional[str] = None, |
| tertiary_device: Optional[str] = None, |
| use_saved_placement: bool = False, |
| ) -> LoadedNestedModel: |
| """Load weights without materializing checkpoint optimizer tensors. |
| |
| By default all inference runs on ``--device``. Model-parallel placement is |
| enabled by supplying ``fine_device`` and ``nested_devices``, or by opting |
| into the placement recorded in the checkpoint. |
| """ |
| path = Path(checkpoint_path).expanduser().resolve() |
| checkpoint = None |
| if path.is_dir(): |
| config_path = path / "config.json" |
| if not config_path.is_file(): |
| raise FileNotFoundError(f"model config not found: {config_path}") |
| config = json.loads(config_path.read_text(encoding="utf-8")) |
| elif path.is_file(): |
| |
| warnings.warn( |
| "Loading a PyTorch .pt checkpoint requires pickle deserialization. " |
| "Only load .pt files that you created or obtained from a trusted source; " |
| "use the published SafeTensors directory for untrusted downloads.", |
| UserWarning, |
| stacklevel=2, |
| ) |
| checkpoint = torch.load(path, map_location="cpu", weights_only=False, mmap=True) |
| config = dict(checkpoint.get("config") or {}) |
| else: |
| raise FileNotFoundError(f"model directory or checkpoint not found: {path}") |
| if config.get("architecture") not in (None, "byte_latent_mamba_nested_jsonl") and config.get( |
| "architecture_label" |
| ) != "byte_latent_mamba_nested_jsonl": |
| raise ValueError(f"{path} is not identified as a nested JSONL checkpoint") |
|
|
| with torch.device("meta"): |
| model = _build_model(config) |
| if bool(getattr(model.forward_model, "mamba2_unfused_inference", False)): |
| print( |
| "Mamba-2 fused causal-conv1d is unavailable; using the portable " |
| "batched convolution + SSD scan path." |
| ) |
| if path.is_dir(): |
| from safetensors.torch import load_file |
|
|
| index_path = path / "model.safetensors.index.json" |
| single_path = path / "model.safetensors" |
| if index_path.is_file(): |
| index = json.loads(index_path.read_text(encoding="utf-8")) |
| weight_map = dict(index.get("weight_map") or {}) |
| expected = set(model.state_dict().keys()) |
| published = set(weight_map.keys()) |
| if expected != published: |
| missing = sorted(expected - published)[:8] |
| unexpected = sorted(published - expected)[:8] |
| raise RuntimeError( |
| f"SafeTensors index does not match architecture; " |
| f"missing={missing}, unexpected={unexpected}" |
| ) |
| for filename in dict.fromkeys(weight_map.values()): |
| shard_path = path / filename |
| shard = load_file(str(shard_path), device="cpu") |
| model.load_state_dict(shard, strict=False, assign=True) |
| del shard |
| elif single_path.is_file(): |
| state = load_file(str(single_path), device="cpu") |
| model.load_state_dict(state, strict=True, assign=True) |
| del state |
| else: |
| raise FileNotFoundError( |
| f"no model.safetensors or model.safetensors.index.json in {path}" |
| ) |
| meta_names = [ |
| name for name, value in model.state_dict().items() if value.device.type == "meta" |
| ] |
| if meta_names: |
| raise RuntimeError(f"unloaded model tensors remain: {meta_names[:8]}") |
| step = 0 |
| trained_tokens = 0 |
| else: |
| model.load_state_dict(checkpoint["model"], strict=True, assign=True) |
| step = int(checkpoint.get("step", 0)) |
| trained_tokens = int(checkpoint.get("trained_tokens", 0)) |
| del checkpoint |
| gc.collect() |
|
|
| dtype = _dtype_for_precision(precision) |
| if int(config.get("mamba_version", 1)) == 2 and dtype == torch.float16: |
| print( |
| "Mamba-2 FP16 inference: cached SSM accumulators will remain FP32; " |
| "BF16 is recommended when the GPU supports it." |
| ) |
| if use_saved_placement: |
| fine_device = fine_device or str(config.get("fine_device") or "cuda:0") |
| if nested_devices is None: |
| saved_nested = config.get("nested_devices") or [] |
| nested_devices = ",".join(str(item) for item in saved_nested) |
| tertiary_device = tertiary_device or ( |
| str(config["tertiary_device"]) if config.get("tertiary_device") else None |
| ) |
|
|
| coarse = _parse_device_list(nested_devices) |
| if coarse: |
| if not fine_device: |
| raise ValueError("--fine-device is required with --nested-devices") |
| fine = torch.device(fine_device) |
| tertiary = torch.device(tertiary_device) if tertiary_device else None |
| requested = [fine, *coarse, *([tertiary] if tertiary else [])] |
| _validate_cuda_devices(requested) |
| |
| |
| model.to(dtype=dtype) |
| model.configure_model_parallel(fine, coarse, tertiary_device=tertiary) |
| root = fine |
| parallel = True |
| else: |
| root = torch.device(device) |
| _validate_cuda_devices([root]) |
| model.to(root, dtype=dtype) |
| parallel = False |
|
|
| model.eval() |
| return LoadedNestedModel( |
| model=model, |
| config=config, |
| checkpoint_step=step, |
| trained_tokens=trained_tokens, |
| checkpoint_path=path, |
| precision=precision, |
| fine_device=root, |
| model_parallel=parallel, |
| ) |
|
|
|
|
| def new_stream(loaded: LoadedNestedModel, maximum_input_bytes: int) -> Dict[str, object]: |
| minimum = max(1, int(loaded.config.get("blt_min_patch_bytes", 16))) |
| max_patches = math.ceil((int(maximum_input_bytes) + 1) / minimum) + 8 |
| return loaded.core.new_stream_state({}, max_patches=max_patches) |
|
|
|
|
| def stream_token( |
| loaded: LoadedNestedModel, |
| stream: Dict[str, object], |
| token_id: int, |
| ) -> torch.Tensor: |
| token = torch.tensor( |
| [[int(token_id)]], dtype=torch.long, device=loaded.fine_device |
| ) |
| return loaded.core.stream_step(stream, {"x": token})[0, 0] |
|
|
|
|
| def apply_sampling_filters( |
| logits: torch.Tensor, |
| *, |
| temperature: float, |
| top_p: float, |
| top_k: int, |
| repeat_penalty: float, |
| recent_tokens: Sequence[int], |
| ) -> torch.Tensor: |
| filtered = logits.float().clone() |
| filtered[PAD] = filtered[BOS] = filtered[UNK] = -torch.inf |
| if repeat_penalty > 1.0: |
| for token_id in set(int(value) for value in recent_tokens): |
| if 0 <= token_id < filtered.numel(): |
| value = filtered[token_id] |
| filtered[token_id] = ( |
| value / repeat_penalty if value >= 0 else value * repeat_penalty |
| ) |
| filtered /= max(1e-5, float(temperature)) |
| if top_k > 0 and top_k < filtered.numel(): |
| threshold = torch.topk(filtered, int(top_k)).values[-1] |
| filtered[filtered < threshold] = -torch.inf |
| if 0.0 < top_p < 1.0: |
| probabilities = torch.softmax(filtered, dim=-1) |
| sorted_probabilities, sorted_indices = torch.sort(probabilities, descending=True) |
| remove = torch.cumsum(sorted_probabilities, dim=0) > float(top_p) |
| remove[0] = False |
| filtered[sorted_indices[remove]] = -torch.inf |
| return filtered |
|
|
|
|
| def choose_token(filtered_logits: torch.Tensor, greedy: bool = False) -> int: |
| finite = torch.isfinite(filtered_logits) |
| if not bool(finite.any().item()): |
| raise FloatingPointError( |
| "sampling has no finite logits; the recurrent inference state became " |
| "non-finite. Retry with --precision bf16 (recommended for Mamba-2) " |
| "or --precision fp32." |
| ) |
| if greedy: |
| return int(filtered_logits.argmax()) |
| probabilities = torch.softmax(filtered_logits, dim=-1) |
| if not bool(torch.isfinite(probabilities).all().item()): |
| raise FloatingPointError( |
| "sampling probabilities became non-finite. Retry with --precision " |
| "bf16 (recommended for Mamba-2) or --precision fp32." |
| ) |
| return int(torch.multinomial(probabilities, 1)) |
|
|
|
|
| def hierarchy_token_attribution( |
| loaded: LoadedNestedModel, |
| stream: Dict[str, object], |
| logits: torch.Tensor, |
| token_id: int, |
| ) -> Dict[str, object]: |
| """Measure direct L2/L3 decoder influence on one selected next byte. |
| |
| This reuses the cached streaming state and only reruns the small decoder |
| and LM head. Positive deltas mean the dynamic hierarchy increased the |
| selected token's log probability relative to its BOE counterfactual. |
| """ |
| parts = list(loaded.core.last_stream_decode_parts) |
| decoder_device = parts[0].device |
| normal_log_probability = float( |
| F.log_softmax(logits.float(), dim=-1)[int(token_id)].detach().cpu().item() |
| ) |
|
|
| def counterfactual(*, remove_l2: bool, remove_l3: bool) -> float: |
| altered = list(parts) |
| if remove_l2: |
| initial_nested = stream["initial_nested_global"] |
| if initial_nested.device != decoder_device: |
| initial_nested = initial_nested.to(decoder_device, non_blocking=True) |
| altered[3] = initial_nested |
| if remove_l3 and len(altered) > 4: |
| initial_tertiary = stream["initial_tertiary_global"] |
| if initial_tertiary.device != decoder_device: |
| initial_tertiary = initial_tertiary.to( |
| decoder_device, non_blocking=True |
| ) |
| altered[4] = initial_tertiary |
| altered_logits = loaded.core.lm_head( |
| loaded.core.decoder(torch.cat(altered, dim=-1)) |
| )[0, 0].float() |
| return float( |
| F.log_softmax(altered_logits, dim=-1)[int(token_id)].detach().cpu().item() |
| ) |
|
|
| level2_active = int(stream["completed_nested_patches"]) > 0 |
| level3_active = int(stream.get("completed_tertiary_patches", 0)) > 0 |
| without_l2 = ( |
| counterfactual(remove_l2=True, remove_l3=False) |
| if level2_active |
| else normal_log_probability |
| ) |
| without_l3 = ( |
| counterfactual(remove_l2=False, remove_l3=True) |
| if level3_active |
| else normal_log_probability |
| ) |
| without_hierarchy = ( |
| counterfactual(remove_l2=True, remove_l3=True) |
| if level3_active |
| else without_l2 |
| ) |
| return { |
| "level2_active": level2_active, |
| "level3_active": level3_active, |
| "level2_delta_logp": normal_log_probability - without_l2, |
| "level3_delta_logp": normal_log_probability - without_l3, |
| "hierarchy_delta_logp": normal_log_probability - without_hierarchy, |
| "selected_logp": normal_log_probability, |
| } |
|
|
|
|
| def hierarchy_mode_logits( |
| loaded: LoadedNestedModel, |
| stream: Dict[str, object], |
| logits: torch.Tensor, |
| mode: str, |
| ) -> torch.Tensor: |
| """Return next-token logits with selected hierarchy readouts disabled.""" |
| normalized = str(mode).lower() |
| if normalized == "full": |
| return logits |
| if normalized not in {"level1", "level12"}: |
| raise ValueError("hierarchy mode must be 'level1', 'level12', or 'full'") |
| parts = list(loaded.core.last_stream_decode_parts) |
| decoder_device = parts[0].device |
| if normalized == "level1": |
| initial_nested = stream["initial_nested_global"] |
| if initial_nested.device != decoder_device: |
| initial_nested = initial_nested.to(decoder_device, non_blocking=True) |
| parts[3] = initial_nested |
| if len(parts) > 4: |
| initial_tertiary = stream["initial_tertiary_global"] |
| if initial_tertiary.device != decoder_device: |
| initial_tertiary = initial_tertiary.to(decoder_device, non_blocking=True) |
| parts[4] = initial_tertiary |
| return loaded.core.lm_head( |
| loaded.core.decoder(torch.cat(parts, dim=-1)) |
| )[0, 0] |
|
|
|
|
| def generate_bytes( |
| loaded: LoadedNestedModel, |
| prompt: bytes, |
| *, |
| max_new_bytes: int, |
| temperature: float = 0.8, |
| top_p: float = 0.9, |
| top_k: int = 0, |
| repeat_penalty: float = 1.05, |
| repeat_window: int = 256, |
| greedy: bool = False, |
| seed: int = 1234, |
| collect_hierarchy_attribution: bool = False, |
| hierarchy_mode: str = "full", |
| ) -> Tuple[bytes, Dict[str, object]]: |
| if collect_hierarchy_attribution and hierarchy_mode != "full": |
| raise ValueError("hierarchy attribution is defined for the full hierarchy rollout") |
| torch.manual_seed(seed) |
| random.seed(seed) |
| stream = new_stream(loaded, len(prompt) + max_new_bytes + 2) |
| with torch.inference_mode(): |
| logits = stream_token(loaded, stream, BOS) |
| for value in prompt: |
| logits = stream_token(loaded, stream, BYTE_OFFSET + int(value)) |
| output = bytearray() |
| attributions: List[Dict[str, object]] = [] |
| recent: List[int] = [BYTE_OFFSET + int(value) for value in prompt[-repeat_window:]] |
| for generated_position in range(max(0, int(max_new_bytes))): |
| sampling_logits = hierarchy_mode_logits( |
| loaded, stream, logits, hierarchy_mode |
| ) |
| if not bool(torch.isfinite(sampling_logits).all().item()): |
| finite_count = int(torch.isfinite(sampling_logits).sum().item()) |
| raise FloatingPointError( |
| "non-finite generation logits before sampling: " |
| f"generated_byte={generated_position} mode={hierarchy_mode} " |
| f"precision={loaded.precision} " |
| f"finite_logits={finite_count}/{sampling_logits.numel()}. " |
| "Retry with --precision bf16 (recommended for Mamba-2) or " |
| "--precision fp32." |
| ) |
| filtered = apply_sampling_filters( |
| sampling_logits, |
| temperature=temperature, |
| top_p=top_p, |
| top_k=top_k, |
| repeat_penalty=repeat_penalty, |
| recent_tokens=recent[-repeat_window:], |
| ) |
| token_id = choose_token(filtered, greedy=greedy) |
| if token_id == EOS: |
| break |
| if not BYTE_OFFSET <= token_id < BYTE_OFFSET + 256: |
| continue |
| if collect_hierarchy_attribution: |
| attributions.append( |
| hierarchy_token_attribution( |
| loaded, stream, sampling_logits, token_id |
| ) |
| ) |
| output.append(token_id - BYTE_OFFSET) |
| recent.append(token_id) |
| logits = stream_token(loaded, stream, token_id) |
| stream["generated_attribution"] = attributions |
| stream["hierarchy_mode"] = hierarchy_mode |
| return bytes(output), stream |
|
|
|
|
| def load_jsonl_texts( |
| path: str, |
| *, |
| text_field: str = "text", |
| max_documents: Optional[int] = None, |
| ) -> Iterable[Tuple[int, str]]: |
| jsonl = Path(path).expanduser().resolve() |
| if not jsonl.is_file(): |
| raise FileNotFoundError(f"JSONL file not found: {jsonl}") |
| yielded = 0 |
| with jsonl.open("r", encoding="utf-8") as handle: |
| for line_number, line in enumerate(handle, 1): |
| if not line.strip(): |
| continue |
| try: |
| record = json.loads(line) |
| except json.JSONDecodeError as error: |
| raise ValueError(f"invalid JSON at {jsonl}:{line_number}: {error.msg}") from error |
| text = record.get(text_field) |
| if not isinstance(text, str): |
| raise ValueError( |
| f"{jsonl}:{line_number} must contain a string field {text_field!r}" |
| ) |
| yield line_number, text |
| yielded += 1 |
| if max_documents is not None and yielded >= int(max_documents): |
| return |
|
|