470 lines
20 KiB
Python
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.")
|