import importlib import importlib.util import logging import math import os import sys from dataclasses import dataclass from typing import Any, Callable, Dict, List, Optional, Tuple, Union import torch from torch import Tensor import torch.nn as nn import torch.nn.functional as F import torch.utils.checkpoint as cp logging.basicConfig( level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s" ) logger = logging.getLogger("XoneLM") class HardwareContext: @staticmethod def get_optimal_device() -> torch.device: if hasattr(torch, "accelerator") and torch.accelerator.is_available(): try: acc_device = torch.accelerator.current_accelerator() if acc_device is not None: idx = ( torch.accelerator.current_device_index() if hasattr(torch.accelerator, "current_device_index") else 0 ) return torch.device(f"{acc_device.type}:{idx}") except Exception: pass if "torch_xla" in sys.modules: try: import torch_xla.core.xla_model as xm return xm.xla_device() except Exception: pass # 3. NVIDIA CUDA GPU if torch.cuda.is_available(): idx = ( torch.cuda.current_device() if hasattr(torch.cuda, "current_device") else 0 ) return torch.device(f"cuda:{idx}") # 4. Intel XPU / Apple MPS / CPU if hasattr(torch, "xpu") and torch.xpu.is_available(): idx = ( torch.xpu.current_device() if hasattr(torch.xpu, "current_device") else 0 ) return torch.device(f"xpu:{idx}") if hasattr(torch.backends, "mps") and torch.backends.mps.is_available(): return torch.device("mps") return torch.device("cpu") @staticmethod def get_optimal_autocast_dtype(device: torch.device) -> torch.dtype: dev_type = device.type if dev_type == "cuda": if torch.cuda.is_available(): major, _ = torch.cuda.get_device_capability(device) if major < 8: return torch.float16 # Ampere (sm_80), Ada (sm_89), Hopper (sm_90), Blackwell (sm_100/sm_120) if torch.cuda.is_bf16_supported(): return torch.bfloat16 return torch.float16 elif dev_type == "xla": return torch.bfloat16 elif dev_type == "xpu": if ( hasattr(torch.xpu, "is_bf16_supported") and torch.xpu.is_bf16_supported() ): return torch.bfloat16 return torch.float16 elif dev_type == "mps": return torch.float16 elif dev_type == "cpu": return torch.bfloat16 return torch.float32 @staticmethod def get_autocast_context(device: torch.device): dev_type = device.type target_dtype = HardwareContext.get_optimal_autocast_dtype(device) if dev_type in ("cuda", "cpu", "xpu"): return torch.amp.autocast( device_type=dev_type, dtype=target_dtype, enabled=(target_dtype != torch.float32), ) elif dev_type == "xla": try: return torch.amp.autocast( device_type="xla", dtype=torch.bfloat16, enabled=True ) except Exception: return torch.nullcontext() elif dev_type == "mps": try: return torch.amp.autocast( device_type="mps", dtype=torch.float16, enabled=True ) except Exception: return torch.nullcontext() return torch.nullcontext() HAS_SDPA = hasattr(F, "scaled_dot_product_attention") HAS_FLASH_ATTN = False _flash_attn_func = None for mod_name, attr_name in [ ("flash_attn", "flash_attn_func"), ("flash_attn_3.flash_attn_interface", "flash_attn_func"), ("flash_attn_4.flash_attn_interface", "flash_attn_func"), ]: try: root_pkg = mod_name.split(".")[0] if importlib.util.find_spec(root_pkg) is not None: mod = importlib.import_module(mod_name) _flash_attn_func = getattr(mod, attr_name, None) if _flash_attn_func is not None: HAS_FLASH_ATTN = True break except Exception: continue HAS_COMPILED_FLEX_ATTENTION = False _compiled_flex_attention_fn = None try: from torch.nn.attention.flex_attention import ( create_block_mask, flex_attention as _raw_flex_attn, ) if torch.cuda.is_available(): major, _ = torch.cuda.get_device_capability() if major >= 8: _compiled_flex_attention_fn = torch.compile(_raw_flex_attn, dynamic=True) HAS_COMPILED_FLEX_ATTENTION = True except Exception: HAS_COMPILED_FLEX_ATTENTION = False HAS_FUSED_LINEAR_CE = hasattr(nn, "LinearCrossEntropyLoss") or hasattr( F, "linear_cross_entropy" ) def create_universal_document_boundary_mask( x_tokens: torch.Tensor, hub_size: int, past_k_len: int, eod_token_id: int, is_dense_with_hub: bool = True, ) -> torch.Tensor: batch_size, text_len = x_tokens.shape cur_seq_len = (hub_size + text_len) if is_dense_with_hub else text_len total_k_len = (hub_size + text_len) if is_dense_with_hub else (past_k_len + text_len) device = x_tokens.device mask = torch.zeros((batch_size, 1, cur_seq_len, total_k_len), device=device, dtype=torch.bool) is_eod = (x_tokens == eod_token_id).long() doc_ids = torch.cumsum(is_eod, dim=-1) doc_ids_shifted = torch.cat( [torch.zeros((batch_size, 1), device=device, dtype=torch.long), doc_ids[:, :-1]], dim=-1 ) doc_mismatch = (doc_ids_shifted.unsqueeze(-1) != doc_ids_shifted.unsqueeze(-2)) rows = torch.arange(text_len, device=device).unsqueeze(1) cols = torch.arange(text_len, device=device).unsqueeze(0) future_mask = (cols > rows).unsqueeze(0).unsqueeze(0) if is_dense_with_hub: mask[:, :, :hub_size, hub_size:] = True combined_text_mask = future_mask | doc_mismatch.unsqueeze(1) mask[:, :, hub_size:, hub_size:] = combined_text_mask else: past_text_len = past_k_len - hub_size combined_text_mask = future_mask | doc_mismatch.unsqueeze(1) mask[:, :, :, past_k_len:] = combined_text_mask if past_text_len > 0: past_doc_mask = (doc_ids_shifted > 0).unsqueeze(-1).expand(-1, -1, past_text_len).unsqueeze(1) mask[:, :, :, hub_size:past_k_len] = past_doc_mask return mask def resolve_head_architecture( dim: int, num_heads: Optional[Union[int, str]] = "auto", d_head: Optional[Union[int, str]] = "auto", ) -> Tuple[int, int]: if ( isinstance(num_heads, int) and num_heads > 0 and isinstance(d_head, int) and d_head > 0 ): return num_heads, d_head if isinstance(num_heads, int) and num_heads > 0: resolved_d_head = max(16, dim // num_heads) return num_heads, resolved_d_head if isinstance(d_head, int) and d_head > 0: resolved_heads = max(1, dim // d_head) return resolved_heads, d_head target_d_head = 2 ** round(math.log2(max(32.0, math.sqrt(2.0 * dim)))) candidate_divisors = [d for d in range(16, dim + 1, 8) if dim % d == 0] if candidate_divisors: resolved_d_head = min( candidate_divisors, key=lambda x: abs(x - target_d_head) ) else: resolved_d_head = 64 if dim % 64 == 0 else (32 if dim % 32 == 0 else 16) resolved_heads = max(1, dim // resolved_d_head) return resolved_heads, resolved_d_head TIER_CONFIGS = { "65M": dict( dim=512, num_layers=12, num_heads=8, d_head=64, hub_size=256, num_specialized_hubs=8, num_terminals=16, slots_per_terminal=4, max_episodic=32, kv_latent_dim=64, lora_rank=32, ), "100M": dict( dim=640, num_layers=14, num_heads=10, d_head=64, hub_size=288, num_specialized_hubs=10, num_terminals=18, slots_per_terminal=4, max_episodic=40, kv_latent_dim=80, lora_rank=40, ), "200M": dict( dim=896, num_layers=18, num_heads=14, d_head=64, hub_size=384, num_specialized_hubs=12, num_terminals=21, slots_per_terminal=4, max_episodic=48, kv_latent_dim=112, lora_rank=56, ), "300M": dict( dim=1280, num_layers=20, num_heads=16, d_head=80, hub_size=448, num_specialized_hubs=14, num_terminals=24, slots_per_terminal=6, max_episodic=56, kv_latent_dim=160, lora_rank=80, ), "500M": dict( dim=1280, num_layers=20, num_heads=16, d_head=80, hub_size=512, num_specialized_hubs=18, num_terminals=28, slots_per_terminal=6, max_episodic=72, kv_latent_dim=160, lora_rank=80, ), "750M": dict( dim=1536, num_layers=22, num_heads=16, d_head=96, hub_size=608, num_specialized_hubs=22, num_terminals=32, slots_per_terminal=6, max_episodic=80, kv_latent_dim=192, lora_rank=96, ), "1.0B": dict( dim=1792, num_layers=24, num_heads=16, d_head=112, hub_size=768, num_specialized_hubs=24, num_terminals=32, slots_per_terminal=8, max_episodic=88, kv_latent_dim=224, lora_rank=112, ), "3.0B": dict( dim=2560, num_layers=32, num_heads=20, d_head=128, hub_size=992, num_specialized_hubs=36, num_terminals=42, slots_per_terminal=8, max_episodic=136, kv_latent_dim=320, lora_rank=160, ), "7.0B": dict( dim=4096, num_layers=36, num_heads=32, d_head=128, hub_size=1312, num_specialized_hubs=52, num_terminals=51, slots_per_terminal=8, max_episodic=192, kv_latent_dim=512, lora_rank=256, ), } @dataclass class XoneLMOutput: loss: Optional[torch.Tensor] = None logits: Optional[torch.Tensor] = None aux_loss: Optional[torch.Tensor] = None z_loss: Optional[torch.Tensor] = None past_key_values: Optional[List[torch.Tensor]] = None soliton_state: Optional[List[torch.Tensor]] = None class RMSNorm(nn.Module): def __init__(self, dim: int, eps: float = 1e-6): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim)) def forward(self, x: torch.Tensor) -> torch.Tensor: input_dtype = x.dtype x_f32 = x.to(torch.float32) variance = x_f32.pow(2).mean(dim=-1, keepdim=True) normed = x_f32 * torch.rsqrt(variance + self.eps) return (normed * self.weight.to(torch.float32)).to(dtype=input_dtype) class SwiGLU(nn.Module): def __init__(self, dim: int, multiple_of: int = 32): super().__init__() hidden_dim = multiple_of * ( (int(2 * (dim * 4) / 3) + multiple_of - 1) // multiple_of ) self.w12 = nn.Linear(dim, 2 * hidden_dim, bias=False) self.w3 = nn.Linear(hidden_dim, dim, bias=False) def forward(self, x: torch.Tensor) -> torch.Tensor: gate, value = self.w12(x).chunk(2, dim=-1) return self.w3(F.silu(gate) * value) class PolyHoPE(nn.Module): def __init__(self, dim: int, degree: int = 180, max_seq_len: int = 8192): super().__init__() self.dim = dim self.degree = degree self.max_seq_len = max_seq_len self.w_poly = nn.Parameter(torch.randn(degree + 1, dim) * 0.02) idx = torch.arange(max_seq_len, dtype=torch.float32) t = (2.0 * idx / (max_seq_len - 1.0)) - 1.0 t = t.clamp(-1.0, 1.0) t_poly = [torch.ones_like(t), t] for n in range(2, degree + 1): t_poly.append(2.0 * t * t_poly[n - 1] - t_poly[n - 2]) t_stack = torch.stack(t_poly, dim=1) self.register_buffer("t_stack", t_stack, persistent=False) def forward( self, seq_len: int, device: torch.device, dtype: torch.dtype, offset: int = 0, ) -> torch.Tensor: grid = self.t_stack[offset : offset + seq_len].to( device=device, dtype=torch.float32 ) pe = torch.matmul(grid, self.w_poly.to(dtype=torch.float32)).to(dtype=dtype) return pe.unsqueeze(0) class LinHoPE(nn.Module): def __init__( self, num_heads: int = 8, hub_size: int = 256, num_hub_heads: Optional[int] = None, min_slope: float = 0.01, max_slope: float = 0.45, mode: str = "geometric", learnable: bool = False, ): super().__init__() self.num_heads = num_heads self.hub_size = hub_size self.mode = mode.lower() self.learnable = learnable if isinstance(num_hub_heads, int) and num_hub_heads > 0: self.num_hub_heads = max(1, min(num_hub_heads, num_heads - 1)) else: self.num_hub_heads = max(1, num_heads // 8) self.num_text_heads = self.num_heads - self.num_hub_heads hub_slopes = torch.full( (1, self.num_hub_heads, 1, 1), 0.015, dtype=torch.float32 ) if self.num_text_heads > 1: text_slopes = torch.linspace( min_slope, max_slope, self.num_text_heads ).view(1, self.num_text_heads, 1, 1) else: text_slopes = torch.full((1, 1, 1, 1), 0.20, dtype=torch.float32) init_slopes = torch.cat([hub_slopes, text_slopes], dim=1) if learnable: self.raw_slopes = nn.Parameter( torch.log(torch.exp(init_slopes) - 1.0 + 1e-6) ) else: self.register_buffer("raw_slopes", init_slopes, persistent=False) log_hub = math.log2(max(2.0, float(hub_size))) hub_m = torch.zeros(1, self.num_hub_heads, 1, 1) if self.num_text_heads > 1: text_m = torch.linspace( 2.0 * log_hub, 8.0 * log_hub, self.num_text_heads ).view(1, self.num_text_heads, 1, 1) else: text_m = torch.full((1, 1, 1, 1), 4.5 * log_hub) hub_damping = torch.cat([hub_m, text_m], dim=1) self.register_buffer("hub_damping", hub_damping, persistent=False) @property def slopes(self) -> torch.Tensor: if self.learnable: return F.softplus(self.raw_slopes) + 1e-4 return self.raw_slopes def forward( self, seq_len: int, num_all: int, device: torch.device, dtype: torch.dtype, is_dense_with_hub: bool = True, past_c_kv: Optional[torch.Tensor] = None, ) -> torch.Tensor: text_k_len = num_all - self.hub_size hub_pos = torch.arange( -self.hub_size, 0, device=device, dtype=torch.float32 ) text_pos_k = torch.arange(0, text_k_len, device=device, dtype=torch.float32) pos_k = torch.cat([hub_pos, text_pos_k], dim=0).unsqueeze(0) if is_dense_with_hub and (past_c_kv is None): text_q_len = seq_len - self.hub_size text_pos_q = torch.arange( 0, text_q_len, device=device, dtype=torch.float32 ) pos_q = torch.cat([hub_pos, text_pos_q], dim=0).unsqueeze(1) else: pos_q = torch.arange( text_k_len - seq_len, text_k_len, device=device, dtype=torch.float32 ).unsqueeze(1) raw_dist = (pos_q - pos_k).clamp(min=0.0) dist_matrix = ( raw_dist.unsqueeze(0) .unsqueeze(0) .repeat(1, self.num_heads, 1, 1) .to(device=device, dtype=torch.float32) ) dist_matrix[:, :, :, : self.hub_size] = self.hub_damping.to( device=device, dtype=torch.float32 ) active_slopes = self.slopes.to(device=device, dtype=torch.float32) if self.mode == "rational": denom = 1.0 + active_slopes * dist_matrix log_bias = -torch.log(denom.clamp(min=1e-8)) else: log_bias = -active_slopes * dist_matrix return log_bias.clamp(min=-10.0, max=0.0).to(dtype=dtype) class Isomorphic3DHoPE(nn.Module): def __init__(self, dim: int, theta: float = 10000.0): super().__init__() self.dim = dim self.dim_z = 2 * (dim // 8) rem = dim - self.dim_z self.dim_y = 2 * (rem // 4) self.dim_x = dim - self.dim_z - self.dim_y self.register_buffer( "inv_freq_z", 1.0 / ( theta ** ( torch.arange(0, self.dim_z, 2, dtype=torch.float32) / self.dim_z ) ), persistent=False, ) self.register_buffer( "inv_freq_y", 1.0 / ( theta ** ( torch.arange(0, self.dim_y, 2, dtype=torch.float32) / self.dim_y ) ), persistent=False, ) self.register_buffer( "inv_freq_x", 1.0 / ( theta ** ( torch.arange(0, self.dim_x, 2, dtype=torch.float32) / self.dim_x ) ), persistent=False, ) def forward( self, p_z: torch.Tensor, p_y: torch.Tensor, p_x: torch.Tensor ) -> torch.Tensor: target_device = self.inv_freq_z.device p_z = p_z.to(device=target_device).float() p_y = p_y.to(device=target_device).float() p_x = p_x.to(device=target_device).float() omega_z = p_z.unsqueeze(-1) * self.inv_freq_z omega_y = p_y.unsqueeze(-1) * self.inv_freq_y omega_x = p_x.unsqueeze(-1) * self.inv_freq_x omega_p = torch.cat([omega_z, omega_y, omega_x], dim=-1) return torch.cat([torch.cos(omega_p), torch.sin(omega_p)], dim=-1) class TopologicalSolitonWaveletState(nn.Module): def __init__(self, kv_latent_dim: int, kappa: float = 0.1): super().__init__() self.kv_latent_dim = kv_latent_dim self.kappa = kappa self.ws = nn.Linear(kv_latent_dim, kv_latent_dim, bias=False) def forward( self, c_kv: torch.Tensor, s_prev: Optional[torch.Tensor] = None ) -> torch.Tensor: batch_size = c_kv.shape[0] if s_prev is None: s_prev = torch.zeros( batch_size, 1, self.kv_latent_dim, device=c_kv.device, dtype=c_kv.dtype, ) c_kv_summary = ( c_kv.mean(dim=1, keepdim=True) if c_kv.ndim == 3 else c_kv.unsqueeze(1) ) tanh_s = torch.tanh(s_prev) sech_sq = (1.0 - tanh_s.pow(2)).clamp(min=1e-6) delta_s = self.kappa * sech_sq * torch.tanh(self.ws(c_kv_summary)) return s_prev + delta_s def compute_fisher_spectral_anisotropy( tensor: torch.Tensor, eps: float = 1e-8 ) -> torch.Tensor: power = tensor.pow(2) total_power = power.sum(dim=-1, keepdim=True) + eps p_c = power / total_power shannon_entropy = -torch.sum(p_c * torch.log(p_c + eps), dim=-1, keepdim=True) max_entropy = math.log(max(tensor.shape[-1], 2)) anisotropy = 1.0 - (shannon_entropy / max_entropy) return anisotropy.clamp(0.0, 1.0) class PoincareHyperbolicTerminalRouter(nn.Module): def __init__( self, dim: int, max_terminals: int = 16, c: float = 1.0, eps: float = 1e-5, ): super().__init__() self.dim = dim self.max_terminals = max_terminals self.c = c self.eps = eps self.query_proj = nn.Linear(dim, dim, bias=False) self.norm = RMSNorm(dim) def forward( self, chunk_summaries: torch.Tensor, hub_query: torch.Tensor ) -> torch.Tensor: batch_size, num_chunks, hidden_dim = chunk_summaries.shape if num_chunks <= self.max_terminals: return self.norm(chunk_summaries) q_mean = hub_query.mean(dim=1, keepdim=True) q_proj = self.query_proj(q_mean) q_norm = q_proj.norm(p=2, dim=-1, keepdim=True) + 1e-8 q_p = q_proj * ( torch.tanh(math.sqrt(self.c) * q_norm) / (math.sqrt(self.c) * q_norm) ) c_norm = chunk_summaries.norm(p=2, dim=-1, keepdim=True) + 1e-8 c_p = chunk_summaries * ( torch.tanh(math.sqrt(self.c) * c_norm) / (math.sqrt(self.c) * c_norm) ) u_sq = (q_p**2).sum(dim=-1, keepdim=True) v_sq = (c_p**2).sum(dim=-1, keepdim=True) diff_sq = ((q_p - c_p) ** 2).sum(dim=-1) denom = torch.clamp((1.0 - u_sq) * (1.0 - v_sq), min=self.eps).squeeze(-1) delta = 2.0 * diff_sq / denom dist_poincare = torch.acosh(1.0 + delta) probs = F.softmax(-dist_poincare, dim=-1) _, top_k_indices = torch.topk(probs, self.max_terminals, dim=-1) anchors = torch.gather( chunk_summaries, 1, top_k_indices.unsqueeze(-1).expand(-1, -1, hidden_dim), ) anchors_norm = F.normalize(anchors, p=2, dim=-1) c_n = F.normalize(chunk_summaries, p=2, dim=-1) assign_sim = torch.matmul(anchors_norm, c_n.transpose(-1, -2)) soft_assignment = F.softmax(assign_sim * 10.0, dim=-1) compacted = torch.matmul(soft_assignment, chunk_summaries) return self.norm(compacted) class HierarchicalEpisodicMemoryBank(nn.Module): def __init__( self, dim: int, max_l1_terminals: int = 16, max_l2_episodic: int = 32, tau_lock: float = 0.65, ): super().__init__() self.dim = dim self.max_l1_terminals = max_l1_terminals self.max_l2_episodic = max_l2_episodic self.tau_lock = tau_lock self.norm = RMSNorm(dim) def forward( self, chunk_summaries: torch.Tensor, hub_query: torch.Tensor ) -> Tuple[torch.Tensor, torch.Tensor]: batch_size, num_chunks, hidden_dim = chunk_summaries.shape q_norm = F.normalize(hub_query.mean(dim=1, keepdim=True), p=2, dim=-1) c_norm = F.normalize(chunk_summaries, p=2, dim=-1) anisotropy = compute_fisher_spectral_anisotropy(chunk_summaries).squeeze(-1) c_l2 = chunk_summaries.norm(p=2, dim=-1) q_align = torch.abs( torch.matmul(c_norm, q_norm.transpose(-1, -2)).squeeze(-1) ) s_fisher = anisotropy * c_l2 * q_align l1_terminals = chunk_summaries[:, : min(num_chunks, self.max_l1_terminals), :] if num_chunks > self.max_l2_episodic: _, l2_top_idx = torch.topk(s_fisher, self.max_l2_episodic, dim=-1) l2_episodic = torch.gather( chunk_summaries, 1, l2_top_idx.unsqueeze(-1).expand(-1, -1, hidden_dim) ) else: l2_mask = (s_fisher > self.tau_lock).unsqueeze(-1) l2_episodic = chunk_summaries * l2_mask return self.norm(l1_terminals), self.norm(l2_episodic) class EpistemicTruthVerifierGate(nn.Module): def __init__(self, dim: int, tau_contra: float = 0.30, beta: float = 1.0): super().__init__() self.dim = dim self.tau_contra = tau_contra self.beta = beta self.wk = nn.Linear(dim, dim, bias=False) self.wv = nn.Linear(dim, dim, bias=False) self.norm = RMSNorm(dim) def forward( self, claim_states: torch.Tensor, parametric_facts: torch.Tensor ) -> Tuple[torch.Tensor, torch.Tensor]: k_facts = self.wk(parametric_facts) v_facts = self.wv(parametric_facts) w_truth = F.softmax( self.beta * torch.matmul(claim_states, k_facts.transpose(-1, -2)) / math.sqrt(self.dim), dim=-1, ) v_expected = torch.matmul(w_truth, v_facts) c_norm = F.normalize(claim_states, p=2, dim=-1) v_norm = F.normalize(v_expected, p=2, dim=-1) s_epistemic = torch.sum(c_norm * v_norm, dim=-1, keepdim=True) gate = torch.sigmoid((s_epistemic - self.tau_contra) * 5.0) verified_states = self.norm(claim_states * gate) fallacy_quarantine = self.norm(claim_states * (1.0 - gate)) return verified_states, fallacy_quarantine class DirectiveCognitiveCompass(nn.Module): def __init__(self, dim: int): super().__init__() self.dim = dim self.wdir = nn.Linear(dim, dim, bias=False) self.norm = RMSNorm(dim) def forward( self, chunk_summaries: torch.Tensor, directive_tokens: torch.Tensor ) -> torch.Tensor: v_compass = self.norm(self.wdir(directive_tokens.mean(dim=1, keepdim=True))) v_comp_norm = F.normalize(v_compass, p=2, dim=-1) c_norm = F.normalize(chunk_summaries, p=2, dim=-1) align = torch.abs(torch.matmul(c_norm, v_comp_norm.transpose(-1, -2))) anisotropy = compute_fisher_spectral_anisotropy(chunk_summaries) c_l2 = chunk_summaries.norm(p=2, dim=-1, keepdim=True) s_dcc = align * anisotropy * c_l2 weighted_chunks = chunk_summaries * (1.0 + torch.sigmoid(s_dcc)) return self.norm(weighted_chunks) class FactPullingAttention(nn.Module): def __init__( self, dim: int, num_heads: int = 8, num_terminals: int = 16, slots_per_terminal: int = 4, ): super().__init__() self.dim = dim self.num_heads = num_heads self.num_terminals = num_terminals self.slots_per_terminal = slots_per_terminal self.total_slots = num_terminals * slots_per_terminal self.aspect_drawers = nn.Parameter( torch.randn(1, num_terminals, slots_per_terminal, dim) * (1.0 / math.sqrt(dim)) ) self.q_proj = nn.Linear(dim, dim, bias=False) self.k_ctx_proj = nn.Linear(dim, dim, bias=False) self.v_ctx_proj = nn.Linear(dim, dim, bias=False) self.k_param_proj = nn.Linear(dim, dim, bias=False) self.v_param_proj = nn.Linear(dim, dim, bias=False) self.out_proj = nn.Linear(dim, dim, bias=False) self.norm = RMSNorm(dim) self.gamma_p = nn.Parameter(torch.ones(1)) self.beta_p = nn.Parameter(torch.zeros(1)) def forward( self, terminal_base: torch.Tensor, context_states: torch.Tensor, parametric_facts: torch.Tensor, ) -> torch.Tensor: batch_size = context_states.shape[0] q_drawers = (terminal_base.unsqueeze(2) + self.aspect_drawers).view( batch_size, self.total_slots, self.dim ) q = self.q_proj(q_drawers) k_ctx = self.k_ctx_proj(context_states) v_ctx = self.v_ctx_proj(context_states) k_param = self.k_param_proj(parametric_facts) v_param = self.v_param_proj(parametric_facts) scores_ctx = torch.matmul(q, k_ctx.transpose(-1, -2)) / math.sqrt(self.dim) scores_param = ( torch.matmul(q, k_param.transpose(-1, -2)) / math.sqrt(self.dim) ) * self.gamma_p + self.beta_p probs = F.softmax( torch.cat([scores_ctx, scores_param], dim=-1), dim=-1 ).to(dtype=q.dtype) v_comb = torch.cat([v_ctx, v_param], dim=-2) attn_out = torch.matmul(probs, v_comb) return self.norm(self.out_proj(attn_out)) class LatentMentalRollout(nn.Module): def __init__(self, dim: int, num_rollout_steps: int = 2): super().__init__() self.dim = dim self.num_rollout_steps = num_rollout_steps self.w_transition = nn.Parameter( torch.randn(dim, dim) * (0.02 / math.sqrt(dim)) ) self.norm = RMSNorm(dim) def forward(self, h_dyn: torch.Tensor) -> torch.Tensor: for _ in range(self.num_rollout_steps): delta = torch.tanh(torch.matmul(h_dyn, self.w_transition)) h_dyn = self.norm(h_dyn + delta) return h_dyn class PhysicsCausalDiffusionAttention(nn.Module): def __init__( self, dim: int, num_heads: int = 8, d_head: int = 64, kv_latent_dim: int = 64, hub_size: int = 256, num_hub_heads: Optional[int] = None, use_sdpa: bool = True, ): super().__init__() self.dim = dim self.num_heads = num_heads self.d_head = d_head self.inner_attn_dim = num_heads * d_head self.kv_latent_dim = kv_latent_dim self.hub_size = hub_size self.use_sdpa = use_sdpa and HAS_SDPA self.kv_down_proj = nn.Linear(dim, kv_latent_dim, bias=False) self.kv_ln = RMSNorm(kv_latent_dim) self.kv_up_proj = nn.Linear( kv_latent_dim, 2 * self.inner_attn_dim, bias=False ) self.q_proj = nn.Linear(dim, self.inner_attn_dim, bias=False) self.out_proj = nn.Linear(self.inner_attn_dim, dim, bias=False) self.q_norm = RMSNorm(self.d_head) self.k_norm = RMSNorm(self.d_head) self.lin_hope = LinHoPE( num_heads=num_heads, hub_size=hub_size, num_hub_heads=num_hub_heads, min_slope=0.01, max_slope=0.45, mode="geometric", ) self.scale_factor = 1.0 / math.sqrt(self.d_head) def forward( self, x: torch.Tensor, attn_mask: Optional[torch.Tensor] = None, is_causal: bool = True, past_c_kv: Optional[torch.Tensor] = None, soliton_state: Optional[torch.Tensor] = None, flex_block_mask=None, is_dense_with_hub: bool = True, ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: batch_size, seq_len, _ = x.shape c_kv_current = self.kv_ln(self.kv_down_proj(x)) if past_c_kv is not None: c_kv_all = torch.cat([past_c_kv, c_kv_current], dim=1) else: c_kv_all = c_kv_current num_all = c_kv_all.shape[1] k_up, v_up = torch.split( self.kv_up_proj(c_kv_all), [self.inner_attn_dim, self.inner_attn_dim], dim=-1, ) k = k_up.reshape( batch_size, num_all, self.num_heads, self.d_head ).permute(0, 2, 1, 3) v = v_up.reshape( batch_size, num_all, self.num_heads, self.d_head ).permute(0, 2, 1, 3) q = self.q_proj(x).reshape( batch_size, seq_len, self.num_heads, self.d_head ).permute(0, 2, 1, 3) q = self.q_norm(q) k = self.k_norm(k) log_decay_bias = self.lin_hope( seq_len=seq_len, num_all=num_all, device=x.device, dtype=q.dtype, is_dense_with_hub=is_dense_with_hub, past_c_kv=past_c_kv, ) full_mask = log_decay_bias.clone() if attn_mask is not None: if attn_mask.dtype == torch.bool: full_mask = full_mask.masked_fill(attn_mask, -10000.0) else: full_mask = full_mask + attn_mask elif is_causal: text_k_len = num_all - self.hub_size causal_m = torch.zeros( (1, 1, seq_len, num_all), device=x.device, dtype=q.dtype ) if is_dense_with_hub and (past_c_kv is None): text_q_len = seq_len - self.hub_size causal_m[:, :, : self.hub_size, self.hub_size :] = -10000.0 rows = torch.arange(text_q_len, device=x.device).unsqueeze(1) cols = torch.arange(text_k_len, device=x.device).unsqueeze(0) causal_m[:, :, self.hub_size :, self.hub_size :].masked_fill_( cols > rows, -10000.0 ) else: pos_q = torch.arange( text_k_len - seq_len, text_k_len, device=x.device, dtype=torch.float32, ).unsqueeze(1) text_pos_k = torch.arange( 0, text_k_len, device=x.device, dtype=torch.float32 ).unsqueeze(0) future_text = (pos_q - text_pos_k) < 0 causal_m[:, :, :, self.hub_size :].masked_fill_( future_text.unsqueeze(0).unsqueeze(0), -10000.0 ) full_mask = full_mask + causal_m if self.use_sdpa: attn_out = F.scaled_dot_product_attention( q, k, v, attn_mask=full_mask, scale=self.scale_factor ) else: scores = ( torch.matmul(q, k.transpose(-1, -2)) * self.scale_factor + full_mask ) probs = F.softmax(scores, dim=-1, dtype=torch.float32).to(dtype=q.dtype) attn_out = torch.matmul(probs, v) attn_out = attn_out.permute(0, 2, 1, 3).reshape( batch_size, seq_len, self.inner_attn_dim ) final_out = self.out_proj(attn_out) return final_out, c_kv_all, None class XoneLMBlock(nn.Module): def __init__( self, dim: int, num_heads: int = 8, d_head: int = 64, kv_latent_dim: int = 64, hub_size: int = 256, num_hub_heads: Optional[int] = None, use_sdpa: bool = True, ): super().__init__() self.ln1 = RMSNorm(dim) self.attn = PhysicsCausalDiffusionAttention( dim=dim, num_heads=num_heads, d_head=d_head, kv_latent_dim=kv_latent_dim, hub_size=hub_size, num_hub_heads=num_hub_heads, use_sdpa=use_sdpa, ) self.ln2 = RMSNorm(dim) self.hub_size = hub_size self.ffn = SwiGLU(dim, multiple_of=32) self.poly_alpha = nn.Parameter(torch.tensor(0.05)) def forward( self, x: torch.Tensor, poly_pe: Optional[torch.Tensor] = None, attn_mask: Optional[torch.Tensor] = None, is_causal: bool = True, past_c_kv: Optional[torch.Tensor] = None, soliton_state: Optional[torch.Tensor] = None, flex_block_mask=None, is_dense_with_hub: bool = True, ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: if poly_pe is not None: x = x + self.poly_alpha * poly_pe attn_out, new_past_c_kv, new_soliton = self.attn( self.ln1(x), attn_mask=attn_mask, is_causal=is_causal, past_c_kv=past_c_kv, soliton_state=soliton_state, is_dense_with_hub=is_dense_with_hub, ) x = x + attn_out x = x + self.ffn(self.ln2(x)) return x, new_past_c_kv, new_soliton class XoneLM(nn.Module): def __init__( self, tier: Optional[str] = None, vocab_size: int = 32000, dim: Optional[int] = None, num_layers: Optional[int] = None, num_heads: Optional[Union[int, str]] = "auto", d_head: Optional[Union[int, str]] = "auto", num_hub_heads: Optional[int] = None, hub_size: Optional[int] = None, num_specialized_hubs: Optional[int] = None, num_terminals: Optional[int] = None, slots_per_terminal: int = 4, max_terminals: Optional[int] = None, tokens_per_terminal: Optional[int] = None, max_episodic: Optional[int] = None, kv_latent_dim: Optional[int] = None, kv_lora_dim: Optional[int] = None, lora_rank: Optional[int] = None, chunk_size: int = 1024, alpha_anchor: float = 0.1, checkpoint_every_n: int = 0, separator_token_id: Optional[int] = None, config: Optional[Any] = None, hub_mode: str = "moh", use_sdpa: bool = True, **kwargs, ): super().__init__() base_params = {} if tier is not None and tier in TIER_CONFIGS: base_params = TIER_CONFIGS[tier].copy() if kv_latent_dim is None and kv_lora_dim is not None: kv_latent_dim = kv_lora_dim self.vocab_size = ( vocab_size if config is None else getattr(config, "vocab_size", vocab_size) ) self.dim = dim if dim is not None else base_params.get("dim", 512) self.num_layers = ( num_layers if num_layers is not None else base_params.get("num_layers", 12) ) cfg_heads = base_params.get("num_heads", "auto") cfg_d_head = base_params.get("d_head", "auto") req_heads = num_heads if num_heads != "auto" else cfg_heads req_d_head = d_head if d_head != "auto" else cfg_d_head self.num_heads, self.d_head = resolve_head_architecture( dim=self.dim, num_heads=req_heads, d_head=req_d_head ) self.hub_size = ( hub_size if hub_size is not None else base_params.get("hub_size", 256) ) self.num_specialized_hubs = ( num_specialized_hubs if num_specialized_hubs is not None else base_params.get("num_specialized_hubs", 8) ) self.num_terminals = ( num_terminals or max_terminals or base_params.get("num_terminals", 16) ) self.slots_per_terminal = ( slots_per_terminal or base_params.get("slots_per_terminal", 4) ) self.total_terminal_slots = ( self.num_terminals * self.slots_per_terminal ) self.max_episodic = ( max_episodic if max_episodic is not None else base_params.get("max_episodic", 32) ) self.kv_latent_dim = ( kv_latent_dim if kv_latent_dim is not None else base_params.get("kv_latent_dim", max(64, self.dim // 8)) ) self.lora_rank = ( lora_rank if lora_rank is not None else base_params.get("lora_rank", max(32, self.kv_latent_dim // 2)) ) self.chunk_size = chunk_size self.alpha_anchor = alpha_anchor self.checkpoint_every_n = max(0, checkpoint_every_n) self.separator_token_id = separator_token_id self.hub_mode = hub_mode self.num_hub_heads = num_hub_heads self.static_hub_slots = self.hub_size // 2 self.dynamic_hub_slots = self.hub_size - self.static_hub_slots mask_static = torch.zeros(1, self.static_hub_slots, 1, dtype=torch.float32) mask_dynamic = torch.ones(1, self.dynamic_hub_slots, 1, dtype=torch.float32) self.register_buffer( "dynamic_slot_mask", torch.cat([mask_static, mask_dynamic], dim=1), persistent=False, ) raw_basis = torch.randn(self.dim, self.hub_size) q_basis, _ = torch.linalg.qr(raw_basis) self.shared_hub_base = nn.Parameter(q_basis.T.unsqueeze(0).contiguous()) self.hub_lora_a = nn.Parameter( torch.randn(self.num_specialized_hubs, self.hub_size, self.lora_rank) * 0.02 ) self.hub_lora_b = nn.Parameter( torch.randn(self.num_specialized_hubs, self.lora_rank, self.dim) * 0.02 ) self.hub_router_gate = nn.Linear( self.dim, self.num_specialized_hubs, bias=False ) self.w_anchor = nn.Linear(self.dim, self.dim, bias=False) self.token_embeddings = nn.Embedding(self.vocab_size, self.dim) nn.init.normal_(self.token_embeddings.weight, mean=0.0, std=0.02) self.poly_hope = PolyHoPE(dim=self.dim, degree=180, max_seq_len=8192) self.hope_3d = Isomorphic3DHoPE(dim=self.dim) self.terminal_router = PoincareHyperbolicTerminalRouter( dim=self.dim, max_terminals=self.num_terminals ) self.episodic_memory = HierarchicalEpisodicMemoryBank( dim=self.dim, max_l1_terminals=self.num_terminals, max_l2_episodic=self.max_episodic, ) self.etvg = EpistemicTruthVerifierGate(dim=self.dim) self.dcc = DirectiveCognitiveCompass(dim=self.dim) self.latent_rollout = LatentMentalRollout(dim=self.dim) self.slot_expander = nn.Linear(self.dim, self.dim, bias=False) self.layers = nn.ModuleList([ XoneLMBlock( dim=self.dim, num_heads=self.num_heads, d_head=self.d_head, kv_latent_dim=self.kv_latent_dim, hub_size=self.hub_size, num_hub_heads=num_hub_heads, use_sdpa=use_sdpa, ) for _ in range(self.num_layers) ]) self.norm = RMSNorm(self.dim) self.head = nn.Linear(self.dim, self.vocab_size, bias=False) self.head.weight = self.token_embeddings.weight self.fused_ce_loss = None if hasattr(nn, "LinearCrossEntropyLoss"): try: self.fused_ce_loss = nn.LinearCrossEntropyLoss(ignore_index=-100) except Exception: self.fused_ce_loss = None self.parametric_facts = nn.Parameter(torch.randn(1, 32, self.dim) * 0.02) self.fact_puller = FactPullingAttention( dim=self.dim, num_heads=self.num_heads, num_terminals=self.num_terminals, slots_per_terminal=self.slots_per_terminal, ) self._block_mask_cache: Dict[Tuple[int, int, str], Any] = {} @property def total_active_hub_slots(self) -> int: return self.hub_size def _init_hub_base( self, batch_size: int, query_rep: Optional[torch.Tensor] = None ) -> torch.Tensor: base = self.shared_hub_base.expand(batch_size, -1, -1) if query_rep is None: return base q_vec = query_rep.mean(dim=1) routing_weights = F.softmax( self.hub_router_gate(q_vec) / math.sqrt(self.dim), dim=-1 ) expert_deltas = torch.bmm(self.hub_lora_a, self.hub_lora_b) combined_delta = torch.einsum("be,esd->bsd", routing_weights, expert_deltas) return base + combined_delta def extract_hub( self, x: torch.Tensor, page_idx: Optional[torch.Tensor] = None, para_idx: Optional[torch.Tensor] = None, sent_idx: Optional[torch.Tensor] = None, directive_tokens: Optional[torch.Tensor] = None, is_sft: bool = False, ) -> torch.Tensor: batch_size, total_len = x.shape chunk_size = min(total_len, self.chunk_size) x_prompt_only = x[:, :chunk_size] prompt_len = x_prompt_only.shape[1] if page_idx is None: page_idx = torch.zeros(batch_size, 1, device=x.device, dtype=torch.long) if para_idx is None: para_idx = torch.zeros(batch_size, 1, device=x.device, dtype=torch.long) if sent_idx is None: sent_idx = torch.zeros(batch_size, 1, device=x.device, dtype=torch.long) hope_anchor = self.hope_3d(page_idx, para_idx, sent_idx) hope_anchor_scaled = F.normalize(hope_anchor, p=2, dim=-1) * 0.05 pe_text = self.poly_hope( prompt_len, x.device, self.token_embeddings.weight.dtype, offset=0 ) h_text_all = self.token_embeddings(x_prompt_only) + pe_text hub_init = self._init_hub_base(batch_size, query_rep=h_text_all) c_term_base = hub_init[:, : self.num_terminals, :] expanded_pf = self.parametric_facts.expand(batch_size, -1, -1) c_hub_out = self.fact_puller(c_term_base, h_text_all, expanded_pf) if directive_tokens is not None: c_hub_out = self.dcc(c_hub_out, directive_tokens) verified_terminals, _ = self.etvg(c_hub_out, expanded_pf) q_norm = F.normalize(hub_init, p=2, dim=-1) k_norm = F.normalize(verified_terminals, p=2, dim=-1) attn_weights = F.softmax( torch.matmul(q_norm, k_norm.transpose(-1, -2)) * 8.0, dim=-1 ) # [B, hub_size, total_terminal_slots] context_features = torch.matmul( attn_weights, verified_terminals ) # [B, hub_size, dim] normed_base = F.normalize(hub_init, p=2, dim=-1) normed_context = F.normalize(context_features, p=2, dim=-1) hub = ( 0.85 * normed_base + 0.15 * normed_context + hope_anchor_scaled ) * math.sqrt(self.dim) if self.dynamic_hub_slots > 0: h_evolved = self.latent_rollout(hub) delta_evolved = h_evolved - hub hub = hub + delta_evolved * self.dynamic_slot_mask.to(dtype=hub.dtype) return hub.to(dtype=self.token_embeddings.weight.dtype) def compute_hub_diversity_loss(self, hub: torch.Tensor) -> torch.Tensor: active_mask = (hub.norm(dim=-1) > 1e-4).float() pair_mask = torch.matmul(active_mask.unsqueeze(-1), active_mask.unsqueeze(-2)) h_norm = F.normalize(hub, p=2, dim=-1, eps=1e-8) sim_matrix = torch.matmul(h_norm, h_norm.transpose(-1, -2)) identity = torch.eye( self.hub_size, device=hub.device, dtype=hub.dtype ).unsqueeze(0) diff_sq = (sim_matrix - identity).pow(2) * pair_mask return diff_sq.sum() / pair_mask.sum().clamp(min=1.0) def _compute_loss_efficient( self, hidden_states: torch.Tensor, labels: torch.Tensor ) -> torch.Tensor: valid_mask = labels != -100 if not valid_mask.any(): return (hidden_states * 0.0).sum() if hasattr(F, "linear_cross_entropy"): try: return F.linear_cross_entropy( hidden_states, self.head.weight, labels, ignore_index=-100, reduction="mean", ) except Exception: pass if self.fused_ce_loss is not None: try: return self.fused_ce_loss(hidden_states, self.head.weight, labels) except Exception: pass logits = self.head(hidden_states) return F.cross_entropy( logits.reshape(-1, self.vocab_size).float(), labels.reshape(-1), ignore_index=-100, ) def _forward_dense( self, x: torch.Tensor, past_c_kv_list: Optional[List[torch.Tensor]] = None, soliton_state_list: Optional[List[torch.Tensor]] = None, override_hub: Optional[torch.Tensor] = None, return_logits: bool = True, past_key_values: Optional[Any] = None, attn_mask: Optional[torch.Tensor] = None, **kwargs, ) -> Tuple: if past_c_kv_list is None and past_key_values is not None: past_c_kv_list = past_key_values batch_size, seq_len = x.shape target_dtype = self.token_embeddings.weight.dtype if past_c_kv_list is None: pe_text = self.poly_hope(seq_len, x.device, target_dtype, offset=0) h_text = self.token_embeddings(x) + pe_text hub = ( override_hub.to(dtype=target_dtype) if override_hub is not None else self._init_hub_base(batch_size, query_rep=h_text).to( dtype=target_dtype ) ) h_initial = torch.cat([hub, h_text], dim=1) h = h_initial hub_zeros = torch.zeros( 1, self.hub_size, self.dim, device=x.device, dtype=target_dtype ) poly_pe_full = torch.cat([hub_zeros, pe_text], dim=1) new_past_c_kv_list = [] for i, layer in enumerate(self.layers): if ( self.training and (self.checkpoint_every_n > 0) and (i % self.checkpoint_every_n == 0) ): def make_checkpoint_fn(l_mod): def forward_fn(hidden_states, pe, m): return l_mod( hidden_states, poly_pe=pe, attn_mask=m, is_causal=True, is_dense_with_hub=True, ) return forward_fn h, layer_c_kv, _ = cp.checkpoint( make_checkpoint_fn(layer), h, poly_pe_full, attn_mask, use_reentrant=False, ) else: h, layer_c_kv, _ = layer( h, poly_pe=poly_pe_full, attn_mask=attn_mask, is_causal=True, is_dense_with_hub=True, ) new_past_c_kv_list.append(layer_c_kv) hub_l0 = h_initial[:, : self.hub_size, :] hub_ln = h[:, : self.hub_size, :] hub_anchored = hub_ln + self.alpha_anchor * self.w_anchor( self.norm(hub_l0) ) text_evolved = self.norm(h[:, self.hub_size :, :]) out_logits = self.head(text_evolved) if return_logits else text_evolved z_loss = self.compute_hub_diversity_loss(hub_anchored) return ( out_logits, torch.tensor(0.0, device=x.device, dtype=target_dtype), z_loss, new_past_c_kv_list, None, ) else: n_past = past_c_kv_list[0].shape[1] pos_offset = max(0, n_past - self.hub_size) pe_step = self.poly_hope( seq_len, x.device, target_dtype, offset=pos_offset ) h = self.token_embeddings(x) + pe_step new_past_c_kv_list = [] for i, layer in enumerate(self.layers): past_layer_c_kv = past_c_kv_list[i] if ( self.training and (self.checkpoint_every_n > 0) and (i % self.checkpoint_every_n == 0) ): def make_checkpoint_fn_kv(l_mod, p_ckv): def forward_fn(hidden_states, pe, m): return l_mod( hidden_states, poly_pe=pe, attn_mask=m, is_causal=True, past_c_kv=p_ckv, is_dense_with_hub=False, ) return forward_fn h, layer_c_kv, _ = cp.checkpoint( make_checkpoint_fn_kv(layer, past_layer_c_kv), h, pe_step, attn_mask, use_reentrant=False, ) else: h, layer_c_kv, _ = layer( h, poly_pe=pe_step, attn_mask=attn_mask, is_causal=True, past_c_kv=past_layer_c_kv, is_dense_with_hub=False, ) new_past_c_kv_list.append(layer_c_kv) h_normed = self.norm(h) out_logits = self.head(h_normed) if return_logits else h_normed return ( out_logits, torch.tensor(0.0, device=x.device, dtype=target_dtype), torch.tensor(0.0, device=x.device, dtype=target_dtype), new_past_c_kv_list, None, ) def forward( self, x: torch.Tensor, labels: Optional[torch.Tensor] = None, scaler: Optional[torch.amp.GradScaler] = None, grad_accum_steps: int = 1, past_key_values: Optional[List] = None, soliton_states: Optional[List] = None, directive_tokens: Optional[torch.Tensor] = None, override_hub: Optional[torch.Tensor] = None, is_sft: bool = False, execute_chunk_backward: bool = False, chunk_callback: Optional[Callable[[int, int], None]] = None, attn_mask: Optional[torch.Tensor] = None, **kwargs, ) -> XoneLMOutput: if past_key_values is not None or x.shape[1] == 1: if chunk_callback is not None: chunk_callback(1, 1) logits, aux_l, z_l, new_kv, new_sol = self._forward_dense( x, past_c_kv_list=past_key_values, soliton_state_list=soliton_states, return_logits=True, attn_mask=attn_mask, ) return XoneLMOutput( logits=logits, aux_loss=aux_l, z_loss=z_l, past_key_values=new_kv, soliton_state=new_sol, ) batch_size, seq_len = x.shape chunk_size = self.chunk_size hub = ( override_hub if override_hub is not None else self.extract_hub( x, directive_tokens=directive_tokens, is_sft=is_sft ) ) target_dtype = self.token_embeddings.weight.dtype if seq_len > chunk_size: num_chunks = math.ceil(seq_len / chunk_size) pad_len = (num_chunks * chunk_size) - seq_len x_padded = F.pad(x, (0, pad_len), value=0) if pad_len > 0 else x labels_padded = ( F.pad(labels, (0, pad_len), value=-100) if (labels is not None and pad_len > 0) else labels ) if self.training and labels is not None and execute_chunk_backward: total_lm_loss_val, total_z_loss_val, past_kv = 0.0, 0.0, None device = x.device for c_idx in range(num_chunks): if chunk_callback is not None: chunk_callback(c_idx + 1, num_chunks) chunk_tokens = x_padded[ :, c_idx * chunk_size : (c_idx + 1) * chunk_size ] chunk_labels = labels_padded[ :, c_idx * chunk_size : (c_idx + 1) * chunk_size ] is_last_chunk = c_idx == num_chunks - 1 with HardwareContext.get_autocast_context(device): if c_idx == 0: chunk_hidden, l_aux, l_z, past_kv, _ = self._forward_dense( chunk_tokens, past_c_kv_list=None, override_hub=hub, return_logits=False, attn_mask=attn_mask, ) else: detached_past_kv = [kv.detach().clone() for kv in past_kv] chunk_hidden, l_aux, l_z, past_kv, _ = self._forward_dense( chunk_tokens, past_c_kv_list=detached_past_kv, override_hub=None, return_logits=False, attn_mask=attn_mask, ) chunk_lm = self._compute_loss_efficient(chunk_hidden, chunk_labels) chunk_total = (chunk_lm + 0.01 * l_z) / ( num_chunks * grad_accum_steps ) retain_flag = not is_last_chunk if scaler is not None: scaler.scale(chunk_total).backward(retain_graph=retain_flag) else: chunk_total.backward(retain_graph=retain_flag) total_lm_loss_val += chunk_lm.item() total_z_loss_val += l_z.item() return XoneLMOutput( loss=torch.tensor(total_lm_loss_val / num_chunks, device=x.device), z_loss=torch.tensor(total_z_loss_val / num_chunks, device=x.device), past_key_values=None, soliton_state=None, ) else: logits_chunks, past_kv = [], None total_moe_loss = torch.tensor(0.0, device=x.device, dtype=target_dtype) total_z_loss = torch.tensor(0.0, device=x.device, dtype=target_dtype) chunk_losses = [] for c_idx in range(num_chunks): if chunk_callback is not None: chunk_callback(c_idx + 1, num_chunks) chunk_tokens = x_padded[ :, c_idx * chunk_size : (c_idx + 1) * chunk_size ] is_last_chunk = c_idx == num_chunks - 1 if c_idx == 0: chunk_hidden_or_logits, l_aux, l_z, past_kv, _ = ( self._forward_dense( chunk_tokens, past_c_kv_list=None, override_hub=hub, return_logits=(labels is None), attn_mask=attn_mask, ) ) else: detached_past_kv = [kv.detach().clone() for kv in past_kv] chunk_hidden_or_logits, l_aux, l_z, past_kv, _ = ( self._forward_dense( chunk_tokens, past_c_kv_list=detached_past_kv, override_hub=None, return_logits=(labels is None), attn_mask=attn_mask, ) ) if labels is not None: chunk_labels = labels_padded[ :, c_idx * chunk_size : (c_idx + 1) * chunk_size ] chunk_lm = self._compute_loss_efficient( chunk_hidden_or_logits, chunk_labels ) chunk_losses.append(chunk_lm) if is_last_chunk or labels is None: logits_chunks.append( chunk_hidden_or_logits if labels is None else self.head(chunk_hidden_or_logits) ) total_moe_loss = total_moe_loss + l_aux total_z_loss = total_z_loss + l_z final_loss = None if labels is not None: final_loss = torch.stack(chunk_losses).mean() + 0.01 * ( total_z_loss / num_chunks ) return XoneLMOutput( loss=final_loss, logits=logits_chunks[-1] if logits_chunks else None, aux_loss=total_moe_loss / num_chunks, z_loss=(total_z_loss / num_chunks) + self.compute_hub_diversity_loss(hub), past_key_values=past_kv, soliton_state=None, ) else: if chunk_callback is not None: chunk_callback(1, 1) if is_sft and labels is not None: hidden_text, aux_l, z_l, new_kv, _ = self._forward_dense( x, override_hub=hub, return_logits=False, attn_mask=attn_mask ) shift_hidden = hidden_text[..., :-1, :].contiguous() shift_labels = labels[..., 1:].contiguous() final_loss = ( self._compute_loss_efficient(shift_hidden, shift_labels) + 0.01 * z_l ) logits = None elif self.training and labels is not None: hidden_text, aux_l, z_l, new_kv, _ = self._forward_dense( x, override_hub=hub, return_logits=False, attn_mask=attn_mask ) final_loss = ( self._compute_loss_efficient(hidden_text, labels) + 0.01 * z_l ) logits = None else: hidden_or_logits, aux_l, z_l, new_kv, _ = self._forward_dense( x, override_hub=hub, return_logits=(labels is None), attn_mask=attn_mask, ) final_loss = None if labels is not None: if is_sft: shift_h = hidden_or_logits[..., :-1, :].contiguous() shift_l = labels[..., 1:].contiguous() final_loss = ( self._compute_loss_efficient(shift_h, shift_l) + 0.01 * z_l ) else: final_loss = ( self._compute_loss_efficient(hidden_or_logits, labels) + 0.01 * z_l ) logits = self.head(hidden_or_logits) else: logits = hidden_or_logits return XoneLMOutput( loss=final_loss, logits=logits, aux_loss=aux_l, z_loss=z_l, past_key_values=new_kv, soliton_state=None, ) @torch.no_grad() def generate( self, prompt_tokens: torch.Tensor, max_new_tokens: int = 64, temperature: float = 0.7, top_k: int = 40, repetition_penalty: float = 1.15, eos_token_id: Optional[int] = None, ) -> torch.Tensor: self.eval() batch_size = prompt_tokens.shape[0] hub = self.extract_hub(prompt_tokens) out = self.forward(prompt_tokens, override_hub=hub) past_kv = out.past_key_values generated = prompt_tokens.clone() logits = out.logits[:, -1, :].clone() / max(temperature, 1e-5) if repetition_penalty != 1.0: for i in range(batch_size): for prev_token in set(generated[i].tolist()): if logits[i, prev_token] < 0: logits[i, prev_token] *= repetition_penalty else: logits[i, prev_token] /= repetition_penalty if top_k > 0: v_top, _ = torch.topk(logits, min(top_k, logits.size(-1))) logits[logits < v_top[:, [-1]]] = -float("Inf") probs = F.softmax(logits, dim=-1) cur_token = torch.multinomial(probs, num_samples=1) generated = torch.cat([generated, cur_token], dim=1) for _ in range(max_new_tokens - 1): if eos_token_id is not None and (cur_token == eos_token_id).all(): break step_out = self.forward(cur_token, past_key_values=past_kv) past_kv = step_out.past_key_values logits = step_out.logits[:, -1, :].clone() / max(temperature, 1e-5) if repetition_penalty != 1.0: for i in range(batch_size): for prev_token in set(generated[i].tolist()): if logits[i, prev_token] < 0: logits[i, prev_token] *= repetition_penalty else: logits[i, prev_token] /= repetition_penalty if top_k > 0: v_top, _ = torch.topk(logits, min(top_k, logits.size(-1))) logits[logits < v_top[:, [-1]]] = -float("Inf") probs = F.softmax(logits, dim=-1) cur_token = torch.multinomial(probs, num_samples=1) generated = torch.cat([generated, cur_token], dim=1) return generated @dataclass class SpecialTokenConfig: pad_token_id: int = 0 bos_token_id: int = 1 eos_token_id: int = 2 unk_token_id: int = 3 eod_token_id: int = 4 im_start_id: Optional[int] = None im_end_id: Optional[int] = None separator_token_id: Optional[int] = None class MultiTurnConversationFormatter: def __init__( self, tokenizer: Any, token_config: Optional[SpecialTokenConfig] = None, ): self.tokenizer = tokenizer self.config = token_config or SpecialTokenConfig() def _get_id(token_str: str) -> Optional[int]: if hasattr(tokenizer, "token_to_id"): return tokenizer.token_to_id(token_str) elif hasattr(tokenizer, "convert_tokens_to_ids"): res = tokenizer.convert_tokens_to_ids(token_str) return res if isinstance(res, int) and res >= 0 else None return None if self.config.im_start_id is None: self.config.im_start_id = _get_id("<|im_start|>") if self.config.im_end_id is None: self.config.im_end_id = _get_id("<|im_end|>") if self.config.eod_token_id is None: self.config.eod_token_id = _get_id("[EOD]") def format_conversation( self, messages: List[Dict[str, str]], max_len: Optional[int] = None ) -> Dict[str, List[int]]: input_ids = [] labels = [] def _encode_text(t: str) -> List[int]: if hasattr(self.tokenizer, "encode"): res = self.tokenizer.encode(t) return res.ids if hasattr(res, "ids") else res elif callable(self.tokenizer): return self.tokenizer(t)["input_ids"] return [] for msg in messages: role = msg["role"] content = msg["content"].strip() header_text = f"<|im_start|>{role}\n" body_text = f"{content}<|im_end|>\n" header_ids = _encode_text(header_text) body_ids = _encode_text(body_text) turn_input_ids = header_ids + body_ids input_ids.extend(turn_input_ids) if role == "assistant": turn_labels = [-100] * len(header_ids) + body_ids labels.extend(turn_labels) else: labels.extend([-100] * len(turn_input_ids)) if self.config.eod_token_id is not None: input_ids.append(self.config.eod_token_id) labels.append(self.config.eod_token_id) if max_len is not None: input_ids = input_ids[:max_len] labels = labels[:max_len] return {"input_ids": input_ids, "labels": labels} class LumiSFTCollator: def __init__(self, seq_len: int = 2048, pad_token_id: int = 0): self.seq_len = seq_len self.pad_token_id = pad_token_id def __call__( self, samples: List[Dict[str, List[int]]] ) -> Dict[str, torch.Tensor]: packed_inputs = [] packed_labels = [] cur_input_buf = [] cur_label_buf = [] for item in samples: inp = item["input_ids"] lbl = item["labels"] doc_len = len(inp) if doc_len > self.seq_len: for s in range(0, doc_len, self.seq_len): chunk_inp = inp[s : s + self.seq_len] chunk_lbl = lbl[s : s + self.seq_len] pad_sz = self.seq_len - len(chunk_inp) packed_inputs.append( torch.tensor( chunk_inp + [self.pad_token_id] * pad_sz, dtype=torch.long ) ) packed_labels.append( torch.tensor(chunk_lbl + [-100] * pad_sz, dtype=torch.long) ) else: if len(cur_input_buf) + doc_len <= self.seq_len: cur_input_buf.extend(inp) cur_label_buf.extend(lbl) else: pad_sz = self.seq_len - len(cur_input_buf) packed_inputs.append( torch.tensor( cur_input_buf + [self.pad_token_id] * pad_sz, dtype=torch.long ) ) packed_labels.append( torch.tensor(cur_label_buf + [-100] * pad_sz, dtype=torch.long) ) cur_input_buf = list(inp) cur_label_buf = list(lbl) if cur_input_buf: pad_sz = self.seq_len - len(cur_input_buf) packed_inputs.append( torch.tensor( cur_input_buf + [self.pad_token_id] * pad_sz, dtype=torch.long ) ) packed_labels.append( torch.tensor(cur_label_buf + [-100] * pad_sz, dtype=torch.long) ) return { "input_ids": torch.stack(packed_inputs), "labels": torch.stack(packed_labels), } class LumiLoaderCollator: def __init__( self, seq_len: int = 8192, eos_token_id: int = 2, pad_token_id: int = 0, ): self.seq_len = seq_len self.eos_token_id = eos_token_id self.pad_token_id = pad_token_id def __call__(self, documents: List[List[int]]) -> Dict[str, torch.Tensor]: packed_batches = [] packed_labels = [] current_buf = [] current_lbl = [] docs_sorted = sorted(documents, key=len, reverse=True) for doc in docs_sorted: doc_with_eos = doc + [self.eos_token_id] doc_len = len(doc_with_eos) if doc_len > self.seq_len: for start_idx in range(0, doc_len, self.seq_len): chunk = doc_with_eos[start_idx : start_idx + self.seq_len] pad_size = self.seq_len - len(chunk) chunk_input = chunk + [self.pad_token_id] * pad_size chunk_label = chunk + [-100] * pad_size packed_batches.append(torch.tensor(chunk_input, dtype=torch.long)) packed_labels.append(torch.tensor(chunk_label, dtype=torch.long)) else: if len(current_buf) + doc_len <= self.seq_len: current_buf.extend(doc_with_eos) current_lbl.extend(doc_with_eos) else: pad_size = self.seq_len - len(current_buf) buf_input = current_buf + [self.pad_token_id] * pad_size buf_label = current_lbl + [-100] * pad_size packed_batches.append(torch.tensor(buf_input, dtype=torch.long)) packed_labels.append(torch.tensor(buf_label, dtype=torch.long)) current_buf = list(doc_with_eos) current_lbl = list(doc_with_eos) if current_buf: pad_size = self.seq_len - len(current_buf) buf_input = current_buf + [self.pad_token_id] * pad_size buf_label = current_lbl + [-100] * pad_size packed_batches.append(torch.tensor(buf_input, dtype=torch.long)) packed_labels.append(torch.tensor(buf_label, dtype=torch.long)) return { "input_ids": torch.stack(packed_batches), "labels": torch.stack(packed_labels), }