Files
JiRackUltra_1b/JiRackTernaryUltra_1b.py
ModelHub XC 42aa1b5282 初始化项目,由ModelHub XC社区提供模型
Model: CMSManhattan/JiRackUltra_1b
Source: Original Platform
2026-08-25 05:59:19 +08:00

470 lines
20 KiB
Python

#%%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.")