#%%writefile JiRackTernaryPyTorch_10b.py # ============================================================================= # COPYRIGHT © 2025-2026 Konstantin Vladimirovich Grabko. ALL RIGHTS RESERVED. # 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. # # Rope fix, stable 99th quantile, and integration of smooth Quantization Warmup. # # Unauthorized commercial use is strictly prohibited. # Contact: grabko@cmsmanhattan.com # ============================================================================= import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.checkpoint import checkpoint # ========================= CONFIG CONSTANTS ========================= VOCAB_SIZE = 128256 HIDDEN_SIZE = 4096 INTERMEDIATE_SIZE = 18432 NUM_LAYERS = 32 NUM_HEADS = 32 # 4096 // 32 = 128 (HEAD_DIM) NUM_KV_HEADS = 8 HEAD_DIM = 128 MAX_SEQ_LEN = 8192 ROPE_THETA = 500000.0 RMS_EPS = 1e-5 ROPE_SCALE_FACTOR = 1.0 # ================================================================= class JiRackConfig10B: 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 def precompute_freqs_cis(dim: int, end: int, theta: float = ROPE_THETA, scale_factor: float = ROPE_SCALE_FACTOR): 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) return torch.cos(freqs), torch.sin(freqs) def apply_rotary_emb(xq, xk, freqs_cos, freqs_sin): def rotate_interleaved(x): x_even = x[..., 0::2] x_odd = x[..., 1::2] return torch.stack((-x_odd, x_even), dim=-1).flatten(-2) cos = freqs_cos[None, None, :, :].repeat_interleave(2, dim=-1) sin = freqs_sin[None, None, :, :].repeat_interleave(2, dim=-1) xq_out = (xq * cos) + (rotate_interleaved(xq) * sin) xk_out = (xk * cos) + (rotate_interleaved(xk) * sin) return xq_out, xk_out class BitLinear(nn.Linear): def __init__(self, in_features, out_features, bias=False, ternary=True): super().__init__(in_features, out_features, bias=bias) self.ternary = ternary self.eps = 1e-5 self.lambda_ = 0.0 # Интеграция плавного прогрева квантования def forward(self, x: torch.Tensor) -> torch.Tensor: if not self.ternary: return F.linear(x, self.weight, self.bias) # Если lambda почти ноль — работаем полностью в full precision, сохраняя когерентность на старте if self.lambda_ < 1e-6: return F.linear(x, self.weight, self.bias) # Троичное квантование весов (Ternary Weights) с плавной интерполяцией w = self.weight gamma = w.abs().mean().clamp(min=self.eps) w_quant = torch.clamp(torch.round(w / gamma), -1.0, 1.0) w_effective = w + self.lambda_ * (w_quant * gamma - w).detach() # Нормализация активаций x_mean = x.mean(dim=-1, keepdim=True) x_variance = x.var(dim=-1, keepdim=True, unbiased=False) x_norm = (x - x_mean) / torch.sqrt(x_variance + self.eps) # Стабильное квантование активаций через 99-й квантиль с плавной интерполяцией x_abs = x_norm.abs() x_quantile = torch.quantile(x_abs.float(), 0.99, dim=-1, keepdim=True).to(x_norm.dtype) x_scale_bound = x_quantile.clamp(min=self.eps) x_scale = 127.0 / x_scale_bound x_quant = torch.clamp(torch.round(x_norm * x_scale), -128.0, 127.0) x_effective = x_norm + self.lambda_ * (x_quant / x_scale - x_norm).detach() # Линейное преобразование out = F.linear(x_effective, w_effective, self.bias) # Декуантование обратно (масштабируется пропорционально вкладу квантования) x_std = torch.sqrt(x_variance + self.eps) # Динамический масштаб декуантования подстраивается под текущий уровень lambda_ scale_factor = (x_std * gamma / 127.0) * self.lambda_ + (1.0 - self.lambda_) return out * scale_factor 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): return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) * self.weight class TransformerBlock(nn.Module): def __init__(self, config, use_checkpoint=False, bias=False , ternary=False): super().__init__() self.ternary=ternary self.bias=bias 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) self.q_proj = BitLinear(config.hidden_size, config.hidden_size, self.bias,self.ternary) self.k_proj = BitLinear(config.hidden_size, self.n_kv_heads * self.head_dim, self.bias,self.ternary) self.v_proj = BitLinear(config.hidden_size, self.n_kv_heads * self.head_dim, self.bias,self.ternary) self.out_proj = BitLinear(config.hidden_size, config.hidden_size, self.bias,self.ternary) self.ffn_w1 = BitLinear(config.hidden_size, config.intermediate_size, self.bias,self.ternary) self.ffn_w3 = BitLinear(config.hidden_size, config.intermediate_size, self.bias,self.ternary) self.ffn_w2 = BitLinear(config.intermediate_size, config.hidden_size, self.bias,self.ternary) 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) 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 JiRackTransformer10B(nn.Module): def __init__(self, config: JiRackConfig10B = None, use_checkpoint=False, bias=False , ternary=True): super().__init__() self.config = config if config is not None else JiRackConfig10B() self.use_checkpoint = use_checkpoint self.ternary=ternary self.bias=bias self.token_emb = nn.Embedding(self.config.vocab_size, self.config.hidden_size) self.blocks = nn.ModuleList([ TransformerBlock(self.config, self.use_checkpoint, self.bias , self.ternary) for _ in range(self.config.num_hidden_layers) ]) self.ln_f = RMSNorm(self.config.hidden_size, eps=self.config.rms_norm_eps) 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) def set_lambda(self, lambda_value: float): """Устанавливает lambda_ для всех BitLinear слоёв модели глобально""" for module in self.modules(): if isinstance(module, BitLinear): module.lambda_ = lambda_value def _set_ternary(self, module): if isinstance(module, BitLinear): module.ternary = True def forward(self, input_ids): seq_len = input_ids.shape[1] x = self.token_emb(input_ids) # RoPE буферы приводятся к типу и девайсу входящих эмбеддингов cos = self.freqs_cos[:seq_len].to(x) sin = self.freqs_sin[:seq_len].to(x) for block in self.blocks: x = block(x, cos, sin) return self.lm_head(self.ln_f(x))