clef-flash-ternary-mlx / packed_format.py
Jakevin's picture
v2.0: measured-KL mixed-bit allocation (3.13 GB)
2d9e93a verified
Raw History Blame Contribute Delete
6.51 kB
"""Packed on-disk format for quantized Clef-Flash and its decoder.
v1.0 (`clef-ternary-v1`):
tern : <p>.trits uint8 (5 trits per byte, base 3, code+1 in {0,1,2}), <p>.scales bf16 [N, K/64], <p>.shape int32
dequantized w = scale * code, code in {-1, 0, +1}
affine4 : <p>.codes uint8 (two 4-bit codes per byte), <p>.scales / <p>.biases bf16 [N, K/64], <p>.shape int32
dequantized w = scale * code + bias
CLoQ : optional <p>.lora_a bf16 [r, K], <p>.lora_b bf16 [N, r]; y = W_q x + B (A x), kept unfolded
v2.0 (`clef-mkl-v1`), same tern/affine4 layout, plus:
bin : <p>.bits uint8 (8 codes/byte, 0/1), <p>.scales bf16 [N, K/64], <p>.shape int32
load as 2-bit QuantizedLinear, w = code*(2s) + (-s)
affine3 : <p>.q uint32 bitstream, <p>.scales / <p>.biases bf16, <p>.shape
native 3-bit QuantizedLinear
Token embedding / lm_head: <name>.q / .scales / .biases = mlx affine 4-bit group 64.
Every other tensor (norms, DeltaNet conv / A_log / dt_bias / in_proj_a / in_proj_b) is the release's
bf16 tensor byte for byte. The joint-schema head ships as the release's own file.
"""
import json
import numpy as np
import torch
import mlx.core as mx
import mlx.nn as nn
from clef_stream import (LM, DecoderLayer, QLoRA, qlinear, sanitize, set_module, FinalNorm, bf16_np_to_mx)
# (de)packing helpers, identical to ternary_rotation.py (copied so the loader needs only mlx, mlx-lm, torch, numpy)
PW3 = torch.tensor([1, 3, 9, 27, 81], dtype=torch.int32)
def pack_trits(codes):
c = codes.flatten().to(torch.int32) + 1
c = torch.cat([c, torch.ones((-c.numel()) % 5, dtype=torch.int32)]).reshape(-1, 5)
return (c * PW3).sum(1).to(torch.uint8)
def unpack_trits(b, n):
v = b.to(torch.int32)[:, None]
return ((v // PW3) % 3 - 1).flatten()[:n].to(torch.int8)
def pack_bits(codes, bits):
per = 8 // bits; c = codes.flatten().to(torch.int32)
c = torch.cat([c, torch.zeros((-c.numel()) % per, dtype=torch.int32)]).reshape(-1, per)
return (c << (torch.arange(per, dtype=torch.int32) * bits)).sum(1).to(torch.uint8)
def unpack_bits(b, n, bits):
per = 8 // bits; v = b.to(torch.int32)[:, None]
return ((v >> (torch.arange(per, dtype=torch.int32) * bits)) & (2 ** bits - 1)).flatten()[:n].to(torch.uint8)
FORMAT = "clef-mkl-v1"
def to_mx_bf16(t): return mx.array(t.float().numpy()).astype(mx.bfloat16)
def pack_linear(pname, kind, codes, scales, lora=None):
"""codes: int8 in {-1,0,1} (tern) or uint8 0..15 (affine4); scales: [N,K/G] (tern) or [N,K/G,2] (affine)."""
out = {pname + ".shape": mx.array(np.array(codes.shape, dtype=np.int32))}
if kind == "tern":
out[pname + ".trits"] = mx.array(pack_trits(codes).numpy()); out[pname + ".scales"] = to_mx_bf16(scales)
else:
out[pname + ".codes"] = mx.array(pack_bits(codes, 4).numpy())
out[pname + ".scales"] = to_mx_bf16(scales[..., 0]); out[pname + ".biases"] = to_mx_bf16(scales[..., 1])
if lora is not None:
out[pname + ".lora_a"], out[pname + ".lora_b"] = lora
return out
def module_from_packed(t, pname, kind, group=64):
N, K = (int(x) for x in np.array(t[pname + ".shape"]))
sc = torch.from_numpy(np.array(t[pname + ".scales"].astype(mx.float32)))
if kind == "tern":
c = (unpack_trits(torch.from_numpy(np.array(t[pname + ".trits"])), N * K).to(torch.int16) + 1).to(torch.uint8).reshape(N, K)
q = qlinear(c, sc, -sc, 2, group)
elif kind == "bin":
c = unpack_bits(torch.from_numpy(np.array(t[pname + ".bits"])), N * K, 1).reshape(N, K)
q = qlinear(c, 2 * sc, -sc, 2, group)
elif kind == "affine4":
c = unpack_bits(torch.from_numpy(np.array(t[pname + ".codes"])), N * K, 4).reshape(N, K)
bi = torch.from_numpy(np.array(t[pname + ".biases"].astype(mx.float32)))
q = qlinear(c, sc, bi, 4, group)
elif kind == "affine3":
q = nn.QuantizedLinear(K, N, bias=False, group_size=group, bits=3)
q.weight = t[pname + ".q"]
q.scales = t[pname + ".scales"]
q.biases = t[pname + ".biases"]
return q
else:
raise ValueError(kind)
if pname + ".lora_a" in t:
return QLoRA(q, t[pname + ".lora_a"], t[pname + ".lora_b"])
return q
def build_layer_from_tensors(args, i, t, kinds):
pre = f"{LM}layers.{i}."
layer = DecoderLayer(args, i)
layer.eval() # Metal GatedDeltaNet kernel (training mode would use the reference ops)
qnames = [p for p in kinds if p.startswith(pre)]
for p in qnames:
set_module(layer, p[len(pre):-len(".weight")], module_from_packed(t, p, kinds[p]))
skip = tuple(p + "." for p in qnames)
rest = [(k[len(pre):], sanitize(k[len(pre):], v)) for k, v in t.items() if k.startswith(pre) and not k.startswith(skip)]
layer.load_weights(rest, strict=False)
from mlx.utils import tree_flatten
have = {k for k, _ in tree_flatten(layer.parameters())}
qpre = tuple(p[len(pre):-len(".weight")] + "." for p in qnames)
missing = {k for k in have if not k.startswith(qpre)} - {k for k, _ in rest}
assert not missing, (i, sorted(missing)[:5])
mx.eval(layer.parameters())
return layer
class PackedBackbone:
"""Resident text backbone decoded from a packed file (quantized linears as mlx QuantizedLinear)."""
def __init__(self, path, args):
t, meta = mx.load(str(path), return_metadata=True)
self.meta = meta; kinds = json.loads(meta["kinds"])
self.args, self.n = args, args.num_hidden_layers
self.layers = [build_layer_from_tensors(args, i, t, kinds) for i in range(self.n)]
E = LM + "embed_tokens.weight"
if E + ".q" in t: self.eq = (t[E + ".q"], t[E + ".scales"], t[E + ".biases"]); self.ebf = None
else: self.eq = None; self.ebf = t[E]
self.lm_head_q = (t["lm_head.weight.q"], t["lm_head.weight.scales"], t["lm_head.weight.biases"]) if "lm_head.weight.q" in t else None
self.norm = FinalNorm(sanitize("norm.weight", t[LM + "norm.weight"]), args.rms_norm_eps)
mx.eval([x for x in (self.eq, self.ebf, self.lm_head_q, self.norm.w) if x is not None])
del t
def __call__(self, i): return self.layers[i]
def embed(self, ids): return embed_rows(self.eq, self.ebf, ids)
def embed_rows(eq, ebf, ids):
ix = mx.array(np.asarray(ids, dtype=np.int32))
if eq is None: return ebf[ix][None]
wq, s, b = eq
return mx.dequantize(wq[ix], s[ix], b[ix], group_size=64, bits=4).astype(mx.bfloat16)[None]