Text Generation
Transformers
Safetensors
English
microloop_diffusion
causal-lm
base-model
small-language-model
custom_code
muon
hummingbird
hummingbird-v2
conversational
Instructions to use juinron/Hummingbird-V2 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use juinron/Hummingbird-V2 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="juinron/Hummingbird-V2", trust_remote_code=True) messages = [ {"role": "user", "content": "Who are you?"}, ] pipe(messages)# Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("juinron/Hummingbird-V2", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use juinron/Hummingbird-V2 with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "juinron/Hummingbird-V2" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "juinron/Hummingbird-V2", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker
docker model run hf.co/juinron/Hummingbird-V2
- SGLang
How to use juinron/Hummingbird-V2 with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "juinron/Hummingbird-V2" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "juinron/Hummingbird-V2", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "juinron/Hummingbird-V2" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "juinron/Hummingbird-V2", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }' - Docker Model Runner
How to use juinron/Hummingbird-V2 with Docker Model Runner:
docker model run hf.co/juinron/Hummingbird-V2
| """Small, HF-compatible causal bootstrap model used by the M0 gate. | |
| The diffusion objective and sampler are deliberately separate modules. This model | |
| provides the shared transformer backbone and a causal forward path so that the | |
| project can validate shape correctness, parameter accounting, and reproducibility | |
| before any expensive data work begins. | |
| """ | |
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| from hashlib import sha256 | |
| from typing import Optional | |
| import torch | |
| from torch import Tensor, nn | |
| from torch.nn import functional as F | |
| from .configuration_microloop import MicroLoopConfig | |
| from .prefix_lm import build_prefix_lm_mask | |
| try: | |
| from transformers import PreTrainedModel | |
| from transformers.generation import GenerationMixin | |
| from transformers.utils import ModelOutput | |
| except ImportError: # pragma: no cover - only used in a minimal environment. | |
| class GenerationMixin: # type: ignore[no-redef] | |
| pass | |
| class ModelOutput: # type: ignore[no-redef] | |
| pass | |
| class PreTrainedModel(nn.Module): # type: ignore[no-redef] | |
| config_class = MicroLoopConfig | |
| base_model_prefix = "microloop" | |
| def __init__(self, config: MicroLoopConfig) -> None: | |
| super().__init__() | |
| self.config = config | |
| class MicroLoopCausalLMOutput(ModelOutput): | |
| """Minimal output object with both attribute and mapping-style access.""" | |
| logits: Optional[Tensor] = None | |
| loss: Optional[Tensor] = None | |
| hidden_states: Optional[Tensor] = None | |
| loop_applications: Optional[int] = None | |
| def __getitem__(self, key: str): | |
| return getattr(self, key) | |
| class RMSNorm(nn.Module): | |
| def __init__(self, hidden_size: int, eps: float) -> None: | |
| super().__init__() | |
| self.weight = nn.Parameter(torch.ones(hidden_size)) | |
| self.eps = eps | |
| def forward(self, hidden_states: Tensor) -> Tensor: | |
| # Explicit computation rather than the fused ``F.rms_norm`` kernel: | |
| # the fused kernel selects implementations based on process-level state | |
| # and produces context-dependent numerics inside the training process | |
| # (the 2026-08-05 provenance incident). This explicit path is | |
| # deterministic everywhere; the small speed cost is acceptable here. | |
| variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True) | |
| return (hidden_states / torch.sqrt(variance + self.eps)) * self.weight | |
| class UnweightedRMSNorm(nn.Module): | |
| """Parameter-free RMS normalization used by the mHC routing projections.""" | |
| def __init__(self, eps: float) -> None: | |
| super().__init__() | |
| self.eps = eps | |
| def forward(self, hidden_states: Tensor) -> Tensor: | |
| variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True) | |
| return hidden_states * torch.rsqrt(variance + self.eps).to(hidden_states.dtype) | |
| class NGramMemory(nn.Module): | |
| """Small causal suffix-N-gram memory for controlled experiments. | |
| The lookup mode keeps raw token IDs unchanged, which is deliberately safe for | |
| digits and mathematical symbols. Each order/head has its own deterministic | |
| hash table; the projected value is gated by the current hidden state and the | |
| value projection is zero-initialized by ``MicroLoopForDiffusionLM``. | |
| """ | |
| def __init__(self, config: MicroLoopConfig) -> None: | |
| super().__init__() | |
| settings = dict(config.ngram_memory) | |
| self.mode = str(settings.get("mode", "lookup")) | |
| self.hash_version = str(settings.get("hash_version", "legacy_v1")) | |
| self.orders = tuple(int(order) for order in settings.get("orders", [2, 3])) | |
| self.num_hash_heads = int(settings.get("num_hash_heads", 2)) | |
| self.num_buckets = int(settings.get("num_buckets", 16_384)) | |
| self.embedding_dim = int(settings.get("embedding_dim", 16)) | |
| self.composition_scale = float(settings.get("composition_scale", 0.1)) | |
| self.eps = config.rms_norm_eps | |
| # Fixed, independent bases; no global RNG or checkpoint tensors needed. | |
| # Arithmetic modulo 2**31-1 keeps int64 products below 2**62. | |
| self.hash_bases = tuple( | |
| 2 + int.from_bytes(sha256(f"microloop-ngram-v2:{i}".encode()).digest()[:4], "big") | |
| % (2**31 - 3) | |
| for i in range(len(self.orders) * self.num_hash_heads) | |
| ) | |
| if self.mode == "lookup": | |
| table_count = len(self.orders) * self.num_hash_heads | |
| self.tables = nn.ModuleList( | |
| nn.Embedding(self.num_buckets, self.embedding_dim) | |
| for _ in range(table_count) | |
| ) | |
| memory_width = table_count * self.embedding_dim | |
| self.key_proj = nn.Linear(memory_width, config.hidden_size, bias=False) | |
| self.value_proj = nn.Linear(memory_width, config.hidden_size, bias=False) | |
| else: | |
| self.tables = nn.ModuleList() | |
| self.key_proj = None | |
| self.value_proj = None | |
| def _hash_indices( | |
| self, | |
| input_ids: Tensor, | |
| order: int, | |
| table_index: int, | |
| document_ids: Tensor | None = None, | |
| ) -> tuple[Tensor, Tensor]: | |
| multiplier = 1_000_003 + table_index * 104_729 | |
| coefficients = torch.arange(1, order + 1, device=input_ids.device, dtype=torch.long) | |
| coefficients = coefficients * multiplier | |
| if self.hash_version == "polynomial_v2": | |
| prime = 2**31 - 1 | |
| base = self.hash_bases[table_index] | |
| padded = F.pad(input_ids, (order - 1, 0), value=0) | |
| windows = padded.unfold(1, order, 1) | |
| hashed = torch.zeros_like(input_ids) | |
| for position in range(order): | |
| hashed = (hashed * base + windows[..., position] + 1) % prime | |
| hashed = hashed % self.num_buckets | |
| elif order in (2, 3) and input_ids.size(1) >= order: | |
| hashed = input_ids * coefficients[-1] | |
| for lag in range(1, order): | |
| shifted = F.pad(input_ids[:, :-lag], (lag, 0), value=0) | |
| hashed = hashed + shifted * coefficients[-lag - 1] | |
| else: | |
| padded = F.pad(input_ids, (order - 1, 0), value=0) | |
| windows = padded.unfold(1, order, 1) | |
| hashed = (windows * coefficients).sum(dim=-1) | |
| if self.hash_version == "legacy_v1": | |
| hashed = (hashed + (table_index + 1) * 7_919) % self.num_buckets | |
| positions = torch.arange(input_ids.size(1), device=input_ids.device) | |
| valid = positions.unsqueeze(0) >= order - 1 | |
| if document_ids is not None: | |
| if order in (2, 3) and input_ids.size(1) >= order: | |
| for lag in range(1, order): | |
| shifted_documents = F.pad(document_ids[:, :-lag], (lag, 0), value=-1) | |
| valid = valid & document_ids.eq(shifted_documents) | |
| else: | |
| padded_documents = F.pad(document_ids, (order - 1, 0), value=-1) | |
| document_windows = padded_documents.unfold(1, order, 1) | |
| valid = valid & document_windows.eq(document_windows[..., :1]).all(dim=-1) | |
| return hashed, valid | |
| def _rms_normalize(self, hidden_states: Tensor) -> Tensor: | |
| variance = hidden_states.to(torch.float32).pow(2).mean(dim=-1, keepdim=True) | |
| return hidden_states * torch.rsqrt(variance + self.eps).to(hidden_states.dtype) | |
| def _forward_flat( | |
| self, | |
| input_ids: Tensor, | |
| hidden_states: Tensor, | |
| token_embeddings: Tensor | None = None, | |
| document_ids: Tensor | None = None, | |
| ) -> Tensor: | |
| if self.mode == "parameter_free": | |
| if token_embeddings is None: | |
| raise ValueError("parameter-free ngram memory requires token embeddings") | |
| embeddings = token_embeddings | |
| outputs = [] | |
| for order in self.orders: | |
| padded = F.pad(embeddings, (0, 0, order - 1, 0)) | |
| windows = padded.unfold(1, order, 1) | |
| outputs.append(windows.mean(dim=-1)) | |
| result = torch.stack(outputs, dim=0).mean(dim=0) | |
| positions = torch.arange(input_ids.size(1), device=input_ids.device) | |
| valid = positions.unsqueeze(0) >= max(self.orders) - 1 | |
| if document_ids is not None: | |
| padded_documents = F.pad( | |
| document_ids, (max(self.orders) - 1, 0), value=-1 | |
| ) | |
| document_windows = padded_documents.unfold(1, max(self.orders), 1) | |
| valid = valid & document_windows.eq(document_windows[..., :1]).all(dim=-1) | |
| return result * valid.unsqueeze(-1).to(result.dtype) * self.composition_scale | |
| retrieved = [] | |
| table_index = 0 | |
| for order in self.orders: | |
| for _ in range(self.num_hash_heads): | |
| indices, valid = self._hash_indices( | |
| input_ids, order, table_index, document_ids=document_ids | |
| ) | |
| values = self.tables[table_index](indices) | |
| retrieved.append(values * valid.unsqueeze(-1).to(values.dtype)) | |
| table_index += 1 | |
| memory = torch.cat(retrieved, dim=-1) | |
| assert self.key_proj is not None and self.value_proj is not None | |
| key = self.key_proj(memory) | |
| value = self.value_proj(memory) | |
| query_norm = self._rms_normalize(hidden_states) | |
| key_norm = self._rms_normalize(key) | |
| gate = torch.sigmoid((query_norm * key_norm).sum(dim=-1) / self.key_proj.out_features**0.5) | |
| return value * gate.unsqueeze(-1) | |
| def forward( | |
| self, | |
| input_ids: Tensor, | |
| hidden_states: Tensor, | |
| token_embeddings: Tensor | None = None, | |
| document_ids: Tensor | None = None, | |
| ) -> Tensor: | |
| """Return a residual update with the same shape as ``hidden_states``.""" | |
| if input_ids.dim() != 2: | |
| raise ValueError("ngram memory input_ids must have shape [batch, sequence]") | |
| if hidden_states.dim() == 3: | |
| if hidden_states.shape[:2] != input_ids.shape: | |
| raise ValueError("ngram memory inputs must have matching batch and sequence") | |
| return self._forward_flat( | |
| input_ids, hidden_states, token_embeddings, document_ids | |
| ) | |
| if hidden_states.dim() == 4: | |
| batch, sequence, streams, hidden = hidden_states.shape | |
| if (batch, sequence) != input_ids.shape: | |
| raise ValueError("ngram memory inputs must have matching batch and sequence") | |
| flat_ids = input_ids.unsqueeze(1).expand(-1, streams, -1).reshape(-1, sequence) | |
| flat_hidden = hidden_states.permute(0, 2, 1, 3).reshape(-1, sequence, hidden) | |
| flat_documents = None | |
| if document_ids is not None: | |
| if document_ids.shape != (batch, sequence): | |
| raise ValueError("document_ids must match input_ids") | |
| flat_documents = document_ids.unsqueeze(1).expand( | |
| -1, streams, -1 | |
| ).reshape(-1, sequence) | |
| flat_embeddings = None | |
| if token_embeddings is not None: | |
| if token_embeddings.shape != (batch, sequence, hidden): | |
| raise ValueError("token_embeddings must match the unstreamed hidden shape") | |
| flat_embeddings = token_embeddings.unsqueeze(1).expand( | |
| -1, streams, -1, -1 | |
| ).reshape(-1, sequence, hidden) | |
| update = self._forward_flat( | |
| flat_ids, flat_hidden, flat_embeddings, flat_documents | |
| ) | |
| return update.view(batch, streams, sequence, hidden).permute(0, 2, 1, 3) | |
| raise ValueError("ngram memory hidden_states must have shape [batch, sequence, hidden]") | |
| class ManifoldHyperConnection(nn.Module): | |
| """Dynamic manifold-constrained routing around one residual sublayer. | |
| The residual state has shape ``[batch, sequence, streams, hidden]``. The | |
| routing matrix is projected onto the Birkhoff polytope with Sinkhorn-Knopp | |
| iterations, while sigmoid-constrained read/write weights collapse the | |
| streams for the inner sublayer and place its output back onto the streams. | |
| """ | |
| def __init__(self, config: MicroLoopConfig) -> None: | |
| super().__init__() | |
| self.multiplier = config.mhc_multiplier | |
| self.sinkhorn_iterations = config.mhc_sinkhorn_iterations | |
| self.eps = config.mhc_eps | |
| self.input_norm = UnweightedRMSNorm(config.rms_norm_eps) | |
| mix = (2 + self.multiplier) * self.multiplier | |
| self.fn = nn.Parameter(torch.empty(mix, self.multiplier * config.hidden_size)) | |
| self.base = nn.Parameter(torch.empty(mix)) | |
| self.scale = nn.Parameter(torch.empty(3)) | |
| def forward(self, hidden_streams: Tensor) -> tuple[Tensor, Tensor, Tensor]: | |
| streams = self.multiplier | |
| flat = self.input_norm(hidden_streams.flatten(start_dim=2).float()) | |
| raw = F.linear(flat, self.fn.float()) | |
| pre_raw, post_raw, residual_raw = raw.split( | |
| [streams, streams, streams * streams], dim=-1 | |
| ) | |
| pre_base, post_base, residual_base = self.base.float().split( | |
| [streams, streams, streams * streams] | |
| ) | |
| pre_scale, post_scale, residual_scale = self.scale.float().unbind(0) | |
| pre = torch.sigmoid(pre_raw * pre_scale + pre_base) + self.eps | |
| post = 2.0 * torch.sigmoid(post_raw * post_scale + post_base) | |
| residual_logits = ( | |
| residual_raw.view(*residual_raw.shape[:-1], streams, streams) * residual_scale | |
| + residual_base.view(streams, streams) | |
| ) | |
| residual = torch.softmax(residual_logits, dim=-1) + self.eps | |
| residual = residual / (residual.sum(dim=-2, keepdim=True) + self.eps) | |
| for _ in range(self.sinkhorn_iterations - 1): | |
| residual = residual / (residual.sum(dim=-1, keepdim=True) + self.eps) | |
| residual = residual / (residual.sum(dim=-2, keepdim=True) + self.eps) | |
| collapsed = (pre.unsqueeze(-1) * hidden_streams.float()).sum(dim=2) | |
| return ( | |
| post.to(hidden_streams.dtype), | |
| residual.to(hidden_streams.dtype), | |
| collapsed.to(hidden_streams.dtype), | |
| ) | |
| def merge( | |
| hidden_streams: Tensor, sublayer_output: Tensor, post: Tensor, residual: Tensor | |
| ) -> Tensor: | |
| mixed = torch.matmul(residual, hidden_streams) | |
| return mixed + post.unsqueeze(-1) * sublayer_output.unsqueeze(2) | |
| class ManifoldHyperHead(nn.Module): | |
| """Dynamically collapse the final mHC residual streams back to model width.""" | |
| def __init__(self, config: MicroLoopConfig) -> None: | |
| super().__init__() | |
| self.multiplier = config.mhc_multiplier | |
| self.eps = config.mhc_eps | |
| self.input_norm = UnweightedRMSNorm(config.rms_norm_eps) | |
| self.fn = nn.Parameter( | |
| torch.empty(self.multiplier, self.multiplier * config.hidden_size) | |
| ) | |
| self.base = nn.Parameter(torch.empty(self.multiplier)) | |
| self.scale = nn.Parameter(torch.empty(1)) | |
| def forward(self, hidden_streams: Tensor) -> Tensor: | |
| flat = self.input_norm(hidden_streams.flatten(start_dim=2).float()) | |
| raw = F.linear(flat, self.fn.float()) | |
| pre = torch.sigmoid(raw * self.scale.float() + self.base.float()) + self.eps | |
| return (pre.unsqueeze(-1) * hidden_streams.float()).sum(dim=2).to(hidden_streams.dtype) | |
| def _rotate_half(x: Tensor) -> Tensor: | |
| x_even = x[..., ::2] | |
| x_odd = x[..., 1::2] | |
| return torch.stack((-x_odd, x_even), dim=-1).flatten(-2) | |
| def _rope_tables(max_position: int, head_dim: int, theta: float) -> tuple[Tensor, Tensor]: | |
| """Build interleaved rotary tables once, in float32 for stable reuse.""" | |
| inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim)) | |
| positions = torch.arange(max_position, device=inv_freq.device, dtype=torch.float32) | |
| angles = positions.unsqueeze(-1) * inv_freq | |
| angles = torch.stack((angles, angles), dim=-1).flatten(-2) | |
| return angles.cos(), angles.sin() | |
| def _apply_rope(q: Tensor, k: Tensor, cos: Tensor, sin: Tensor) -> tuple[Tensor, Tensor]: | |
| """Apply cached interleaved rotary embeddings to query and key tensors.""" | |
| return q * cos + _rotate_half(q) * sin, k * cos + _rotate_half(k) * sin | |
| class GroupedQueryAttention(nn.Module): | |
| """Grouped-query self-attention with explicit Q/K/V projections. | |
| Attention uses an explicit scaled-dot-product implementation (matmul + | |
| softmax) rather than ``F.scaled_dot_product_attention`` because the fused | |
| kernels select implementations based on process-level state and produce | |
| deterministic but context-dependent results: evaluations inside the | |
| training process then disagree with evaluations of the same saved | |
| checkpoint in a fresh process (the 2026-08-05 provenance incident). The | |
| explicit math path is deterministic everywhere at the cost of a small | |
| amount of speed, which is acceptable at this model size. | |
| """ | |
| def __init__(self, config: MicroLoopConfig) -> None: | |
| super().__init__() | |
| self.num_heads = config.num_attention_heads | |
| self.num_key_value_heads = config.num_key_value_heads | |
| self.head_dim = config.head_dimension | |
| self.num_groups = self.num_heads // self.num_key_value_heads | |
| self.q_proj = nn.Linear(config.hidden_size, self.num_heads * self.head_dim, bias=False) | |
| self.k_proj = nn.Linear( | |
| config.hidden_size, self.num_key_value_heads * self.head_dim, bias=False | |
| ) | |
| self.v_proj = nn.Linear( | |
| config.hidden_size, self.num_key_value_heads * self.head_dim, bias=False | |
| ) | |
| self.o_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=False) | |
| self.output_gate = ( | |
| nn.Linear(config.hidden_size, self.num_heads, bias=False) | |
| if config.attention_output_gate | |
| else None | |
| ) | |
| self.output_gate_activation = config.attention_output_gate_activation | |
| self.rope_theta = config.rope_theta | |
| self.attention_implementation = config.attention_implementation | |
| self.qk_norm_position = config.qk_norm_position | |
| if config.qk_norm == "per_head": | |
| # One affine scale is shared by all Q heads and one by all K heads; | |
| # normalization itself is applied independently over each head. | |
| self.q_norm: RMSNorm | None = RMSNorm(self.head_dim, config.rms_norm_eps) | |
| self.k_norm: RMSNorm | None = RMSNorm(self.head_dim, config.rms_norm_eps) | |
| else: | |
| self.q_norm = None | |
| self.k_norm = None | |
| def forward( | |
| self, | |
| hidden_states: Tensor, | |
| attention_mask: Tensor | None = None, | |
| position_ids: Tensor | None = None, | |
| rope_embeddings: tuple[Tensor, Tensor] | None = None, | |
| first_value_states: Tensor | None = None, | |
| value_residual_scales: Tensor | None = None, | |
| ) -> Tensor: | |
| batch, sequence, _ = hidden_states.shape | |
| q = self.q_proj(hidden_states).view(batch, sequence, self.num_heads, self.head_dim) | |
| k = self.k_proj(hidden_states).view( | |
| batch, sequence, self.num_key_value_heads, self.head_dim | |
| ) | |
| v = self.v_proj(hidden_states).view( | |
| batch, sequence, self.num_key_value_heads, self.head_dim | |
| ) | |
| q = q.transpose(1, 2) | |
| k = k.transpose(1, 2) | |
| v = v.transpose(1, 2) | |
| if first_value_states is not None: | |
| if value_residual_scales is None: | |
| raise ValueError("value residual states require value_residual_scales") | |
| if first_value_states.shape != v.shape: | |
| raise ValueError("first_value_states must match the current value-state shape") | |
| v = value_residual_scales[0] * v + value_residual_scales[1] * first_value_states | |
| if position_ids is None: | |
| position_ids = torch.arange(sequence, device=hidden_states.device).expand(batch, -1) | |
| if position_ids.shape != (batch, sequence): | |
| raise ValueError( | |
| f"position_ids must have shape [batch, sequence], got {tuple(position_ids.shape)}" | |
| ) | |
| if rope_embeddings is None: | |
| table_cos, table_sin = _rope_tables(sequence, self.head_dim, self.rope_theta) | |
| cos = table_cos[position_ids].unsqueeze(1).to(q.dtype) | |
| sin = table_sin[position_ids].unsqueeze(1).to(q.dtype) | |
| else: | |
| cos, sin = rope_embeddings | |
| if self.q_norm is not None and self.qk_norm_position == "pre_rope": | |
| # Normalize in the unrotated head basis. RMS normalization with | |
| # learned per-coordinate scales does not commute with RoPE; doing | |
| # it first preserves RoPE's relative-position shift equivariance. | |
| q = self.q_norm(q) | |
| assert self.k_norm is not None | |
| k = self.k_norm(k) | |
| q, k = _apply_rope(q, k, cos, sin) | |
| if self.q_norm is not None and self.qk_norm_position == "post_rope": | |
| q = self.q_norm(q) | |
| assert self.k_norm is not None | |
| k = self.k_norm(k) | |
| if self.num_heads % self.num_key_value_heads != 0: | |
| raise ValueError( | |
| f"num_heads {self.num_heads} must divide num_key_value_heads " | |
| f"{self.num_key_value_heads}" | |
| ) | |
| if attention_mask is None: | |
| visible: Tensor | None = None | |
| is_causal = True | |
| else: | |
| # A 2-D mask uses the conventional HF meaning: one means visible. | |
| if attention_mask.shape == (batch, sequence): | |
| causal = torch.tril( | |
| torch.ones(sequence, sequence, device=q.device, dtype=torch.bool) | |
| ) | |
| visible = causal.unsqueeze(0).unsqueeze(0) & attention_mask.bool().unsqueeze( | |
| 1 | |
| ).unsqueeze(2) | |
| is_causal = False | |
| elif attention_mask.shape == (batch, sequence, sequence): | |
| visible = attention_mask.bool().unsqueeze(1) | |
| is_causal = False | |
| else: | |
| raise ValueError( | |
| "attention_mask must have shape [batch, sequence] or " | |
| "[batch, sequence, sequence], " | |
| f"got {tuple(attention_mask.shape)}" | |
| ) | |
| # GQA: repeat the KV heads so every query head has its own K/V. | |
| repeat = self.num_heads // self.num_key_value_heads | |
| if repeat > 1: | |
| k = k.repeat_interleave(repeat, dim=1) | |
| v = v.repeat_interleave(repeat, dim=1) | |
| if self.attention_implementation == "sdpa": | |
| attended = F.scaled_dot_product_attention( | |
| q, k, v, attn_mask=visible, dropout_p=0.0, is_causal=is_causal | |
| ) | |
| else: | |
| scores = torch.matmul(q, k.transpose(-2, -1)) / float(self.head_dim) ** 0.5 | |
| if is_causal: | |
| seq_ids = torch.arange(sequence, device=scores.device) | |
| visible = seq_ids.unsqueeze(0) <= seq_ids.unsqueeze(1) | |
| visible = visible.expand(batch, self.num_heads, sequence, sequence) | |
| if visible is not None: | |
| scores = scores.masked_fill(~visible, float("-inf")) | |
| probs = torch.softmax(scores, dim=-1) | |
| attended = torch.matmul(probs, v) | |
| if self.output_gate is not None: | |
| gate_logits = self.output_gate(hidden_states) | |
| gate = ( | |
| torch.sigmoid(gate_logits) | |
| if self.output_gate_activation == "sigmoid" | |
| else F.silu(gate_logits) | |
| ) | |
| gate = gate.transpose(1, 2).unsqueeze(-1) | |
| attended = attended * gate | |
| attended = attended.transpose(1, 2).contiguous().view(batch, sequence, -1) | |
| return self.o_proj(attended) | |
| class LowRankFFNProjection(nn.Module): | |
| """CoLA-inspired B(activation(A(x))); outer SwiGLU remains in the block.""" | |
| def __init__(self, in_features: int, out_features: int, rank: int, activation: str): | |
| super().__init__() | |
| self.reduce = nn.Linear(in_features, rank, bias=False) | |
| self.expand = nn.Linear(rank, out_features, bias=False) | |
| self.activation = nn.SiLU() if activation == "silu" else nn.Identity() | |
| def forward(self, hidden_states: Tensor) -> Tensor: | |
| return self.expand(self.activation(self.reduce(hidden_states))) | |
| class MicroLoopBlock(nn.Module): | |
| """Pre-norm transformer block with SwiGLU feed-forward network.""" | |
| def __init__(self, config: MicroLoopConfig, layer_index: int) -> None: | |
| super().__init__() | |
| self.attn_norm = RMSNorm(config.hidden_size, config.rms_norm_eps) | |
| self.attn = GroupedQueryAttention(config) | |
| self.ffn_norm = RMSNorm(config.hidden_size, config.rms_norm_eps) | |
| if config.ffn_rank is None: | |
| self.ffn_gate = nn.Linear(config.hidden_size, config.intermediate_size, bias=False) | |
| self.ffn_up = nn.Linear(config.hidden_size, config.intermediate_size, bias=False) | |
| self.ffn_down = nn.Linear(config.intermediate_size, config.hidden_size, bias=False) | |
| else: | |
| self.ffn_gate = LowRankFFNProjection( | |
| config.hidden_size, config.intermediate_size, | |
| config.ffn_rank, config.ffn_factor_activation, | |
| ) | |
| self.ffn_up = LowRankFFNProjection( | |
| config.hidden_size, config.intermediate_size, | |
| config.ffn_rank, config.ffn_factor_activation, | |
| ) | |
| self.ffn_down = LowRankFFNProjection( | |
| config.intermediate_size, config.hidden_size, | |
| config.ffn_rank, config.ffn_factor_activation, | |
| ) | |
| self.use_mhc = config.mhc_multiplier > 1 | |
| if self.use_mhc: | |
| self.attn_mhc = ManifoldHyperConnection(config) | |
| self.ffn_mhc = ManifoldHyperConnection(config) | |
| self.use_attn_residuals = config.attn_res_block_size is not None | |
| self.value_residual_scales = ( | |
| nn.Parameter(torch.tensor([1.0, 0.0])) | |
| if config.value_residual.get("enabled", False) and layer_index > 0 | |
| else None | |
| ) | |
| if self.use_attn_residuals: | |
| self.attn_res_norm = RMSNorm(config.hidden_size, config.rms_norm_eps) | |
| self.ffn_res_norm = RMSNorm(config.hidden_size, config.rms_norm_eps) | |
| self.attn_res_proj = nn.Linear(config.hidden_size, 1, bias=False) | |
| self.ffn_res_proj = nn.Linear(config.hidden_size, 1, bias=False) | |
| clamp = dict(config.swiglu_clamp) | |
| self.swiglu_clamp_enabled = bool(clamp.get("enabled", False)) | |
| self.swiglu_linear_min = float(clamp.get("linear_min", -10.0)) | |
| self.swiglu_linear_max = float(clamp.get("linear_max", 10.0)) | |
| self.swiglu_gate_max = float(clamp.get("gate_max", 10.0)) | |
| def forward( | |
| self, | |
| hidden_states: Tensor, | |
| attention_mask: Tensor | None = None, | |
| position_ids: Tensor | None = None, | |
| rope_embeddings: tuple[Tensor, Tensor] | None = None, | |
| block_residuals: list[Tensor] | None = None, | |
| first_value_states: Tensor | None = None, | |
| ) -> Tensor: | |
| if self.use_mhc: | |
| attn_post, attn_residual, attn_input = self.attn_mhc(hidden_states) | |
| attn_output = self.attn( | |
| self.attn_norm(attn_input), | |
| attention_mask, | |
| position_ids, | |
| rope_embeddings, | |
| first_value_states=first_value_states, | |
| value_residual_scales=self.value_residual_scales, | |
| ) | |
| hidden_states = self.attn_mhc.merge( | |
| hidden_states, attn_output, attn_post, attn_residual | |
| ) | |
| ffn_post, ffn_residual, ffn_input = self.ffn_mhc(hidden_states) | |
| normalized = self.ffn_norm(ffn_input) | |
| gate_linear = self.ffn_gate(normalized) | |
| up_linear = self.ffn_up(normalized) | |
| if self.swiglu_clamp_enabled: | |
| gate_linear = gate_linear.clamp(self.swiglu_linear_min, self.swiglu_linear_max) | |
| up_linear = up_linear.clamp(self.swiglu_linear_min, self.swiglu_linear_max) | |
| gate = F.silu(gate_linear) | |
| if self.swiglu_clamp_enabled: | |
| gate = gate.clamp(max=self.swiglu_gate_max) | |
| ffn_output = self.ffn_down(gate * up_linear) | |
| return self.ffn_mhc.merge(hidden_states, ffn_output, ffn_post, ffn_residual) | |
| if self.use_attn_residuals and block_residuals: | |
| residual_stack = torch.stack(block_residuals, dim=-2) | |
| scores = torch.cat( | |
| [self.attn_res_proj(self.attn_res_norm(state)) for state in block_residuals], dim=-1 | |
| ) | |
| hidden_states = hidden_states + ( | |
| torch.softmax(scores, dim=-1).unsqueeze(-1) * residual_stack | |
| ).sum(dim=-2) | |
| hidden_states = hidden_states + self.attn( | |
| self.attn_norm(hidden_states), | |
| attention_mask, | |
| position_ids, | |
| rope_embeddings, | |
| first_value_states=first_value_states, | |
| value_residual_scales=self.value_residual_scales, | |
| ) | |
| if self.use_attn_residuals and block_residuals: | |
| residual_stack = torch.stack(block_residuals, dim=-2) | |
| scores = torch.cat( | |
| [self.ffn_res_proj(self.ffn_res_norm(state)) for state in block_residuals], dim=-1 | |
| ) | |
| hidden_states = hidden_states + ( | |
| torch.softmax(scores, dim=-1).unsqueeze(-1) * residual_stack | |
| ).sum(dim=-2) | |
| ffn_input = self.ffn_norm(hidden_states) | |
| gate_linear = self.ffn_gate(ffn_input) | |
| up_linear = self.ffn_up(ffn_input) | |
| if self.swiglu_clamp_enabled: | |
| gate_linear = gate_linear.clamp(self.swiglu_linear_min, self.swiglu_linear_max) | |
| up_linear = up_linear.clamp(self.swiglu_linear_min, self.swiglu_linear_max) | |
| gate = F.silu(gate_linear) | |
| if self.swiglu_clamp_enabled: | |
| gate = gate.clamp(max=self.swiglu_gate_max) | |
| ffn_output = self.ffn_down(gate * up_linear) | |
| return hidden_states + ffn_output | |
| class MicroLoopPreTrainedModel(PreTrainedModel): | |
| config_class = MicroLoopConfig | |
| base_model_prefix = "microloop" | |
| class MicroLoopForDiffusionLM(MicroLoopPreTrainedModel, GenerationMixin): | |
| """Backbone plus tied output head for causal bootstrap and diffusion training.""" | |
| _tied_weights_keys = {"lm_head.weight": "embed_tokens.weight"} | |
| def __init__(self, config: MicroLoopConfig) -> None: | |
| config.validate() | |
| super().__init__(config) | |
| self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size) | |
| self.layers = nn.ModuleList( | |
| [MicroLoopBlock(config, index) for index in range(config.num_hidden_layers)] | |
| ) | |
| looping = config.looping | |
| self.loop_gates = ( | |
| nn.Parameter( | |
| torch.zeros( | |
| max( | |
| 0, | |
| int( | |
| looping.get( | |
| "max_loop_count", looping.get("maximum_serving_loops", 3) | |
| ) | |
| ) | |
| - 1, | |
| ), | |
| config.hidden_size, | |
| ) | |
| ) | |
| if looping.get("mode", "layer") == "block" and looping.get("gated", False) | |
| else None | |
| ) | |
| self.ngram_memory = ( | |
| NGramMemory(config) if config.ngram_memory.get("enabled", False) else None | |
| ) | |
| self.ngram_memory_layer = ( | |
| int(config.ngram_memory["insertion_layer"]) | |
| if self.ngram_memory is not None | |
| else None | |
| ) | |
| digit_settings = config.digit_position_embedding | |
| self.digit_position_embedding = ( | |
| nn.Embedding( | |
| int(digit_settings.get("max_positions", 128)), config.hidden_size | |
| ) | |
| if digit_settings.get("enabled", False) | |
| else None | |
| ) | |
| self.digit_token_ids = frozenset( | |
| int(token_id) for token_id in digit_settings.get("digit_token_ids", []) | |
| ) | |
| self.register_buffer( | |
| "digit_token_ids_tensor", | |
| torch.tensor(sorted(self.digit_token_ids), dtype=torch.long), | |
| persistent=False, | |
| ) | |
| self.mhc_head = ManifoldHyperHead(config) if config.mhc_multiplier > 1 else None | |
| self.norm = RMSNorm(config.hidden_size, config.rms_norm_eps) | |
| self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) | |
| self.mtp_proj = ( | |
| nn.Linear(config.hidden_size, config.hidden_size, bias=False) | |
| if config.mtp_enabled | |
| else None | |
| ) | |
| rope_cos, rope_sin = _rope_tables( | |
| config.max_position_embeddings, config.head_dimension, config.rope_theta | |
| ) | |
| # PERSISTENT buffers on purpose: transformers 5.x ``from_pretrained`` | |
| # re-initializes non-persistent buffers that are missing from the | |
| # checkpoint (``_initialize_missing_keys`` -> ``initialize_weights``), | |
| # overwriting the rotary tables with garbage. That made every fresh- | |
| # process evaluation of saved checkpoints compute with corrupted rope | |
| # tables (the historical "external collapse" at ~0.03 accuracy was this | |
| # artifact). Persisting the tables makes the saved checkpoint carry | |
| # exactly the tables used in training. | |
| self.register_buffer("rope_cos", rope_cos, persistent=True) | |
| self.register_buffer("rope_sin", rope_sin, persistent=True) | |
| if config.tie_word_embeddings: | |
| self.lm_head.weight = self.embed_tokens.weight | |
| # Canonical HF pattern: post_init() installs all_tied_weights_keys and | |
| # dispatches _initialize_weights per module. | |
| self.post_init() | |
| if self.digit_position_embedding is not None: | |
| nn.init.zeros_(self.digit_position_embedding.weight) | |
| if self.ngram_memory is not None and self.ngram_memory.value_proj is not None: | |
| nn.init.zeros_(self.ngram_memory.value_proj.weight) | |
| def _reset_rope_buffers(self) -> None: | |
| """Recompute the rotary tables from the config (they are deterministic).""" | |
| rope_cos, rope_sin = _rope_tables( | |
| self.config.max_position_embeddings, | |
| self.config.head_dimension, | |
| self.config.rope_theta, | |
| ) | |
| self.rope_cos.copy_(rope_cos) | |
| self.rope_sin.copy_(rope_sin) | |
| def _initialize_weights(self, module: nn.Module, is_custom_code: bool = False) -> None: | |
| # The caller controls the RNG through seed_everything; this method performs | |
| # no hidden reseeding and is therefore reproducible by construction. | |
| if getattr(module, "_is_hf_initialized", False): | |
| return | |
| if module is self: | |
| # The main module owns the rotary tables. transformers 5.x | |
| # re-initializes buffers that are missing from a loaded checkpoint | |
| # ("_initialize_missing_keys"), which zeroes/garbles the tables; | |
| # restore the deterministic canonical tables here instead. | |
| self._reset_rope_buffers() | |
| if isinstance(module, (nn.Linear, nn.Embedding)): | |
| nn.init.normal_(module.weight, mean=0.0, std=0.02) | |
| if getattr(module, "bias", None) is not None: | |
| nn.init.zeros_(module.bias) | |
| elif isinstance(module, ManifoldHyperConnection): | |
| nn.init.normal_(module.fn, mean=0.0, std=0.02) | |
| nn.init.zeros_(module.base) | |
| nn.init.constant_(module.scale, self.config.mhc_init_scale) | |
| elif isinstance(module, ManifoldHyperHead): | |
| nn.init.normal_(module.fn, mean=0.0, std=0.02) | |
| nn.init.zeros_(module.base) | |
| nn.init.constant_(module.scale, self.config.mhc_init_scale) | |
| elif isinstance(module, RMSNorm): | |
| nn.init.ones_(module.weight) | |
| def get_input_embeddings(self) -> nn.Embedding: | |
| return self.embed_tokens | |
| def get_output_embeddings(self) -> nn.Linear: | |
| return self.lm_head | |
| def set_input_embeddings(self, value: nn.Embedding) -> None: | |
| self.embed_tokens = value | |
| if self.config.tie_word_embeddings: | |
| self.lm_head.weight = self.embed_tokens.weight | |
| def digit_position_ids( | |
| self, | |
| input_ids: Tensor, | |
| document_ids: Tensor | None = None, | |
| prior_run_length: Tensor | None = None, | |
| prior_active: Tensor | None = None, | |
| ) -> tuple[Tensor, Tensor, Tensor]: | |
| """Return zero-based contiguous digit positions and streaming state. | |
| Positions reset after non-digit tokens and document boundaries. The two | |
| state tensors contain the run length and whether the final token is a | |
| digit, allowing the cache to continue a number across blocks. | |
| """ | |
| batch, sequence = input_ids.shape | |
| device = input_ids.device | |
| enabled = self.digit_position_embedding is not None | |
| if not enabled: | |
| zeros = torch.zeros((batch, sequence), dtype=torch.long, device=device) | |
| return zeros, zeros[:, -1], zeros[:, -1].bool() | |
| is_digit = torch.isin(input_ids, self.digit_token_ids_tensor) | |
| positions = torch.zeros_like(input_ids, dtype=torch.long) | |
| run = ( | |
| prior_run_length.to(device=device, dtype=torch.long) | |
| if prior_run_length is not None | |
| else torch.zeros(batch, dtype=torch.long, device=device) | |
| ) | |
| active = ( | |
| prior_active.to(device=device, dtype=torch.bool) | |
| if prior_active is not None | |
| else torch.zeros(batch, dtype=torch.bool, device=device) | |
| ) | |
| previous_digit = F.pad(is_digit[:, :-1], (1, 0), value=False) | |
| same_document = torch.ones_like(is_digit) | |
| if document_ids is not None: | |
| previous_document = F.pad(document_ids[:, :-1], (1, 0), value=-1) | |
| same_document = document_ids.eq(previous_document) | |
| same_document[:, 0] = False | |
| continuation = is_digit & previous_digit & same_document | |
| if prior_active is not None: | |
| continuation[:, 0] = is_digit[:, 0] & active | |
| starts = is_digit & ~continuation | |
| max_positions = self.digit_position_embedding.num_embeddings | |
| base = sequence + max_positions + 1 | |
| markers = torch.where(starts, torch.arange(sequence, device=device) + base, 0) | |
| if prior_active is not None: | |
| prior_start = base - run | |
| markers[:, 0] = torch.where( | |
| is_digit[:, 0] & active, prior_start, markers[:, 0] | |
| ) | |
| last_start = torch.cummax(markers, dim=1).values - base | |
| positions = (torch.arange(sequence, device=device) - last_start).clamp_min(0) | |
| positions = positions.expand(batch, -1).clamp_max(max_positions - 1) | |
| positions = torch.where(is_digit, positions, torch.zeros_like(positions)) | |
| run = torch.where(is_digit, positions + 1, 0)[:, -1] | |
| active = is_digit[:, -1] | |
| return positions, run, active | |
| def add_digit_position_embeddings( | |
| self, | |
| input_ids: Tensor, | |
| token_embeddings: Tensor, | |
| document_ids: Tensor | None = None, | |
| prior_run_length: Tensor | None = None, | |
| prior_active: Tensor | None = None, | |
| ) -> tuple[Tensor, Tensor, Tensor]: | |
| positions, run, active = self.digit_position_ids( | |
| input_ids, document_ids=document_ids, prior_run_length=prior_run_length | |
| , prior_active=prior_active | |
| ) | |
| if self.digit_position_embedding is not None: | |
| is_digit = torch.isin(input_ids, self.digit_token_ids_tensor) | |
| token_embeddings = token_embeddings + self.digit_position_embedding( | |
| positions | |
| ) * is_digit.unsqueeze(-1) | |
| return token_embeddings, run, active | |
| def forward( | |
| self, | |
| input_ids: Tensor | None = None, | |
| inputs_embeds: Tensor | None = None, | |
| attention_mask: Tensor | None = None, | |
| document_ids: Tensor | None = None, | |
| prefix_lengths: Tensor | None = None, | |
| position_ids: Tensor | None = None, | |
| labels: Tensor | None = None, | |
| logit_mask: Tensor | None = None, | |
| mtp_loss_weight: float = 0.0, | |
| use_cut_cross_entropy: bool = False, | |
| loop_count: int = 1, | |
| output_hidden_states: bool = False, | |
| **_: object, | |
| ) -> MicroLoopCausalLMOutput: | |
| if (input_ids is None) == (inputs_embeds is None): | |
| raise ValueError("exactly one of input_ids or inputs_embeds must be provided") | |
| if loop_count < 1: | |
| raise ValueError("loop_count must be at least one") | |
| if use_cut_cross_entropy and labels is None: | |
| raise ValueError("use_cut_cross_entropy requires labels") | |
| if input_ids is not None: | |
| if input_ids.dim() != 2: | |
| raise ValueError( | |
| f"input_ids must have shape [batch, sequence], got {input_ids.dim()}-D" | |
| ) | |
| batch, sequence = input_ids.shape | |
| else: | |
| if inputs_embeds is None or inputs_embeds.dim() != 3: | |
| raise ValueError( | |
| "inputs_embeds must have shape [batch, sequence, hidden], got " | |
| f"{None if inputs_embeds is None else inputs_embeds.dim()}-D" | |
| ) | |
| batch, sequence, _ = inputs_embeds.shape | |
| if self.ngram_memory is not None and input_ids is None: | |
| raise ValueError("ngram memory requires input_ids for deterministic suffix lookup") | |
| device = (input_ids if input_ids is not None else inputs_embeds).device | |
| if document_ids is not None: | |
| if document_ids.shape != (batch, sequence): | |
| raise ValueError("document_ids must have the same shape as the input") | |
| if attention_mask is not None and attention_mask.dim() != 2: | |
| raise ValueError( | |
| "document_ids cannot be combined with a precomputed attention mask" | |
| ) | |
| valid = ( | |
| attention_mask.bool() | |
| if attention_mask is not None | |
| else torch.ones((batch, sequence), dtype=torch.bool, device=device) | |
| ) | |
| if prefix_lengths is None: | |
| causal = torch.tril(torch.ones(sequence, sequence, dtype=torch.bool, device=device)) | |
| attention_mask = ( | |
| causal.unsqueeze(0) | |
| & document_ids.unsqueeze(2).eq(document_ids.unsqueeze(1)) | |
| & valid.unsqueeze(1) | |
| & valid.unsqueeze(2) | |
| ) | |
| else: | |
| attention_mask = build_prefix_lm_mask( | |
| prefix_lengths.to(device), | |
| sequence, | |
| valid=valid, | |
| document_ids=document_ids, | |
| ) | |
| elif prefix_lengths is not None: | |
| if attention_mask is not None and attention_mask.dim() != 2: | |
| raise ValueError("prefix_lengths requires a 2-D padding mask") | |
| valid = ( | |
| attention_mask.bool() | |
| if attention_mask is not None | |
| else torch.ones((batch, sequence), dtype=torch.bool, device=device) | |
| ) | |
| attention_mask = build_prefix_lm_mask( | |
| prefix_lengths.to(device), sequence, valid=valid | |
| ) | |
| if position_ids is None: | |
| position_ids = torch.arange(sequence, device=device).expand(batch, -1) | |
| rope_dtype = ( | |
| torch.get_autocast_dtype("cuda") | |
| if device.type == "cuda" and torch.is_autocast_enabled("cuda") | |
| else self.embed_tokens.weight.dtype | |
| ) | |
| rope_embeddings = ( | |
| self.rope_cos[position_ids].unsqueeze(1).to(rope_dtype), | |
| self.rope_sin[position_ids].unsqueeze(1).to(rope_dtype), | |
| ) | |
| token_embeddings = self.embed_tokens(input_ids) if input_ids is not None else inputs_embeds | |
| if input_ids is not None: | |
| token_embeddings, _, _ = self.add_digit_position_embeddings( | |
| input_ids, token_embeddings, document_ids=document_ids | |
| ) | |
| hidden_states = token_embeddings | |
| if self.config.mhc_multiplier > 1: | |
| hidden_states = hidden_states.unsqueeze(2).expand( | |
| -1, -1, self.config.mhc_multiplier, -1 | |
| ).contiguous() | |
| first_value_states = None | |
| if self.config.value_residual.get("enabled", False): | |
| first_attention = self.layers[0].attn | |
| first_values = first_attention.v_proj(self.layers[0].attn_norm(hidden_states)) | |
| first_value_states = first_values.view( | |
| batch, | |
| sequence, | |
| first_attention.num_key_value_heads, | |
| first_attention.head_dim, | |
| ).transpose(1, 2) | |
| loop_layers = list(self.config.looping.get("layers", [4, 5, 6])) | |
| loop_mode = str(self.config.looping.get("mode", "layer")) | |
| if ( | |
| loop_mode == "block" | |
| and self.config.looping.get("gated", False) | |
| and loop_count | |
| > int( | |
| self.config.looping.get( | |
| "max_loop_count", self.config.looping.get("maximum_serving_loops", 3) | |
| ) | |
| ) | |
| ): | |
| raise ValueError("loop_count exceeds configured max_loop_count") | |
| block_residuals: list[Tensor] = [] | |
| block_size = self.config.attn_res_block_size | |
| total_applications = 0 | |
| if loop_mode == "block": | |
| valid_layers = sorted( | |
| {layer for layer in loop_layers if 1 <= int(layer) <= len(self.layers)} | |
| ) | |
| if valid_layers: | |
| loop_set = set(valid_layers) | |
| layer_order = [ | |
| *range(1, valid_layers[0]), | |
| *valid_layers * loop_count, | |
| *range(valid_layers[-1] + 1, len(self.layers) + 1), | |
| ] | |
| else: | |
| loop_set = set() | |
| layer_order = list(range(1, len(self.layers) + 1)) | |
| else: | |
| loop_set = {int(layer) for layer in loop_layers} | |
| layer_order = [ | |
| layer_number | |
| for layer_number in range(1, len(self.layers) + 1) | |
| for _ in range(loop_count if layer_number in loop_set else 1) | |
| ] | |
| block_start = min(loop_set) if loop_set else None | |
| block_end = max(loop_set) if loop_set else None | |
| block_size_layers = len(loop_set) | |
| block_pass = 0 | |
| block_input: Tensor | None = None | |
| for application_index, layer_number in enumerate(layer_order): | |
| layer = self.layers[layer_number - 1] | |
| if ( | |
| loop_mode == "block" | |
| and block_start is not None | |
| and application_index >= block_start - 1 | |
| and (application_index - (block_start - 1)) % block_size_layers == 0 | |
| ): | |
| block_input = hidden_states | |
| if block_size is not None and (layer_number - 1) % block_size == 0: | |
| block_residuals.append(hidden_states) | |
| prior_block_residuals = block_residuals[:-1] if block_size is not None else None | |
| hidden_states = layer( | |
| hidden_states, | |
| attention_mask=attention_mask, | |
| position_ids=position_ids, | |
| rope_embeddings=rope_embeddings, | |
| block_residuals=prior_block_residuals, | |
| first_value_states=first_value_states if layer_number > 1 else None, | |
| ) | |
| if self.ngram_memory is not None and layer_number == self.ngram_memory_layer: | |
| assert input_ids is not None | |
| hidden_states = hidden_states + self.ngram_memory( | |
| input_ids, | |
| hidden_states, | |
| token_embeddings=token_embeddings, | |
| document_ids=document_ids, | |
| ) | |
| total_applications += 1 | |
| if ( | |
| loop_mode == "block" | |
| and block_end is not None | |
| and layer_number == block_end | |
| and block_start is not None | |
| and (application_index - (block_start - 1) + 1) % block_size_layers == 0 | |
| ): | |
| block_pass += 1 | |
| if block_pass < loop_count: | |
| assert block_input is not None | |
| hidden_states = self._apply_loop_gate( | |
| block_input, hidden_states, block_pass - 1 | |
| ) | |
| if self.mhc_head is not None: | |
| hidden_states = self.mhc_head(hidden_states) | |
| hidden_states = self.norm(hidden_states) | |
| if use_cut_cross_entropy: | |
| if logit_mask is not None: | |
| raise ValueError("use_cut_cross_entropy cannot be combined with logit_mask") | |
| logits = None | |
| elif logit_mask is not None: | |
| # Diffusion training: only the masked positions carry loss, so the | |
| # output head runs on those rows instead of the full sequence. | |
| if labels is not None: | |
| raise ValueError("logit_mask cannot be combined with labels") | |
| if logit_mask.shape != (batch, sequence): | |
| raise ValueError("logit_mask must have the same shape as the input") | |
| flat = hidden_states.reshape(-1, self.config.hidden_size) | |
| logits = self.lm_head(flat[logit_mask.reshape(-1)]) | |
| else: | |
| logits = self.lm_head(hidden_states) | |
| loss = None | |
| if labels is not None: | |
| if labels.shape != (batch, sequence): | |
| raise ValueError("labels must have the same shape as the input") | |
| if labels.size(1) < 2: | |
| raise ValueError("causal training requires sequences with at least two tokens") | |
| # Causal next-token prediction: position t predicts the label at t + 1. | |
| if use_cut_cross_entropy: | |
| try: | |
| from cut_cross_entropy import linear_cross_entropy | |
| except ImportError as exc: # pragma: no cover - optional dependency | |
| raise RuntimeError( | |
| "Cut Cross Entropy is optional; install the fused-loss project extra" | |
| ) from exc | |
| cce_embeddings = hidden_states[:, :-1, :] | |
| if torch.is_autocast_enabled("cuda"): | |
| cce_embeddings = cce_embeddings.to(torch.get_autocast_dtype("cuda")) | |
| if cce_embeddings.dtype not in (torch.float16, torch.bfloat16): | |
| raise ValueError("Cut Cross Entropy requires CUDA fp16 or bf16 autocast") | |
| loss = linear_cross_entropy( | |
| cce_embeddings, | |
| self.lm_head.weight, | |
| labels[:, 1:], | |
| ignore_index=-100, | |
| ) | |
| else: | |
| assert logits is not None | |
| loss = F.cross_entropy( | |
| logits[:, :-1, :].reshape(-1, logits.size(-1)), | |
| labels[:, 1:].reshape(-1), | |
| ignore_index=-100, | |
| ) | |
| if mtp_loss_weight: | |
| if self.mtp_proj is None: | |
| raise ValueError("mtp_loss_weight requires mtp_enabled=true") | |
| if mtp_loss_weight < 0: | |
| raise ValueError("mtp_loss_weight must be non-negative") | |
| mtp_logits = self.lm_head(self.mtp_proj(hidden_states[:, :-2, :])) | |
| mtp_loss = F.cross_entropy( | |
| mtp_logits.reshape(-1, mtp_logits.size(-1)), | |
| labels[:, 2:].reshape(-1), | |
| ignore_index=-100, | |
| ) | |
| loss = loss + float(mtp_loss_weight) * mtp_loss | |
| return MicroLoopCausalLMOutput( | |
| logits=logits, | |
| loss=loss, | |
| hidden_states=hidden_states if output_hidden_states else None, | |
| loop_applications=total_applications, | |
| ) | |
| def _apply_loop_gate( | |
| self, previous: Tensor, candidate: Tensor, pass_index: int | |
| ) -> Tensor: | |
| """Blend a repeated shared block into the previous pass.""" | |
| if self.loop_gates is None: | |
| return candidate | |
| if pass_index < 0 or pass_index >= self.loop_gates.size(0): | |
| raise ValueError("loop_count exceeds configured max_loop_count") | |
| gate = torch.tanh(self.loop_gates[pass_index]).to(candidate.dtype) | |
| return previous + gate.view(1, 1, -1) * (candidate - previous) | |
| def generate_greedy( | |
| self, input_ids: Tensor, max_new_tokens: int, eos_token_id: int | None = None | |
| ) -> Tensor: | |
| """Small causal smoke decoder; diffusion sampling belongs in ``sampler.py``.""" | |
| return self.generate_causal( | |
| input_ids, max_new_tokens=max_new_tokens, eos_token_id=eos_token_id, do_sample=False | |
| ) | |
| def generate_causal( | |
| self, | |
| input_ids: Tensor, | |
| *, | |
| max_new_tokens: int, | |
| eos_token_id: int | None = None, | |
| do_sample: bool = False, | |
| temperature: float = 1.0, | |
| top_k: int | None = None, | |
| ) -> Tensor: | |
| """Generate a batched causal continuation for M2 validation and serving. | |
| This deliberately recomputes the context on every step. KV caching is a | |
| later optimization; keeping this reference path simple makes M2 output | |
| semantics straightforward to test. | |
| """ | |
| if max_new_tokens < 0: | |
| raise ValueError("max_new_tokens must be non-negative") | |
| if temperature <= 0: | |
| raise ValueError("temperature must be positive") | |
| if top_k is not None and top_k <= 0: | |
| raise ValueError("top_k must be positive when supplied") | |
| generated = input_ids | |
| finished = torch.zeros(input_ids.size(0), dtype=torch.bool, device=input_ids.device) | |
| for _ in range(max_new_tokens): | |
| next_logits = self(generated).logits[:, -1, :] | |
| if do_sample: | |
| next_logits = next_logits / temperature | |
| if top_k is not None and top_k < next_logits.size(-1): | |
| threshold = torch.topk(next_logits, top_k, dim=-1).values[:, -1:] | |
| next_logits = next_logits.masked_fill(next_logits < threshold, float("-inf")) | |
| next_token = torch.multinomial(torch.softmax(next_logits, dim=-1), 1) | |
| else: | |
| next_token = next_logits.argmax(dim=-1, keepdim=True) | |
| if eos_token_id is not None: | |
| next_token = torch.where( | |
| finished.unsqueeze(1), | |
| torch.full_like(next_token, eos_token_id), | |
| next_token, | |
| ) | |
| finished |= next_token.squeeze(1).eq(eos_token_id) | |
| generated = torch.cat((generated, next_token), dim=1) | |
| if eos_token_id is not None and bool(finished.all()): | |
| break | |
| return generated | |