#%%writefile JiRackTernaryUltra_1p5b.py # ============================================================================= # COPYRIGHT © 2026 Konstantin Vladimirovich Grabko. ALL RIGHTS RESERVED. # JiRack Ultra Ternary Transformer # # CMS Manhattan JiRack Technology — PATENT PENDING # # This code is proprietary. # Personal and non-commercial research use is allowed. # Any commercial use, derivative works for profit, or distribution # requires a paid license and 5% royalty. # # Unauthorized commercial use is strictly prohibited. # Contact: grabko@cmsmanhattan.com # ============================================================================= import math import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.checkpoint import checkpoint # ==================== CONFIG CONSTANTS [DS1.5-1] ==================== VOCAB_SIZE = 151936 HIDDEN_SIZE = 1536 INTERMEDIATE_SIZE = 8960 NUM_LAYERS = 28 NUM_HEADS = 12 NUM_KV_HEADS = 2 HEAD_DIM = 128 # 12 * 128 = 1536 = hidden (q); kv dim = 2*128 = 256 MAX_SEQ_LEN = 4096 # [DS-6] raise for long-context (ckpt supports 131072) ROPE_THETA = 10000.0 # Qwen2.5-Math value (NOT Llama-3's 500000) RMS_EPS = 1e-6 # Qwen2 value (NOT Llama-3's 1e-5) ROPE_SCALE_FACTOR = 1.0 INIT_STD = 0.02 ATTN_QKV_BIAS = True # [DS-2] Qwen2: bias on q/k/v only # ================================================================= # [FIX-6] Feature-detect native GQA support in SDPA (PyTorch >= 2.5). def _detect_sdpa_gqa() -> bool: try: q = torch.zeros(1, 2, 1, 8) kv = torch.zeros(1, 1, 1, 8) F.scaled_dot_product_attention(q, kv, kv, enable_gqa=True) return True except TypeError: return False except Exception: return False _SDPA_HAS_GQA = _detect_sdpa_gqa() class JiRackConfig: def __init__(self): self.vocab_size = VOCAB_SIZE self.hidden_size = HIDDEN_SIZE self.intermediate_size = INTERMEDIATE_SIZE self.num_hidden_layers = NUM_LAYERS self.num_attention_heads = NUM_HEADS self.num_key_value_heads = NUM_KV_HEADS self.head_dim = HEAD_DIM self.max_seq_len = MAX_SEQ_LEN self.rope_theta = ROPE_THETA self.rms_norm_eps = RMS_EPS self.rope_scale_factor = ROPE_SCALE_FACTOR self.init_std = INIT_STD self.attn_qkv_bias = ATTN_QKV_BIAS # ==================== RoPE — HALF-SPLIT (HF convention) [DS-3] ==================== def precompute_freqs_cis( dim: int, end: int, theta: float = ROPE_THETA, scale_factor: float = ROPE_SCALE_FACTOR, ): """cos/sin of shape (end, dim), HF half-split layout: the (dim/2) frequency vector is CONCATENATED with itself (torch.cat), not interleaved. Matches transformers' LlamaRotaryEmbedding/Qwen2.""" freqs = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim)) if scale_factor > 1.0: freqs = freqs / scale_factor t = torch.arange(end, dtype=torch.float32) freqs = torch.outer(t, freqs) # (end, dim/2) emb = torch.cat((freqs, freqs), dim=-1) # (end, dim) — half-split return torch.cos(emb), torch.sin(emb) def rotate_half(x): """HF convention: (-x2, x1) where x1/x2 are the two HALVES of head_dim.""" x1 = x[..., : x.shape[-1] // 2] x2 = x[..., x.shape[-1] // 2:] return torch.cat((-x2, x1), dim=-1) def apply_rotary_emb(xq, xk, cos, sin): """Half-split RoPE, identical math to transformers.apply_rotary_pos_emb. cos/sin: (T, head_dim); q/k: (B, H, T, head_dim).""" cos = cos[None, None, :, :] sin = sin[None, None, :, :] xq_out = (xq * cos) + (rotate_half(xq) * sin) xk_out = (xk * cos) + (rotate_half(xk) * sin) return xq_out, xk_out class BitLinear(nn.Linear): """BitNet b1.58-style fake-quant linear with lambda warmup. Identical to the 10B version ([FIX-1..5] preserved); bias — when present ([DS-2]) — stays full precision ([DS-5]).""" def __init__(self, in_features, out_features, bias=False): super().__init__(in_features, out_features, bias=bias) self.eps = 1e-5 # [FIX-4] Buffer -> saved in state_dict, survives checkpoint resume. self.register_buffer("lambda_", torch.zeros(()), persistent=True) def forward(self, x: torch.Tensor) -> torch.Tensor: # Fast path (exact at lambda=0 by continuity). if not self.training and float(self.lambda_) < 1e-6: return F.linear(x, self.weight, self.bias) lam = self.lambda_.to(x.dtype) # === Weights: per-tensor absmean ternary (b1.58) === w = self.weight gamma = w.float().abs().mean().clamp(min=self.eps).to(w.dtype) # [FIX-5] w_quant = torch.clamp(torch.round(w / gamma), -1.0, 1.0) * gamma w_effective = w + lam * (w_quant - w).detach() # === Activations: per-token absmax int8 ([FIX-2],[FIX-3]) === x_scale = 127.0 / x.abs().max(dim=-1, keepdim=True).values.clamp(min=self.eps) x_quant = torch.clamp(torch.round(x * x_scale), -128.0, 127.0) / x_scale x_effective = x + lam * (x_quant - x).detach() # [FIX-1] Dequantized operands -> no post-matmul rescale. # [DS-5] bias added in full precision by F.linear. return F.linear(x_effective, w_effective, self.bias) class RMSNorm(nn.Module): def __init__(self, dim, eps=RMS_EPS): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim)) def forward(self, x): # [FIX-5] Compute statistics in fp32, cast back to input dtype. dtype = x.dtype x = x.float() x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) return (x * self.weight.float()).to(dtype) class TransformerBlock(nn.Module): def __init__(self, config, use_checkpoint=False): super().__init__() self.use_checkpoint = use_checkpoint self.n_heads = config.num_attention_heads self.n_kv_heads = config.num_key_value_heads self.head_dim = config.head_dim self.n_rep = self.n_heads // self.n_kv_heads self.norm1 = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.norm2 = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) qkv_bias = config.attn_qkv_bias # [DS-2] self.q_proj = BitLinear(config.hidden_size, self.n_heads * self.head_dim, bias=qkv_bias) self.k_proj = BitLinear(config.hidden_size, self.n_kv_heads * self.head_dim, bias=qkv_bias) self.v_proj = BitLinear(config.hidden_size, self.n_kv_heads * self.head_dim, bias=qkv_bias) self.out_proj = BitLinear(self.n_heads * self.head_dim, config.hidden_size, bias=False) self.ffn_w1 = BitLinear(config.hidden_size, config.intermediate_size, bias=False) # gate self.ffn_w3 = BitLinear(config.hidden_size, config.intermediate_size, bias=False) # up self.ffn_w2 = BitLinear(config.intermediate_size, config.hidden_size, bias=False) # down def forward(self, x, freqs_cos, freqs_sin): if self.use_checkpoint and self.training: return checkpoint( self._forward_impl, x, freqs_cos, freqs_sin, use_reentrant=False ) return self._forward_impl(x, freqs_cos, freqs_sin) def _forward_impl(self, x, freqs_cos, freqs_sin): h = self.norm1(x) B, T, _ = h.shape q = self.q_proj(h).view(B, T, self.n_heads, self.head_dim).transpose(1, 2) k = self.k_proj(h).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2) v = self.v_proj(h).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2) q, k = apply_rotary_emb(q, k, freqs_cos, freqs_sin) # [DS-3] half-split if self.n_rep > 1 and _SDPA_HAS_GQA: # [FIX-6] attn_out = F.scaled_dot_product_attention( q, k, v, is_causal=True, enable_gqa=True ) else: if self.n_rep > 1: k = k.repeat_interleave(self.n_rep, dim=1) v = v.repeat_interleave(self.n_rep, dim=1) attn_out = F.scaled_dot_product_attention(q, k, v, is_causal=True) attn_out = attn_out.transpose(1, 2).contiguous().view(B, T, -1) x = x + self.out_proj(attn_out) m = self.norm2(x) gate = F.silu(self.ffn_w1(m)) up = self.ffn_w3(m) x = x + self.ffn_w2(gate * up) return x class JiRackTransformer(nn.Module): def __init__(self, config: JiRackConfig = None, use_checkpoint=False): super().__init__() self.config = config if config is not None else JiRackConfig() self.use_checkpoint = use_checkpoint self.token_emb = nn.Embedding(self.config.vocab_size, self.config.hidden_size) self.blocks = nn.ModuleList([ TransformerBlock(self.config, self.use_checkpoint) for _ in range(self.config.num_hidden_layers) ]) self.ln_f = RMSNorm(self.config.hidden_size, eps=self.config.rms_norm_eps) # [DS1.5-2] tie_word_embeddings = False (verified) -> separate lm_head. self.lm_head = nn.Linear(self.config.hidden_size, self.config.vocab_size, bias=False) cos, sin = precompute_freqs_cis( dim=self.config.head_dim, end=self.config.max_seq_len, theta=self.config.rope_theta, scale_factor=self.config.rope_scale_factor, ) self.register_buffer("freqs_cos", cos, persistent=False) self.register_buffer("freqs_sin", sin, persistent=False) # [FIX-8] Only relevant when training from scratch; harmless before # load_hf_state_dict() overwrites everything. self._init_weights() def _init_weights(self): std = self.config.init_std resid_std = std / math.sqrt(2 * self.config.num_hidden_layers) nn.init.normal_(self.token_emb.weight, mean=0.0, std=std) nn.init.normal_(self.lm_head.weight, mean=0.0, std=std) for block in self.blocks: for lin in (block.q_proj, block.k_proj, block.v_proj, block.ffn_w1, block.ffn_w3): nn.init.normal_(lin.weight, mean=0.0, std=std) if lin.bias is not None: nn.init.zeros_(lin.bias) for lin in (block.out_proj, block.ffn_w2): nn.init.normal_(lin.weight, mean=0.0, std=resid_std) if lin.bias is not None: nn.init.zeros_(lin.bias) # ---------------- lambda warmup hooks (unchanged) ---------------- def set_lambda(self, lambda_value: float): for module in self.modules(): if isinstance(module, BitLinear): module.lambda_.fill_(lambda_value) def get_lambda(self) -> float: for module in self.modules(): if isinstance(module, BitLinear): return float(module.lambda_) return 0.0 def forward(self, input_ids): seq_len = input_ids.shape[1] x = self.token_emb(input_ids) cos = self.freqs_cos[:seq_len].to(device=x.device, dtype=x.dtype) sin = self.freqs_sin[:seq_len].to(device=x.device, dtype=x.dtype) for block in self.blocks: x = block(x, cos, sin) return self.lm_head(self.ln_f(x)) # ------------------------------------------------------------------ # [DS-4] HF Qwen2 -> JiRack weight mapping. # Usage: # from transformers import AutoModelForCausalLM # hf = AutoModelForCausalLM.from_pretrained( # "deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B", torch_dtype=torch.float32) # model.load_hf_state_dict(hf.state_dict()) # or load safetensors shards directly and merge them into one dict. # No RoPE permutation is needed: this model now uses the same # half-split rotation as HF ([DS-3]). # ------------------------------------------------------------------ @torch.no_grad() def load_hf_state_dict(self, hf_sd: dict, strict: bool = True): mapped = {} mapped["token_emb.weight"] = hf_sd["model.embed_tokens.weight"] mapped["ln_f.weight"] = hf_sd["model.norm.weight"] if "lm_head.weight" in hf_sd: mapped["lm_head.weight"] = hf_sd["lm_head.weight"] else: # fallback for third-party re-uploads that strip lm_head # (official 1.5B distill DOES ship lm_head.weight — untied) mapped["lm_head.weight"] = hf_sd["model.embed_tokens.weight"] for i in range(self.config.num_hidden_layers): hf = f"model.layers.{i}" jr = f"blocks.{i}" mapped[f"{jr}.norm1.weight"] = hf_sd[f"{hf}.input_layernorm.weight"] mapped[f"{jr}.norm2.weight"] = hf_sd[f"{hf}.post_attention_layernorm.weight"] mapped[f"{jr}.q_proj.weight"] = hf_sd[f"{hf}.self_attn.q_proj.weight"] mapped[f"{jr}.k_proj.weight"] = hf_sd[f"{hf}.self_attn.k_proj.weight"] mapped[f"{jr}.v_proj.weight"] = hf_sd[f"{hf}.self_attn.v_proj.weight"] mapped[f"{jr}.q_proj.bias"] = hf_sd[f"{hf}.self_attn.q_proj.bias"] mapped[f"{jr}.k_proj.bias"] = hf_sd[f"{hf}.self_attn.k_proj.bias"] mapped[f"{jr}.v_proj.bias"] = hf_sd[f"{hf}.self_attn.v_proj.bias"] mapped[f"{jr}.out_proj.weight"] = hf_sd[f"{hf}.self_attn.o_proj.weight"] mapped[f"{jr}.ffn_w1.weight"] = hf_sd[f"{hf}.mlp.gate_proj.weight"] mapped[f"{jr}.ffn_w3.weight"] = hf_sd[f"{hf}.mlp.up_proj.weight"] mapped[f"{jr}.ffn_w2.weight"] = hf_sd[f"{hf}.mlp.down_proj.weight"] missing, unexpected = self.load_state_dict(mapped, strict=False) # lambda_ buffers are OURS (not in HF) — they legitimately stay missing. real_missing = [k for k in missing if not k.endswith("lambda_")] if strict: assert not real_missing, f"missing from HF checkpoint: {real_missing}" assert not unexpected, f"unexpected keys: {unexpected}" print(f"✅ HF weights loaded: {len(mapped)} tensors " f"({len(real_missing)} missing, {len(unexpected)} unexpected)") return real_missing, unexpected # ------------------------------------------------------------------ # [FIX-9] Ternary export (biases included, full precision, [DS-5]). # ------------------------------------------------------------------ @torch.no_grad() def export_ternary_state_dict(self): out = {} for name, module in self.named_modules(): if isinstance(module, BitLinear): w = module.weight.float() gamma = w.abs().mean().clamp(min=module.eps) codes = torch.clamp(torch.round(w / gamma), -1.0, 1.0).to(torch.int8) out[f"{name}.codes"] = codes out[f"{name}.gamma"] = gamma if module.bias is not None: out[f"{name}.bias"] = module.bias.detach().clone() out["token_emb.weight"] = self.token_emb.weight.detach().clone() out["lm_head.weight"] = self.lm_head.weight.detach().clone() out["ln_f.weight"] = self.ln_f.weight.detach().clone() for name, module in self.named_modules(): if isinstance(module, RMSNorm) and name != "ln_f": out[f"{name}.weight"] = module.weight.detach().clone() return out # Convenience aliases so existing training scripts barely change: JiRackConfig = JiRackConfig # drop-in name compat (optional) JiRackTransformer = JiRackTransformer JiRackConfig = JiRackConfig # lets ds7b-style imports work JiRackTransformer = JiRackTransformer # ============================================================================= # Smoke test: python JiRackTernaryPyTorch_ds1p5b.py (tiny config, CPU, seconds) # ============================================================================= if __name__ == "__main__": class TinyConfig(JiRackConfig): def __init__(self): super().__init__() self.vocab_size = 256 self.hidden_size = 64 self.intermediate_size = 128 self.num_hidden_layers = 2 self.num_attention_heads = 4 self.num_key_value_heads = 2 self.head_dim = 16 self.max_seq_len = 64 torch.manual_seed(0) model = JiRackTransformer(TinyConfig()).eval() ids = torch.randint(0, 256, (2, 32)) with torch.no_grad(): model.set_lambda(0.0) y0 = model(ids) model.set_lambda(1e-4) y_eps = model(ids) model.set_lambda(1.0) y1 = model(ids) # 1) Continuity in lambda. rel_jump = (y_eps - y0).norm() / y0.norm() print(f"relative change at lambda=1e-4: {rel_jump.item():.2e} (must be ~1e-4)") assert rel_jump < 1e-2, "lambda warmup is not continuous!" # 2) Output scale sanity at full quantization. ratio = y1.std() / y0.std() print(f"std ratio lambda=1 vs lambda=0: {ratio.item():.3f} (must be O(1))") assert 0.1 < ratio.item() < 10.0, "output scale collapsed or exploded!" # 3) Gradients flow through STE at lambda=1 (weights AND qkv biases). model.train() model.set_lambda(1.0) loss = model(ids).float().pow(2).mean() loss.backward() g = model.blocks[0].q_proj.weight.grad gb = model.blocks[0].q_proj.bias.grad assert g is not None and torch.isfinite(g).all() and g.abs().sum() > 0 assert gb is not None and torch.isfinite(gb).all(), "qkv bias got no grad!" print(f"grad norms q_proj: weight={g.norm().item():.4f}, bias={gb.norm().item():.4f}") # 4) lambda survives a state_dict round-trip. sd = model.state_dict() model2 = JiRackTransformerDS1p5B(TinyConfig()) model2.load_state_dict(sd) assert abs(model2.get_lambda() - 1.0) < 1e-9, "lambda_ not serialized!" print("lambda serialization: OK") # 5) HF name mapping round-trip on the tiny config: build a fake HF # dict from our own weights, load it back, outputs must match. fake_hf = { "model.embed_tokens.weight": model.token_emb.weight.detach().clone(), "model.norm.weight": model.ln_f.weight.detach().clone(), "lm_head.weight": model.lm_head.weight.detach().clone(), } for i, blk in enumerate(model.blocks): p = f"model.layers.{i}" fake_hf[f"{p}.input_layernorm.weight"] = blk.norm1.weight.detach().clone() fake_hf[f"{p}.post_attention_layernorm.weight"] = blk.norm2.weight.detach().clone() fake_hf[f"{p}.self_attn.q_proj.weight"] = blk.q_proj.weight.detach().clone() fake_hf[f"{p}.self_attn.k_proj.weight"] = blk.k_proj.weight.detach().clone() fake_hf[f"{p}.self_attn.v_proj.weight"] = blk.v_proj.weight.detach().clone() fake_hf[f"{p}.self_attn.q_proj.bias"] = blk.q_proj.bias.detach().clone() fake_hf[f"{p}.self_attn.k_proj.bias"] = blk.k_proj.bias.detach().clone() fake_hf[f"{p}.self_attn.v_proj.bias"] = blk.v_proj.bias.detach().clone() fake_hf[f"{p}.self_attn.o_proj.weight"] = blk.out_proj.weight.detach().clone() fake_hf[f"{p}.mlp.gate_proj.weight"] = blk.ffn_w1.weight.detach().clone() fake_hf[f"{p}.mlp.up_proj.weight"] = blk.ffn_w3.weight.detach().clone() fake_hf[f"{p}.mlp.down_proj.weight"] = blk.ffn_w2.weight.detach().clone() model3 = JiRackTransformer(TinyConfig()).eval() model3.load_hf_state_dict(fake_hf) model3.set_lambda(0.0) model.eval(); model.set_lambda(0.0) with torch.no_grad(): y_ref = model(ids) y_map = model3(ids) assert torch.allclose(y_ref, y_map, atol=1e-5), "HF mapping mismatch!" print("HF name-mapping round-trip: OK") # 6) Export produces genuinely ternary codes + preserved biases. exported = model.export_ternary_state_dict() codes = exported["blocks.0.q_proj.codes"] assert set(codes.unique().tolist()) <= {-1, 0, 1} assert "blocks.0.q_proj.bias" in exported, "qkv bias lost in export!" print(f"export: {sum('codes' in k for k in exported)} ternary tensors, " f"biases preserved, OK") print("\nAll smoke tests passed.")