File size: 41,640 Bytes
a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 1e100d1 a0dc6f0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 | #!/usr/bin/env python3
"""Standalone decoder for the COMPACT-PACKED Qwen3.6-35B-A3B quant (Mach-1-Small).
NO modal. PURE NUMPY for every tier -- experts, NE linears, lm_head, embedding --
with no torch and no GPL'd code on the default path. This is the SINGLE canonical
decoder: both pack jobs (trellis_experts_canon_aw.py pack arms, trellis_ne_pack.py)
import from here and assert their packed output round-trips through THESE functions,
so the shipped decoder is the one that is actually verified.
TWO codecs, per block:
- EXPERTS: rotated trellis codes (sign/scale side-streams per expert); two rate
populations per the file manifest.
- NE linears + lm_head: transform-free trellis codes with a per-matrix fp16 LUT and
per-group fp16 scales; self-describing per-shard manifests.
- Embed: int3 asymmetric group codes.
Packed layout (HF repo):
packed/experts/L{LL}.safetensors -- per expert e, proj in gate/up/down:
e{e}.{proj}.trellis int16 [ntiles, 16*K] (pack_trellis bitstream, ~K bpw)
e{e}.{proj}.su fp16 [n]
e{e}.{proj}.sv fp16 [m] (Wscale absorbed)
or (demoted K1 experts)
e{e}.{proj}.SU int8 [n]; e{e}.{proj}.SV int8 [m]; e{e}.{proj}.Wscale fp16 []
plus, when the manifest carries "basis" (rung-3 shared cold-expert low-rank basis):
basis.{proj}.A fp16 [r, n]; basis.{proj}.B fp16 [m, r]
e{e}.{proj}.c fp16 [r] (demoted experts only)
metadata["manifest"]: {"cb2":..., "cb1":..., "demoted":[e,...], "geom":{proj:[m,n]},
optional "basis": {"r":r, "dtype":"fp16", "shared_A_gu":false}}
packed/experts/codebook.safetensors: tlut fp16 [512, 2] (shared by K2/K1 — the tlut
depends only on (tlut_bits=9, V=2); K only changes trellis transitions).
packed/ne/L{LL}.safetensors (zero-padded LNN -- the LIVE l64 spine tier) -- canon
int-lattice codec "canon_rht_bitshift_trellis_intlattice"; NO manifest key. Per-file
metadata: cb_params {K:4,L:16,V:2,tlut_bits:9,quantlut_sym,td 16x16} + dims
{name: [m0, n0, m, n]}. Keys per tensor:
{name}|trellis int16, {name}|SU int8 [n], {name}|SV int8 [m], {name}|Wscale fp16
tier-shared codebook: packed/ne/tlut.safetensors tlut fp16 [512, 2]
packed/ne/{L{i}|head_c{c}of8}.safetensors (single-digit L{i}: legacy builds only) --
transform-free tier:
{name}|packed uint8 [B, T*K/8], {name}|gscale fp16 [B, T/128], {name}|lut fp16 [4096]
metadata["manifest"]: {"codec":"lloyd_bitshift_trellis","L":12,"group":128,
"pattern":[K]*4,"tensors":{name:{shape:[m0,n0],transposed:bool}}}
packed/head/head_c{c}of8.safetensors (LIVE head tier) -- codec "int5g64_packed":
LMHEADCHUNK:{r0}:{r1}|qp uint8 [rows, n/8*5], |gscale fp16 [rows, n/64]
(+ optional |prot_rows int32 / |prot_dense exact rows); metadata dims {name:[rows,n]}
packed/ne/embed_int4.safetensors (LIVE embed tier; decode_embed(bits=4)) --
affine int4-g64: q_packed uint8 [rows, hid/2], mn/mx fp16 [rows, hid/64];
bf16 cast of the decode is bit-exact vs the served embedding.
packed/ne/embed_packed.safetensors (superseded 8-bpw lossless Lloyd-LUT container;
decode_embed_packed) and embed_int3.safetensors (legacy int3):
q_packed uint8 [rows, hid*3/8], mn fp16 [rows, hid/64], mx fp16 [rows, hid/64]
Usage:
from decode import (decode_trellis, decode_expert_layer, decode_ne_shard,
decode_head, decode_embed_packed, decode_embed)
License note: the default decode path is self-contained numpy written from the
mathematical spec (documented inline below) — NO GPL dependency. Setting
QTIP_LEGACY_DECODE=1 switches decode_trellis to the upstream Cornell QTIP lib
(GPLv3, resolved via QTIP_DIR) purely as a cross-check oracle; the two paths are
gated bit-exact (bf16 AND fp32) across all 40 expert layers (gate_np_decode.py).
"""
import json
import math
import os
import sys
import numpy as np
HIDDEN, INTER, NEXP, NLAYERS = 2048, 512, 256, 40
NHEADC = 8
CB2 = dict(K=2, L=16, V=2, tlut_bits=9, decode_mode="quantlut_sym", td_x=16, td_y=16)
CB1 = dict(K=1, L=16, V=2, tlut_bits=9, decode_mode="quantlut_sym", td_x=16, td_y=16)
CB4 = dict(K=4, L=16, V=2, tlut_bits=9, decode_mode="quantlut_sym", td_x=16, td_y=16)
_QTIP_READY = False
def _legacy_decode():
"""The vendored/upstream qtip (GPLv3) path is OPT-IN, for verification only."""
return os.environ.get("QTIP_LEGACY_DECODE") == "1"
def _setup_qtip(qtip_dir=None):
"""Idempotent: stub the optional CUDA kernels and put the qtip lib on sys.path."""
global _QTIP_READY
if _QTIP_READY:
return
import types
if "fast_hadamard_transform" not in sys.modules:
_f = types.ModuleType("fast_hadamard_transform")
_f.hadamard_transform = lambda x, s=1.0: x * s
sys.modules["fast_hadamard_transform"] = _f
if "qtip_kernels" not in sys.modules:
sys.modules["qtip_kernels"] = types.ModuleType("qtip_kernels")
# qtip's lib/utils/__init__ transitively imports `datasets` (calibration-only) — stub it.
if "datasets" not in sys.modules:
_d = types.ModuleType("datasets")
_d.load_dataset = None
sys.modules["datasets"] = _d
if "glog" not in sys.modules: # logging-only, avoid the dependency
_g = types.ModuleType("glog")
_g.debug = _g.info = _g.warn = _g.warning = _g.error = lambda *a, **k: None
sys.modules["glog"] = _g
cands = []
if qtip_dir:
cands.append(qtip_dir)
if os.environ.get("QTIP_DIR"):
cands.append(os.environ["QTIP_DIR"])
_here = os.path.dirname(os.path.abspath(__file__))
cands += ["/qtip", os.path.join(_here, "vendor", "qtip"), # HF repo layout
os.path.join(_here, "..", "qtip_gemma4", "vendor", "qtip")]
for c in cands:
if c and os.path.isdir(os.path.join(c, "lib")) and c not in sys.path:
sys.path.insert(0, c)
import torch # noqa: F401
torch.compile = lambda fn=None, **kw: (fn if fn is not None else (lambda g: g))
_QTIP_READY = True
def _supported_dim(d):
from lib.utils.matmul_had import get_hadK
try:
get_hadK(d)
return True
except Exception:
return False
def padto(d):
"""Dim -> RHT-supported padded dim. Every dim in this pack (512/1024/2048) is a
power of two, where padto is the identity; non-2^k dims would need Kronecker
Hadamard factors, which only the legacy qtip path provides."""
if d > 0 and (d & (d - 1)) == 0:
return d
if _legacy_decode():
_setup_qtip()
return d if _supported_dim(d) else (1 << math.ceil(math.log2(d)))
return 1 << math.ceil(math.log2(d))
def build_codebook(cb_params, tlut, device="cuda"):
"""Rebuild the bitshift codebook with the PERSISTED tlut (kmeans init is
non-deterministic, so the encode-time LUT is saved and reused verbatim)."""
_setup_qtip()
import torch
from lib.codebook import bitshift
tl = torch.as_tensor(np.asarray(tlut, np.float32)).to(device).float()
cb = bitshift.bitshift_codebook(
L=cb_params["L"], K=cb_params["K"], V=cb_params["V"],
tlut_bits=cb_params["tlut_bits"], decode_mode=cb_params["decode_mode"],
tlut=tl).to(device).float()
return cb
def _get_hatWr(cb, trellis, m, n, td_x, td_y, device):
"""Packed int16 -> rotated codebook-unit weights [m,n] (BitshiftLinear.get_hatW)."""
import torch
tr = torch.as_tensor(np.asarray(trellis)).to(device)
if tr.dtype != torch.int16:
tr = tr.view(torch.int16)
unpacked = cb.unpack_trellis(tr, td_x * td_y) # [ntiles, T//V]
return cb.recons(unpacked).transpose(0, 1).transpose(1, 2).reshape(
m // td_x, n // td_y, td_x, td_y).transpose(1, 2).reshape(m, n)
def _decode_trellis_qtip(trellis, su, sv, tlut, m0, n0, cb_params, wscale=None,
cb=None, device="cuda"):
"""LEGACY verification oracle (QTIP_LEGACY_DECODE=1): decode through the upstream
Cornell QTIP lib (GPLv3). Kept ONLY to cross-check the numpy path; gated bit-exact
against it on every expert of every layer."""
_setup_qtip()
import torch
from lib import utils
td_x, td_y = cb_params["td_x"], cb_params["td_y"]
m, n = padto(m0), padto(n0)
if cb is None:
cb = build_codebook(cb_params, tlut, device=device)
hatWr = _get_hatWr(cb, trellis, m, n, td_x, td_y, device)
hatWr = hatWr.half().float() # fp16 hatWr spec (see decode_trellis)
if wscale is not None:
hatWr = hatWr * float(wscale)
su_t = torch.as_tensor(np.asarray(su, np.float32)).to(device).float()
sv_t = torch.as_tensor(np.asarray(sv, np.float32)).to(device).float()
hatW = (utils.matmul_hadU((utils.matmul_hadU(hatWr) * su_t).T) * sv_t).T
return hatW[:m0, :n0].float().cpu().numpy()
# ============================================================================ #
# Expert tier, PURE NUMPY (default path). Clean-room implementation from the
# mathematical spec of the bitshift-trellis format; no qtip code, no torch.
#
# Format spec (as verified bit-exact against the reference on all 40 layers):
# * A weight matrix is padded to (m, n) = (padto(m0), padto(n0)) and cut into
# td_x x td_y tiles, row-major over the (m/td_x, n/td_y) grid.
# * Each tile is a length-T scalar sequence (T = td_x*td_y, row-major inside the
# tile) produced by an L-bit shift register: the register emits one V-vector
# per step (T/V steps), then shifts in K*V fresh bits. State recurrence:
# reg_i = ((reg_{i-1} << K*V) | fresh_i) & (2^L - 1)
# * Bitstream per tile: the L bits of reg_0 (MSB first), then the K*V fresh bits
# of each later step (MSB first) -- T*K bits total, packed into BIG-ENDIAN
# 16-bit words. The last L-K*V register bits are not stored: the sequence is
# tail-biting, they wrap around to the start of the stream.
# * A register state s maps to a V-vector through a hashed symmetric LUT
# ("quantlut_sym"): with p = s*(s+1) exact-integer,
# row(s) = (p >> (16 - tlut_bits - 1)) & (2^tlut_bits - 1)
# vec(s) = tlut[row(s)] with component 0 negated iff bit 15 of p is set
# * Un-rotation (two-sided RHT): with H_d the orthonormal Walsh-Hadamard matrix
# (Sylvester order, scaled 1/sqrt(d); symmetric),
# W = diag(sv) . H_m . Wunit . H_n . diag(su), restricted to [:m0, :n0]
# computed as: FWHT each row over n then scale columns by su, then FWHT each
# column over m and scale rows by sv. For K1 packs Wunit is first multiplied
# by the scalar `wscale`.
#
# Bit-exactness note: fp32 elementwise add/sub/mul/div are IEEE-exact, so matching
# the reference's OPERATION ORDER (butterfly pairing stride 1,2,4,...; one final
# division by fp32 sqrt(dim); scale-then-rotate ordering above) makes the numpy
# output bit-identical to the torch/qtip decode, not merely close.
# ============================================================================ #
_FULL_LUT_CACHE = {}
def _np_full_lut(tlut, L, tlut_bits):
"""Expand the persisted [2^tlut_bits, V] codebook to the full [2^L, V] fp32
decoder table via the hashed-symmetric-LUT spec above."""
small = np.asarray(tlut, np.float32)
s = np.arange(1 << L, dtype=np.int64)
p = s * (s + 1) # exact in int64
row = (p >> (16 - tlut_bits - 1)) & ((1 << tlut_bits) - 1)
table = small[row].copy() # [2^L, V]
table[:, 0] *= (1 - ((p >> 15) & 1) * 2).astype(np.float32)
return table
def _np_full_lut_cached(tlut, L, tlut_bits):
key = (L, tlut_bits, np.asarray(tlut).tobytes())
if key not in _FULL_LUT_CACHE:
_FULL_LUT_CACHE[key] = _np_full_lut(tlut, L, tlut_bits)
return _FULL_LUT_CACHE[key]
def _np_rate_bits(T, K, V):
"""(bits per shift step, bits per tile). K need not be an integer: the format's only
rate constraint is that K*V and K*T be whole numbers of bits, so (K=1/2, V=2) and
(K=1/2, V=4) are legal 0.5-bpw rungs in this same bitstream."""
step, nbits = float(K) * V, float(K) * T
if step != int(step) or nbits != int(nbits):
raise ValueError(f"rate K={K} with V={V}, T={T} needs whole-bit steps "
f"(K*V={step}) and a whole-bit tile (K*T={nbits})")
return int(step), int(nbits)
def _np_unpack_trellis(stream, T, L, K, V):
"""Packed bitstream [rows, T*K/16] (u)int16 -> register states [rows, T//V] int32,
per the shift-register + tail-biting spec above."""
words = np.ascontiguousarray(stream)
if words.dtype != np.uint16:
words = words.view(np.uint16)
step, nbits = _np_rate_bits(T, K, V)
if step > L:
raise ValueError(f"K*V={step} exceeds register width L={L}")
rows, nstep = words.shape[0], T // V
bits = ((words[:, :, None].astype(np.int64) >> np.arange(15, -1, -1)) & 1)
bits = bits.reshape(rows, -1)[:, :nbits] # MSB-first big-endian words
bits = np.concatenate([bits, bits[:, :L - step]], axis=1) # tail-biting wrap
seed_w = 1 << np.arange(L - 1, -1, -1, dtype=np.int64)
step_w = 1 << np.arange(step - 1, -1, -1, dtype=np.int64)
fresh = bits[:, L:L + (nstep - 1) * step].reshape(rows, nstep - 1, step) @ step_w
states = np.empty((rows, nstep), np.int32)
reg = bits[:, :L] @ seed_w
states[:, 0] = reg
mask = (1 << L) - 1
for i in range(1, nstep):
reg = ((reg << step) & mask) | fresh[:, i - 1]
states[:, i] = reg
return states
def _np_recons(states, table, m, n, td_x, td_y):
"""Register states [ntiles, T//V] + full LUT -> rotated codebook-unit weights
[m, n] fp32, per the tile-layout spec above (V scalars per state, row-major
tiles, row-major tile grid)."""
vals = table[states] # [ntiles, T//V, V]
return np.ascontiguousarray(
vals.reshape(m // td_x, n // td_y, td_x, td_y).transpose(0, 2, 1, 3)
).reshape(m, n)
def _np_hadamard(x):
"""Orthonormal Walsh-Hadamard transform (Sylvester order) along the LAST axis:
y = FWHT(x) / sqrt(dim). Butterfly combines adjacent pairs at stride 1, then
2, 4, ... in fp32, with a single fp32 division by sqrt(dim) after the final
pass -- the exact operation order of the reference decode (bit-exact)."""
dim = x.shape[-1]
if dim & (dim - 1):
raise ValueError(f"pure-numpy RHT needs a power-of-2 dim, got {dim} "
"(set QTIP_LEGACY_DECODE=1 for Kronecker dims)")
cur = np.ascontiguousarray(x, dtype=np.float32).reshape(-1, dim)
span = 1
while span < dim:
blk = cur.reshape(-1, dim // (2 * span), 2, span)
nxt = np.empty_like(blk)
nxt[:, :, 0, :] = blk[:, :, 0, :] + blk[:, :, 1, :]
nxt[:, :, 1, :] = blk[:, :, 0, :] - blk[:, :, 1, :]
cur = nxt.reshape(-1, dim)
span *= 2
return (cur / np.float32(np.sqrt(np.float32(dim)))).reshape(x.shape)
def decode_trellis(trellis, su, sv, tlut, m0, n0, cb_params, wscale=None, cb=None,
device="cuda"):
"""Canonical trellis decode -> fp32 [m0, n0]. DEFAULT: pure numpy (no qtip, no
torch); QTIP_LEGACY_DECODE=1 routes through the upstream qtip lib instead
(verification oracle -- gated bit-exact against this path).
su [n] / sv [m] are the RHT vectors over the PADDED dims: int8 +/-1 signs for raw
PTQ tensors (then `wscale` must be given), or continuous fp16 FT'd vectors for the
block-FT'd K2 experts (Wscale absorbed into sv; pass wscale=None).
Reverses: unpack states -> tiles -> [m,n] -> *wscale -> *su-side FWHT -> *sv-side
FWHT -> unpad. `cb` may carry a prebuilt table (numpy [2^L, V] LUT, or a qtip
codebook object in legacy mode); anything else is ignored and rebuilt from tlut.
"""
if _legacy_decode():
qcb = None if isinstance(cb, np.ndarray) else cb
return _decode_trellis_qtip(trellis, su, sv, tlut, m0, n0, cb_params,
wscale=wscale, cb=qcb, device=device)
mode = cb_params.get("decode_mode", "quantlut_sym")
if mode != "quantlut_sym":
raise NotImplementedError(f"numpy decode implements quantlut_sym, got {mode}")
td_x, td_y = cb_params["td_x"], cb_params["td_y"]
L, K, V = cb_params["L"], cb_params["K"], cb_params["V"]
m, n = padto(m0), padto(n0)
table = cb if isinstance(cb, np.ndarray) else \
_np_full_lut_cached(tlut, L, cb_params["tlut_bits"])
states = _np_unpack_trellis(np.asarray(trellis), td_x * td_y, L, K, V)
unit = _np_recons(states, table, m, n, td_x, td_y)
# SPEC: hatWr is defined at fp16 precision — the qlstate's stored dtype, which
# every training/fold/eval artifact decodes from. Round the fp32 recons to fp16
# here so the shipped decode is bit-identical to those artifacts.
unit = unit.astype(np.float16).astype(np.float32)
if wscale is not None:
unit = unit * np.float32(wscale)
rowside = _np_hadamard(unit) * np.asarray(su, np.float32) # over n, then *su
colside = _np_hadamard(rowside.T) * np.asarray(sv, np.float32) # over m, then *sv
return np.ascontiguousarray(colside.T[:m0, :n0])
# ============================================================================ #
# packed-dir glue: whole-layer / whole-shard reconstruction.
# ============================================================================ #
def _read_safetensors_np(path):
"""Read a pack file. v3 files carry a zstd sidecar (`__zsc__` + metadata["zsc"])
holding every non-code-stream tensor byte-exactly; expand it transparently so
all decode paths work identically on v2 (raw) and v3 (compressed) packs."""
from safetensors import safe_open
out = {}
with safe_open(path, framework="numpy") as fh:
meta = fh.metadata() or {}
for k in fh.keys():
out[k] = fh.get_tensor(k)
if "__zsc__" in out:
import zstandard
man = json.loads(meta["zsc"])
buf = zstandard.ZstdDecompressor().decompress(
out.pop("__zsc__").tobytes(), max_output_size=man["raw_len"])
for key, dt, shape, off, nb in man["entries"]:
out[key] = np.frombuffer(buf, dtype=np.dtype(dt),
count=nb // np.dtype(dt).itemsize,
offset=off).reshape(shape)
return out, meta
_ST_TORCH_DTYPES = {"BF16": "bfloat16", "F16": "float16", "F32": "float32",
"F64": "float64", "I8": "int8", "U8": "uint8", "I16": "int16",
"I32": "int32", "I64": "int64", "BOOL": "bool"}
def read_safetensors_torch(path):
"""Torch-side v2/v3 reader for dtype-opaque files (bf16 vision/extras). v3 files
hold ONLY `__zsc__`; entries carry safetensors dtype tokens ("BF16", ...), expanded
with torch.frombuffer. Returns ({name: torch.Tensor}, metadata)."""
import torch
from safetensors import safe_open
out = {}
with safe_open(path, framework="pt") as fh:
meta = fh.metadata() or {}
for k in fh.keys():
out[k] = fh.get_tensor(k)
if "__zsc__" in out:
import zstandard
man = json.loads(meta["zsc"])
buf = zstandard.ZstdDecompressor().decompress(
out.pop("__zsc__").numpy().tobytes(), max_output_size=man["raw_len"])
for key, dt, shape, off, nb in man["entries"]:
if dt in _ST_TORCH_DTYPES: # safetensors dtype token
t = torch.frombuffer(bytearray(buf[off:off + nb]),
dtype=getattr(torch, _ST_TORCH_DTYPES[dt]))
out[key] = t.reshape(shape)
else: # numpy dtype str (np repack)
a = np.frombuffer(buf, dtype=np.dtype(dt),
count=nb // np.dtype(dt).itemsize, offset=off)
out[key] = torch.from_numpy(a.reshape(shape).copy())
return out, meta
def decode_expert_layer(packed_dir, layer, device="cuda"):
"""Reassemble fused gate_up_proj [256,1024,2048] bf16-ready fp32 + down_proj
[256,2048,512]. K2-ft experts decode with continuous su/sv (no wscale);
demoted K1 experts decode with int8 signs + Wscale, plus the shared low-rank
basis residual (B*c)@A when the manifest carries "basis"."""
path = os.path.join(packed_dir, "experts", f"L{layer:02d}.safetensors")
t, meta = _read_safetensors_np(path)
# CONTAINER DISPATCH: trained chunked shards (metadata fields includes
# wave_gamma) route to decode_expert_layer_v3t, so calling THIS entry on
# any shipped payload is safe rather than a KeyError on "manifest".
if "wave_gamma" in (meta.get("fields") or ""):
return decode_expert_layer_v3t(packed_dir, layer)
man = json.loads(meta["manifest"])
tlut, _ = _read_safetensors_np(os.path.join(packed_dir, "experts", "codebook.safetensors"))
tlut = tlut["tlut"]
if _legacy_decode(): # qtip oracle: prebuild its codebook objects
cb2 = build_codebook(man.get("cb2", CB2), tlut, device=device)
cb1 = build_codebook(man.get("cb1", CB1), tlut, device=device)
else: # numpy: K only changes transitions, LUT shared
cb2 = cb1 = _np_full_lut_cached(tlut, man.get("cb2", CB2)["L"],
man.get("cb2", CB2)["tlut_bits"])
demoted = set(man["demoted"])
geom = {p: tuple(v) for p, v in man["geom"].items()} # proj -> (m, n)
basis = man.get("basis")
if basis is not None: # rung-3 shared cold-expert low-rank basis
bA = {p: t[f"basis.{p}.A"].astype(np.float32) for p in geom}
bB = {p: t[f"basis.{p}.B"].astype(np.float32) for p in geom}
gate_up = np.empty((NEXP, 2 * INTER, HIDDEN), np.float32)
down = np.empty((NEXP, HIDDEN, INTER), np.float32)
for e in range(NEXP):
for proj, dst in (("gate", gate_up[e, :INTER]), ("up", gate_up[e, INTER:]),
("down", down[e])):
m0, n0 = geom[proj]
if e in demoted:
ws = float(np.asarray(t[f"e{e}.{proj}.Wscale"]).ravel()[0])
w = decode_trellis(t[f"e{e}.{proj}.trellis"], t[f"e{e}.{proj}.SU"],
t[f"e{e}.{proj}.SV"], tlut, m0, n0, man.get("cb1", CB1),
wscale=ws, cb=cb1, device=device)
if basis is not None:
# rs[m,n] = sum_r B[m,r]*c[r]*A[r,n], fp32, added BEFORE any bf16
# cast — op order identical to gkd_fold + pack_assemble's honesty
c = t[f"e{e}.{proj}.c"].astype(np.float32)
w = w + (bB[proj] * c[None, :]) @ bA[proj]
else:
w = decode_trellis(t[f"e{e}.{proj}.trellis"], t[f"e{e}.{proj}.su"],
t[f"e{e}.{proj}.sv"], tlut, m0, n0, man.get("cb2", CB2),
wscale=None, cb=cb2, device=device)
dst[:] = w
return {"gate_up_proj": gate_up, "down_proj": down}
# ============================================================================ #
# NE transform-free tier: Lloyd bitshift trellis L=12, uniform K, group-128 fp16
# scales, per-matrix 4096-entry fp16 LUT. Pure numpy — no qtip/GPL dependency.
# ============================================================================ #
NELL_L = 12
def _nell_bits_to_states(packed, T, k, L=NELL_L):
"""Inverse of the encoder's states_to_bits: step t's k NEW trellis bits are stored
MSB-first, bytes are big-endian packbits. State recurrence (s_{-1} = 0):
s_t = ((s_{t-1} << k) | b_t) & (2^L - 1)."""
B = packed.shape[0]
bits = np.unpackbits(packed, axis=1, count=T * k).astype(np.int64)
mask = (1 << L) - 1
states = np.empty((B, T), np.int32)
s = np.zeros(B, np.int64)
pos = 0
for t in range(T):
b = np.zeros(B, np.int64)
for j in range(k):
b = (b << 1) | bits[:, pos]
pos += 1
s = ((s << k) | b) & mask
states[:, t] = s
return states
def decode_ne_ll_tensor(t, man, name):
"""One transform-free NE matrix -> fp32 [m0, n0]. Bit-identical to the encoder's
recon: fp16 LUT gathered as fp32, * per-group fp16 scale in fp32, transpose back."""
g, k = man["group"], man["pattern"][0]
geom = man["tensors"][name]
m0, n0 = geom["shape"]
B, T = (n0, m0) if geom["transposed"] else (m0, n0)
lut = t[f"{name}|lut"].astype(np.float32)
gs = t[f"{name}|gscale"].astype(np.float32)
states = _nell_bits_to_states(t[f"{name}|packed"], T, k)
W = lut[states] * np.repeat(gs, g, axis=1)
return np.ascontiguousarray(W.T) if geom["transposed"] else W
def decode_ne_shard_canon(t, meta, packed_dir, device="cuda", subdir="ne"):
"""Canon int-lattice NE spine shard (codec canon_rht_bitshift_trellis_intlattice,
the live l64-tier zero-padded L00-L39 files). These shards carry no manifest key:
per-file metadata holds cb_params (K=4, L=16, V=2, tlut_bits=9, quantlut_sym) and
dims (name -> [m0, n0, m, n]); keys are <tensor>|{trellis,SU,SV,Wscale} with int8
sign SU/SV + scalar Wscale; the tier-shared codebook is <subdir>/tlut.safetensors."""
cbp = json.loads(meta["cb_params"])
dims = json.loads(meta["dims"])
tlut, _ = _read_safetensors_np(os.path.join(packed_dir, subdir, "tlut.safetensors"))
tlut = tlut["tlut"]
cb = build_codebook(cbp, tlut, device=device) if _legacy_decode() else \
_np_full_lut_cached(tlut, cbp["L"], cbp["tlut_bits"])
out = {}
for name in sorted({key.rsplit("|", 1)[0] for key in t}):
m0, n0 = dims[name][0], dims[name][1]
out[name] = decode_trellis(t[f"{name}|trellis"], t[f"{name}|SU"], t[f"{name}|SV"],
tlut, m0, n0, cbp,
wscale=float(np.asarray(t[f"{name}|Wscale"]).ravel()[0]),
cb=cb, device=device)
return out
def decode_ne_shard(packed_dir, shard, device="cuda", subdir="ne"):
"""Decode every NE tensor in one shard -> {name: fp32 [m0,n0]}. Dispatches on the
shard's own metadata: canon int-lattice spine (no manifest key; codec/cb_params/dims
metadata), transform-free Lloyd tier, or legacy canonical QTIP (manifest codecs).
subdir picks the size variant: "ne" (K=5, default) or "ne-4bit" (K=4, 8.8GB build)."""
path = os.path.join(packed_dir, subdir, f"{shard}.safetensors")
t, meta = _read_safetensors_np(path)
if "manifest" not in meta:
codec = meta.get("codec")
assert codec == "canon_rht_bitshift_trellis_intlattice", \
f"NE shard {shard}: no manifest and unknown codec {codec!r}"
return decode_ne_shard_canon(t, meta, packed_dir, device=device, subdir=subdir)
man = json.loads(meta["manifest"])
if man.get("codec") == "lloyd_bitshift_trellis":
return {name: decode_ne_ll_tensor(t, man, name) for name in man["tensors"]}
cb_params = man["cb"]
tlut, _ = _read_safetensors_np(os.path.join(packed_dir, "ne", "codebook.safetensors"))
tlut = tlut["tlut"]
cb = build_codebook(cb_params, tlut, device=device) if _legacy_decode() else \
_np_full_lut_cached(tlut, cb_params["L"], cb_params["tlut_bits"])
out = {}
for name, geom in man["tensors"].items():
out[name] = decode_trellis(t[f"{name}|trellis"], t[f"{name}|SU"], t[f"{name}|SV"],
tlut, geom["m0"], geom["n0"], cb_params,
wscale=float(np.asarray(t[f"{name}|Wscale"]).ravel()[0]),
cb=cb, device=device)
return out
# ============================================================================ #
# EMBED: int3 asymmetric group-64 (Hessian-free, fully deterministic). Pure numpy.
# ============================================================================ #
def pack_embed_q(q, bits=3):
"""q uint8 [rows, hid] with values < 2**bits -> packed uint8 [rows, hid*bits/8]."""
rows, hid = q.shape
b = np.unpackbits(q[..., None], axis=-1, count=8)[..., 8 - bits:] # [rows,hid,bits] MSB-first
return np.packbits(b.reshape(rows, hid * bits), axis=1)
def unpack_embed_q(packed, hid, bits=3):
rows = packed.shape[0]
b = np.unpackbits(packed, axis=1, count=hid * bits).reshape(rows, hid, bits)
q = np.zeros((rows, hid), np.uint8)
for j in range(bits):
q = (q << 1) | b[..., j]
return q
def decode_embed(packed_dir, bits=3, group=None):
"""embed_int{bits} codes -> fp32 [rows, hid]. Bit-exact vs the encoder's decode:
mn/mx stored fp16, step computed in fp32 EXACTLY as at encode time. `group` is
inferred from the stored shapes when not given (hid from q_packed, g = hid/ngroups),
so g64 and g128 packs decode identically.
Optional exception tensors close encoder-device rounding: the int4 grid was
materialized on GPU (fused multiply-add in mn + q*step), and on ~1e-3 of elements
the separately-rounded CPU product crosses a bf16 boundary. Those elements ship as
exc_idx int32 (flat index) + exc_bits uint16 (the exact bf16 bit pattern, expanded
to fp32 here) and overwrite the grid decode."""
path = os.path.join(packed_dir, "ne", f"embed_int{bits}.safetensors")
t, meta = _read_safetensors_np(path)
mn = t["mn"].astype(np.float32)[..., None] # [rows, hid/g, 1]
mx = t["mx"].astype(np.float32)[..., None]
rows, ng = mn.shape[0], mn.shape[1]
if group is None:
group = (t["q_packed"].shape[1] * 8 // bits) // ng
hid = ng * group
q = unpack_embed_q(t["q_packed"], hid, bits=bits).astype(np.float32)
lv = float(2 ** bits - 1)
step = np.maximum(mx - mn, 1e-8) / lv
dec = (mn + q.reshape(rows, ng, group) * step).reshape(rows, hid)
if "exc_idx" in t:
vals = (np.asarray(t["exc_bits"]).astype(np.uint32) << 16).view(np.float32)
dec.reshape(-1)[np.asarray(t["exc_idx"], dtype=np.int64)] = vals
return dec
def decode_embed_packed(packed_dir, subdir="ne"):
"""embed_packed.safetensors (LIVE embed tier: per-group Lloyd LUT + 4-bit nibble
codes, chunk keys EMBEDCHUNK:{r0}:{r1}.{codes|lut}) -> fp32 [rows, hid].
LOSSLESS: the lut stores the source bf16 bit patterns, so the fp32 output equals
the served bf16 embedding exactly. Low nibble = even column. Uses the torch-side
reader because the lut is bf16 (numpy cannot represent it)."""
path = os.path.join(packed_dir, subdir, "embed_packed.safetensors")
t, meta = read_safetensors_torch(path)
group = int(meta.get("group", 128))
chunks = sorted({k.rsplit(".", 1)[0] for k in t},
key=lambda c: int(c.split(":")[1]))
r_end, parts = 0, []
for c in chunks:
r0, r1 = int(c.split(":")[1]), int(c.split(":")[2])
assert r0 == r_end, f"non-contiguous embed chunks at {c}"
r_end = r1
codes = np.asarray(t[f"{c}.codes"]) # uint8 [rows, cols/2]
lut = t[f"{c}.lut"].float().numpy() # [G, 16] exact bf16 values
rows, half = codes.shape
cols = half * 2
q = np.empty((rows, cols), np.uint8)
q[:, 0::2] = codes & 0x0F
q[:, 1::2] = codes >> 4
w = np.take_along_axis(lut, q.reshape(-1, group).astype(np.int64), axis=1)
parts.append(w.reshape(rows, cols))
return np.concatenate(parts, axis=0)
if __name__ == "__main__":
import argparse
ap = argparse.ArgumentParser(description="spot-verify a packed dir against decoded refs")
ap.add_argument("--packed-dir", required=True)
ap.add_argument("--ref-experts", default=None)
ap.add_argument("--ref-ne", default=None)
ap.add_argument("--layers", default="0")
ap.add_argument("--ne-shards", default="L0")
ap.add_argument("--device", default="cuda")
a = ap.parse_args()
def _re(x, y):
x, y = np.asarray(x, np.float32).ravel(), np.asarray(y, np.float32).ravel()
return float(np.linalg.norm(x - y) / max(np.linalg.norm(y), 1e-30))
if a.ref_experts:
for L in [int(x) for x in a.layers.split(",") if x.strip()]:
dec = decode_expert_layer(a.packed_dir, L, device=a.device)
ref, _ = _read_safetensors_np(os.path.join(a.ref_experts, f"L{L:02d}.safetensors"))
for key in ("gate_up_proj", "down_proj"):
print(f"[experts L{L} {key}] relerr_vs_ref={_re(dec[key], ref[key]):.2e}")
if a.ref_ne:
for shard in [x for x in a.ne_shards.split(",") if x.strip()]:
dec = decode_ne_shard(a.packed_dir, shard, device=a.device)
refp = os.path.join(a.ref_ne, f"{shard}.safetensors")
ref, _ = _read_safetensors_np(refp) if os.path.exists(refp) else ({}, {})
for name, w in dec.items():
if name in ref:
print(f"[NE {shard} {name}] relerr_vs_ref={_re(w, ref[name]):.4f}")
else:
print(f"[NE {shard} {name}] shape={w.shape} (no ref)")
# ============================================================================ #
# int5-g64 head tier (shift-add class): symmetric int5 codes, one fp16 scale
# per 64 reduction-dim weights, RAW domain (no rotation). Storage: 8 codes
# packed into 5 little-endian bytes (code i occupies bits [5i, 5i+5) of the
# 40-bit block; stored value = q + 16, q in [-16, 15]).
# Optional protected rows: |prot_rows int32 + |prot_dense bf16 overwrite the
# listed rows with exact dense values (termination-critical vocab rows).
# ============================================================================ #
def pack_int5(q):
"""int8 [m, n] in [-16, 15] -> uint8 [m, n//8*5] little-endian 5-bit pack."""
m, n = q.shape
assert n % 8 == 0, n
u = (q.astype(np.int64) + 16).astype(np.uint64)
assert u.max() < 32 and u.min() >= 0, (int(u.min()), int(u.max()))
blocks = u.reshape(m, n // 8, 8)
word = np.zeros((m, n // 8), dtype=np.uint64)
for i in range(8):
word |= blocks[:, :, i] << np.uint64(5 * i)
by = word.astype("<u8").view(np.uint8).reshape(m, n // 8, 8)[:, :, :5]
return np.ascontiguousarray(by.reshape(m, n // 8 * 5))
def unpack_int5(qp, n):
"""Inverse of pack_int5 -> int8 [m, n] in [-16, 15]."""
m = qp.shape[0]
assert qp.shape[1] == n // 8 * 5, (qp.shape, n)
by = qp.reshape(m, n // 8, 5)
full = np.zeros((m, n // 8, 8), dtype=np.uint8)
full[:, :, :5] = by
word = full.reshape(m, n // 8 * 8).view("<u8").reshape(m, n // 8)
out = np.zeros((m, n // 8, 8), dtype=np.int8)
for i in range(8):
out[:, :, i] = ((word >> np.uint64(5 * i)) & np.uint64(31)).astype(np.int8) - 16
return out.reshape(m, n)
def decode_int5g64(qp, gscale, m0, n0, group=64, prot_rows=None, prot_dense=None):
"""Shipped int5-g64 head decode -> fp32 [m0, n0]. Pure numpy.
W[r, j] = q[r, j] * gscale[r, j // group]; protected rows are then
overwritten with their exact dense values. Accumulation against activations
is shift-adds (|q| <= 16) with one scale multiply per group of 64.
"""
q = unpack_int5(np.asarray(qp), n0).astype(np.float32)
s = np.asarray(gscale, dtype=np.float32)
W = q * np.repeat(s, group, axis=1)[:, :n0]
if prot_rows is not None and len(prot_rows):
W[np.asarray(prot_rows, dtype=np.int64)] = np.asarray(prot_dense,
dtype=np.float32)
return np.ascontiguousarray(W[:m0, :n0])
def decode_head(packed_dir, subdir="head"):
"""Shipped int5-g64 lm_head (packed/{subdir}/head_c{c}of8.safetensors, codec
"int5g64_packed") -> fp32 [vocab, hid]. Chunk keys LMHEADCHUNK:{r0}:{r1}|{qp|gscale}
(+ optional |prot_rows / |prot_dense exact-row overwrite); each file's dims metadata
gives the chunk's [rows, n]; row chunks assemble in r0 order."""
d = os.path.join(packed_dir, subdir)
files = sorted(f for f in os.listdir(d)
if f.startswith("head_c") and f.endswith(".safetensors"))
assert files, f"no head chunk files under {d}"
pieces = []
for f in files:
t, meta = _read_safetensors_np(os.path.join(d, f))
group = int(meta.get("group", 64))
dims = json.loads(meta["dims"])
for name, (m0, n0) in dims.items():
r0 = int(name.split(":")[1])
w = decode_int5g64(t[f"{name}|qp"], t[f"{name}|gscale"], m0, n0,
group=group, prot_rows=t.get(f"{name}|prot_rows"),
prot_dense=t.get(f"{name}|prot_dense"))
pieces.append((r0, w))
pieces.sort(key=lambda x: x[0])
return np.concatenate([w for _, w in pieces], axis=0)
# ============================================================================ #
# encode-v2 expert extension: per-(expert, wavefront) loading gammas.
# LDLQ feedback inflates quantization targets position-dependently along the
# anti-diagonal wavefront; v2 experts normalize each wave to the codebook's
# design radius and ship gamma [n_waves] fp16 per (expert, proj). Decode:
# unpack states -> codebook gather -> MULTIPLY each 16x16 tile by
# gamma[wave(tile)] -> * Wscale -> sign-flip Hadamard un-rotation.
# ============================================================================ #
def wave_index_map(Mb, Nb):
"""Tile grid [Mb, Nb] -> the LAST wavefront index that wrote each tile
(the encoder's starts recurrence, including the reference schedule's
duplicated top-right-starting wave)."""
starts = ([(Mb - i - 1, Nb - 1) for i in range(Mb)]
+ [(0, Nb - i - 1) for i in range(Nb)])
idx = np.zeros((Mb, Nb), dtype=np.int32)
for w, (jm, jn) in enumerate(starts):
while 0 <= jm < Mb and 0 <= jn < Nb:
idx[jm, jn] = w
jm += 1
jn -= 1
return idx
def apply_wave_gamma(Wr, gamma, td=16):
"""Multiply each td x td tile of Wr [m, n] by gamma[wave(tile)]."""
m, n = Wr.shape
Mb, Nb = m // td, n // td
g = np.asarray(gamma, np.float32)[wave_index_map(Mb, Nb)]
return np.ascontiguousarray(
(Wr.reshape(Mb, td, Nb, td) * g[:, None, :, None]).reshape(m, n))
# ============================================================================ #
# Trained-container expert decode (payload rdsl64j_s300 / rdsl64i_s100 class).
# Layout: 32-expert CHUNK-STACKED keys e{c0}.{proj}.{trellis|su|sv|wave_gamma}
# (c0 in 0,32,...,224), NO Wscale / int8 signs. su/sv are continuous fp16 over
# the PADDED dims with Wscale absorbed into sv; wave_gamma is fp16 [n_chunk,
# Mb+Nb] indexed by wave_index_map. Discriminator: file metadata carries
# fields="trellis|su|sv|wave_gamma" and susv="trained_fp16:<export>".
# EXACT op order (gated bit-faithful vs the served checkpoint, bf16 rounding):
# states -> recons -> cast fp16 -> apply_wave_gamma (BEFORE both Hadamards)
# -> hadamard over n -> * su (between the Hadamards)
# -> transpose -> hadamard over m -> * sv -> transpose -> crop
# ============================================================================ #
CB_V3T = dict(K=1.5, L=16, V=8, tlut_bits=15, decode_mode="quantlut_sym",
td_x=16, td_y=16)
def decode_expert_v3t(t, tlut, proj, e, m0, n0, cb_params=None, table=None):
"""Decode ONE expert's projection from a trained-container layer dict."""
cbp = cb_params or CB_V3T
c0, off = (e // 32) * 32, e % 32
if table is None:
table = _np_full_lut_cached(tlut, cbp["L"], cbp["tlut_bits"])
m, n = padto(m0), padto(n0)
tr = np.asarray(t[f"e{c0}.{proj}.trellis"][off])
su = np.asarray(t[f"e{c0}.{proj}.su"][off], np.float32)
sv = np.asarray(t[f"e{c0}.{proj}.sv"][off], np.float32)
states = _np_unpack_trellis(tr, cbp["td_x"] * cbp["td_y"],
cbp["L"], cbp["K"], cbp["V"])
unit = _np_recons(states, table, m, n, cbp["td_x"], cbp["td_y"])
unit = unit.astype(np.float16).astype(np.float32)
gk = f"e{c0}.{proj}.wave_gamma"
if gk in t:
unit = apply_wave_gamma(unit, np.asarray(t[gk][off], np.float32),
td=cbp["td_x"])
rowside = _np_hadamard(unit) * su
colside = _np_hadamard(rowside.T) * sv
return np.ascontiguousarray(colside.T[:m0, :n0])
def decode_expert_layer_v3t(packed_dir, layer):
"""Whole-layer reassembly for the trained container -> gate_up/down fp32."""
path = os.path.join(packed_dir, "experts", f"L{layer:02d}.safetensors")
t, meta = _read_safetensors_np(path)
assert "wave_gamma" in (meta.get("fields") or ""), "not a trained-container shard (fields metadata lacks wave_gamma)"
tlut, _ = _read_safetensors_np(os.path.join(packed_dir, "experts",
"codebook.safetensors"))
tlut = tlut["tlut"]
table = _np_full_lut_cached(tlut, CB_V3T["L"], CB_V3T["tlut_bits"])
gate_up = np.empty((NEXP, 2 * INTER, HIDDEN), np.float32)
down = np.empty((NEXP, HIDDEN, INTER), np.float32)
for e in range(NEXP):
gate_up[e, :INTER] = decode_expert_v3t(t, tlut, "gate", e, INTER, HIDDEN,
table=table)
gate_up[e, INTER:] = decode_expert_v3t(t, tlut, "up", e, INTER, HIDDEN,
table=table)
down[e] = decode_expert_v3t(t, tlut, "down", e, HIDDEN, INTER, table=table)
return {"gate_up_proj": gate_up, "down_proj": down}
|