File size: 32,167 Bytes
9d3ee31
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Standalone numpy decoder for the packed Mach-1 checkpoint.

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]      (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" (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) -- 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}: older builds) --
    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 -- 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 (decode_embed(bits=4)) --
    affine int4-g64: q_packed uint8 [rows, hid/2], mn/mx fp16 [rows, hid/64]
  packed/ne/embed_packed.safetensors (8-bpw lossless Lloyd-LUT container;
    decode_embed_packed) and embed_int3.safetensors (older 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)
"""
import json
import math
import os

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)

def padto(d):
    """Dim -> padded dim: identity for powers of two, else next power of two."""
    if d > 0 and (d & (d - 1)) == 0:
        return d
    return 1 << math.ceil(math.log2(d))


# ============================================================================ #
# Expert tier. Format:
#   * 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`.
# The operation order above (and in _np_hadamard) is part of the format: fp32
# elementwise add/sub/mul/div are IEEE-exact, so following it reproduces the
# encoder's decode exactly.
# ============================================================================ #
_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 per the hashed-symmetric-LUT rule 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 only
    rate constraint is that K*V and K*T are whole numbers of bits."""
    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."""
    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 -> codebook-unit weights [m, n] fp32
    (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). Butterflies pair at stride 1, then 2, 4, ... in fp32,
    with a single fp32 division by sqrt(dim) after the final pass."""
    dim = x.shape[-1]
    if dim & (dim - 1):
        raise ValueError(f"pure-numpy RHT needs a power-of-2 dim, got {dim}")
    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"):
    """Trellis decode -> fp32 [m0, n0].

    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 vectors for the
    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);
    anything else is ignored and rebuilt from tlut.
    """
    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)
    # the unit tensor is defined at fp16 precision
    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."""
    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). v3 files hold ONLY
    `__zsc__`; entries carry safetensors dtype tokens ("BF16", ...).
    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
                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] fp32 + down_proj [256,2048,512].
    K2 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)
    # chunked shards (metadata fields includes wave_gamma) route to
    # decode_expert_layer_v3t
    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"]
    # K only changes transitions, so the [2^L, V] LUT is 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:              # 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 -- the op order is part of the format
                    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.
# ============================================================================ #
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]: 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"):
    """NE spine shard (codec canon_rht_bitshift_trellis_intlattice, zero-padded
    L00-L39 files). No manifest key: per-file metadata holds cb_params 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 = _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. subdir picks the size variant: "ne" (default) or "ne-4bit"."""
    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 = _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


# ============================================================================ #
# Embedding: int{3,4} asymmetric group codes, and the lossless Lloyd-LUT container.
# ============================================================================ #
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]. mn/mx stored fp16, step computed
    in fp32 exactly as at encode time. `group` is inferred from the stored shapes
    when not given, so g64 and g128 packs decode identically.
    Optional exception tensors: elements listed in exc_idx int32 (flat index) are
    overwritten with the exact bf16 bit patterns in exc_bits uint16 (expanded to
    fp32 here)."""
    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 (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. Low nibble = even column. Uses the torch-side
    reader because the lut is bf16."""
    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: 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.
# ============================================================================ #
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):
    """int5-g64 head decode -> fp32 [m0, n0]: W[r, j] = q[r, j] * gscale[r, j // group];
    protected rows are then overwritten with their exact dense values."""
    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"):
    """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); 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)


# ============================================================================ #
# Chunked-container expert decode with per-(expert, wavefront) gammas.
# 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".
# Op order (part of the format):
#   states -> recons -> cast fp16 -> apply_wave_gamma   (BEFORE both Hadamards)
#   -> hadamard over n -> * su    (between the Hadamards)
#   -> transpose -> hadamard over m -> * sv -> transpose -> crop
# ============================================================================ #
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 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))


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 chunked-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 chunked 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 chunked-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}