""" Modeling file for Ivme-Conversate-S-v1. Architecture ported directly from the training script (train_ivme_s_v1.py), verified correct there via extensive isolated testing during development: - Factorized, untied token embeddings (separate small-rank projections for input embedding and output head -- NOT tied, unlike most tiny LMs). - GQA (grouped-query attention), 4:1 query:kv head ratio. - DIFF attention V2 (per microsoft/unilm Diff-Transformer-V2): Q has 2x heads, K/V unchanged, single fused attention call, interleaved head split (NOT a half-split -- verified against the reference blog's explicit "Wrong Implementation" ablation warning), lambda is a per-token per-head sigmoid-projected value. - nGPT-style hypersphere normalization: weights renormalized onto the unit hypersphere after every optimizer step during training (a training-time concern, not present in this inference-only file), EXCLUDING the output head and token embedding -- confirmed via a direct overfitting test that including them creates a hard, unmovable floor on achievable loss (~2.1 on a trivially overfittable 8-token batch, vs 0.0005 when excluded). - Immediate block-wise weight sharing: `n_unique_layers` distinct blocks, each executed `share_factor` times in a row, giving an effective depth of n_unique_layers * share_factor at the parameter cost of n_unique_layers. - Learnable meta/register tokens prepended to the sequence, dropped before the output head. - RoPE positional encoding, applied at full head_dim (not split -- DIFF V2 doesn't split head_dim, unlike V1). FlashAttention-2 (via HF Kernels, pinned specifically because SDPA's FLASH_ATTENTION label was found to silently resolve to FA4 on Blackwell-class GPUs and regress for this model's shape profile) is used opportunistically when available and the GPU meets its Ampere+ compute-capability floor; otherwise this falls back to SDPA's default (unrestricted) backend selection, which works correctly on any GPU including pre-Ampere hardware, just without the fused-kernel speedup. """ import math import torch import torch.nn as nn import torch.nn.functional as F from transformers import PreTrainedModel from transformers.modeling_outputs import CausalLMOutput try: from .configuration_ivme_s_v1 import IvmeConversateSConfig except ImportError: # Fallback for non-package imports (e.g. cloning the repo and running # `import modeling_ivme_s_v1` directly rather than through HF's # trust_remote_code dynamic-module loader, which resolves the relative # import above correctly via its own transformers_modules.* packaging). from configuration_ivme_s_v1 import IvmeConversateSConfig try: from torch.nn.attention import SDPBackend, sdpa_kernel _HAS_SDPA_KERNEL_CONTEXT = True except ImportError: _HAS_SDPA_KERNEL_CONTEXT = False _HF_FLASH_ATTN2 = None _HF_FLASH_ATTN2_IMPORT_ERROR = None try: from kernels import get_kernel as _get_kernel _HF_FLASH_ATTN2 = _get_kernel("kernels-community/flash-attn2", version=2) except Exception as _e: _HF_FLASH_ATTN2_IMPORT_ERROR = _e _FA2_MIN_COMPUTE_CAPABILITY = (8, 0) # Ampere+ _fa2_capability_cache = {} def _cuda_supports_fa2(device): key = str(device) if key not in _fa2_capability_cache: try: cap = torch.cuda.get_device_capability(device) _fa2_capability_cache[key] = cap >= _FA2_MIN_COMPUTE_CAPABILITY except Exception: _fa2_capability_cache[key] = False return _fa2_capability_cache[key] # --------------------------------------------------------------------- # RoPE # --------------------------------------------------------------------- def build_rope_cache(dim, max_seq_len, base=10000.0, device="cpu"): assert dim % 2 == 0 inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2, device=device).float() / dim)) t = torch.arange(max_seq_len, device=device).float() freqs = torch.outer(t, inv_freq) emb = torch.cat([freqs, freqs], dim=-1) return emb.cos(), emb.sin() def rotate_half(x): x1, x2 = x.chunk(2, dim=-1) return torch.cat([-x2, x1], dim=-1) def apply_rope(x, cos, sin): T = x.shape[-2] # cos/sin cast to match x's dtype at the point of use (not stored that way) # -- multiplying a bf16 autocast tensor against permanently-fp32 buffers # silently upcasts the RESULT back to fp32 via normal type promotion, # which propagates downstream with no error. Confirmed by a real crash # when this reached a bf16-only FA2 kernel. cos = cos[:T].unsqueeze(0).unsqueeze(0).to(x.dtype) sin = sin[:T].unsqueeze(0).unsqueeze(0).to(x.dtype) return x * cos + rotate_half(x) * sin def l2norm(x, dim=-1, eps=1e-6): return x / (x.norm(dim=dim, keepdim=True) + eps) # --------------------------------------------------------------------- # Factorized, untied embeddings # --------------------------------------------------------------------- class FactorizedEmbedding(nn.Module): def __init__(self, vocab_size, r, d_model): super().__init__() self.embed = nn.Embedding(vocab_size, r) self.proj = nn.Linear(r, d_model, bias=False) def forward(self, ids): return self.proj(self.embed(ids)) class FactorizedHead(nn.Module): def __init__(self, vocab_size, r, d_model): super().__init__() self.proj = nn.Linear(d_model, r, bias=False) self.unembed = nn.Linear(r, vocab_size, bias=False) def forward(self, h): return self.unembed(self.proj(h)) # --------------------------------------------------------------------- # DIFF attention V2 + GQA # --------------------------------------------------------------------- class DiffGQAAttention(nn.Module): def __init__(self, d_model, n_heads, n_kv_heads, layer_idx, n_layers): super().__init__() assert d_model % n_heads == 0 assert n_heads % n_kv_heads == 0 self.n_heads = n_heads self.n_kv_heads = n_kv_heads self.head_dim = d_model // n_heads self.wq = nn.Linear(d_model, 2 * n_heads * self.head_dim, bias=False) self.wk = nn.Linear(d_model, n_kv_heads * self.head_dim, bias=False) self.wv = nn.Linear(d_model, n_kv_heads * self.head_dim, bias=False) self.wo = nn.Linear(n_heads * self.head_dim, d_model, bias=False) self.lam_proj = nn.Linear(d_model, n_heads, bias=True) def forward(self, x, rope_cos, rope_sin): B, T, D = x.shape q = self.wq(x).view(B, T, 2 * self.n_heads, self.head_dim).transpose(1, 2) k = self.wk(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2) v = self.wv(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2) q = apply_rope(q, rope_cos, rope_sin) k = apply_rope(k, rope_cos, rope_sin) q, k = l2norm(q), l2norm(k) if x.is_cuda and _HF_FLASH_ATTN2 is not None and _cuda_supports_fa2(x.device): target_dtype = torch.bfloat16 if x.dtype != torch.float16 else torch.float16 qt = q.transpose(1, 2).to(target_dtype) kt = k.transpose(1, 2).to(target_dtype) vt = v.transpose(1, 2).to(target_dtype) attn = _HF_FLASH_ATTN2.flash_attn_func(qt, kt, vt, causal=True) attn = attn.transpose(1, 2) elif x.is_cuda and _HAS_SDPA_KERNEL_CONTEXT and _cuda_supports_fa2(x.device): with sdpa_kernel([SDPBackend.FLASH_ATTENTION, SDPBackend.EFFICIENT_ATTENTION]): attn = F.scaled_dot_product_attention(q, k, v, is_causal=True, enable_gqa=True) else: # Unrestricted SDPA -- correct on any hardware (falls back to the # MATH backend where no fused kernel is available, e.g. pre-Ampere # GPUs). Confirmed necessary: restricting to fused-only backends # on such hardware leaves SDPA with nothing to fall back to and # raises "No available kernel." attn = F.scaled_dot_product_attention(q, k, v, is_causal=True, enable_gqa=True) attn = attn.transpose(1, 2) # (B, T, 2h, head_dim) attn1, attn2 = attn[:, :, 0::2, :], attn[:, :, 1::2, :] # interleaved, not halved lam_val = torch.sigmoid(self.lam_proj(x)).unsqueeze(-1) out = attn1 - lam_val * attn2 out = out.reshape(B, T, self.n_heads * self.head_dim) return self.wo(out) class SwiGLU(nn.Module): def __init__(self, d_model, d_ff): super().__init__() self.w_gate_up = nn.Linear(d_model, 2 * d_ff, bias=False) self.w_down = nn.Linear(d_ff, d_model, bias=False) self.d_ff = d_ff def forward(self, x): gate, up = self.w_gate_up(x).split(self.d_ff, dim=-1) return self.w_down(F.silu(gate) * up) class Block(nn.Module): def __init__(self, d_model, n_heads, n_kv_heads, d_ff, layer_idx, n_layers): super().__init__() self.attn = DiffGQAAttention(d_model, n_heads, n_kv_heads, layer_idx, n_layers) self.ffn = SwiGLU(d_model, d_ff) self.alpha_attn = nn.Parameter(torch.full((d_model,), 1.0 / math.sqrt(d_model))) self.alpha_ffn = nn.Parameter(torch.full((d_model,), 1.0 / math.sqrt(d_model))) def forward(self, x, rope_cos, rope_sin): h = self.attn(l2norm(x), rope_cos, rope_sin) x = l2norm(x + self.alpha_attn * (l2norm(h) - x)) h = self.ffn(l2norm(x)) x = l2norm(x + self.alpha_ffn * (l2norm(h) - x)) return x class IvmeConversateSModel(PreTrainedModel): """HF-compatible wrapper. Load with: AutoModelForCausalLM.from_pretrained(repo_id, trust_remote_code=True) """ config_class = IvmeConversateSConfig def __init__(self, config: IvmeConversateSConfig): super().__init__(config) n_eff = config.n_unique_layers * config.share_factor self.tok_embed = FactorizedEmbedding(config.vocab_size, config.embed_rank, config.d_model) self.head = FactorizedHead(config.vocab_size, config.embed_rank, config.d_model) self.meta_tokens = nn.Parameter(torch.randn(config.n_meta_tokens, config.d_model) * 0.02) self.blocks = nn.ModuleList([ Block(config.d_model, config.n_heads, config.n_kv_heads, config.d_ff, i, n_eff) for i in range(config.n_unique_layers) ]) self.execution_order = [ b for b in range(config.n_unique_layers) for _ in range(config.share_factor) ] head_dim = config.d_model // config.n_heads cos, sin = build_rope_cache(head_dim, config.max_seq_len + config.n_meta_tokens) # NOTE: persistent=True (not False). HF's from_pretrained() uses a # fast/meta-device init path by default that SKIPS real __init__ # buffer computation for non-persistent buffers -- this is a # documented transformers behavior (see huggingface/transformers # issue #33326: sinusoidal/positional buffers computed in __init__ # are "rendered completely ineffective" under this path, while # persistent buffers/weights ARE correctly restored from the # checkpoint's state_dict). Confirmed by a real bug: with # persistent=False, model.from_pretrained(model.save_pretrained(...)) # produced NaN logits because rope_cos/rope_sin were left as # uninitialized memory. persistent=True saves this small deterministic # buffer in the checkpoint and lets the normal state_dict-loading path # (which works correctly) restore it, sidestepping the meta-device # gap entirely. self.register_buffer("rope_cos", cos, persistent=True) self.register_buffer("rope_sin", sin, persistent=True) self.post_init() def get_input_embeddings(self): return self.tok_embed.embed def set_input_embeddings(self, value): self.tok_embed.embed = value def can_generate(self): # No KV-cache in this architecture -- forward() always recomputes # attention over the full sequence. .generate() would technically run # (each step re-does the full forward pass) but is O(n^2) rather than # the O(n) a cached model gets, so it's slow, not broken. True either # way; documented here rather than silently pretending otherwise. return True def forward(self, input_ids, labels=None, **kwargs): B, T = input_ids.shape tok = self.tok_embed(input_ids) meta = self.meta_tokens.unsqueeze(0).expand(B, -1, -1) x = torch.cat([meta, tok], dim=1) x = l2norm(x) for b_idx in self.execution_order: x = self.blocks[b_idx](x, self.rope_cos, self.rope_sin) x = x[:, self.config.n_meta_tokens:, :] logits = self.head(x) loss = None if labels is not None: loss = F.cross_entropy( logits[:, :-1, :].reshape(-1, self.config.vocab_size), labels[:, 1:].reshape(-1), ) return CausalLMOutput(loss=loss, logits=logits)