"""Packed on-disk format for quantized Clef-Flash and its decoder. v1.0 (`clef-ternary-v1`): tern :

.trits uint8 (5 trits per byte, base 3, code+1 in {0,1,2}),

.scales bf16 [N, K/64],

.shape int32 dequantized w = scale * code, code in {-1, 0, +1} affine4 :

.codes uint8 (two 4-bit codes per byte),

.scales /

.biases bf16 [N, K/64],

.shape int32 dequantized w = scale * code + bias CLoQ : optional

.lora_a bf16 [r, K],

.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 :

.bits uint8 (8 codes/byte, 0/1),

.scales bf16 [N, K/64],

.shape int32 load as 2-bit QuantizedLinear, w = code*(2s) + (-s) affine3 :

.q uint32 bitstream,

.scales /

.biases bf16,

.shape native 3-bit QuantizedLinear Token embedding / lm_head: .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]