"""Read-only, hashed mixed-bundle loader. Never expands a whole packed matrix.""" from __future__ import annotations import hashlib import json import math from pathlib import Path import numpy as np import mlx.core as mx from .gptq_q4_provenance import FORMAT as GPTQ_Q4_FORMAT, verify_provenance def sha256(path): h = hashlib.sha256() with Path(path).open('rb') as f: for block in iter(lambda: f.read(1024 * 1024), b''): h.update(block) return h.hexdigest() def require(ok, message): if not ok: raise ValueError(message) def unpack_lsb(data, bits, columns): position = np.arange(columns, dtype=np.int64) * bits padded = np.pad(data, ((0, 0), (0, 1))) byte, shift = position // 8, position % 8 words = padded[:, byte].astype(np.uint16) | (padded[:, byte + 1].astype(np.uint16) << 8) return ((words >> shift) & ((1 << bits) - 1)).astype(np.uint8) def pack_lsb(codes, bits): rows, columns = codes.shape output = np.zeros((rows, (columns * bits + 7) // 8), dtype=np.uint8) # Independent bounded-code packer, also used for legacy seven-level Q3. for col in range(columns): byte, shift = divmod(col * bits, 8) output[:, byte] |= codes[:, col] << shift if shift + bits > 8: output[:, byte + 1] |= codes[:, col] >> (8 - shift) return output class Dense: def __init__(self, values): self.values = values self.rows, self.cols = values.shape if values.ndim == 2 else (None, None) def __call__(self, inputs): require(self.values.ndim == 2 and inputs.shape[-1] == self.cols, 'dense projection shape') return inputs @ self.values.astype(inputs.dtype).T def embedding(self, ids, dtype): require(self.values.ndim == 2 and len(ids) <= 1024 and all(type(i) is int and 0 <= i < self.rows for i in ids), 'embedding IDs') return self.values[mx.array(ids, dtype=mx.int32)].astype(dtype) @property def nbytes(self): return self.values.nbytes class Packed: def __init__(self, words, scales, *, bits, columns, group_size=64): self.words, self.scales = words, scales self.affine_offset = 4 if bits == 3 else 8 self.bits, self.cols, self.group_size = bits, columns, group_size self.rows = words.shape[0] require(bits in (3, 4) and group_size == 64 and columns % 64 == 0, 'unsupported packed shape') require(words.dtype == mx.uint32 and words.shape == (self.rows, columns * bits // 32), 'word shape') require(scales.shape == (self.rows, columns // group_size), 'scale shape') def __call__(self, inputs): require(inputs.shape[-1] == self.cols, 'packed projection input width') # MLX dequantization arithmetic follows metadata dtype. Promote metadata # explicitly: F16 metadata followed by an F32 result cast is not exact. scale = self.scales.astype(inputs.dtype) bias = -self.affine_offset * scale return mx.quantized_matmul(inputs, self.words, scale, bias, transpose=True, group_size=self.group_size, bits=self.bits, mode='affine') def embedding(self, ids, dtype): require(len(ids) <= 1024 and all(type(i) is int and 0 <= i < self.rows for i in ids), 'embedding IDs') index = mx.array(ids, dtype=mx.int32) scale = self.scales[index].astype(mx.float32) return mx.dequantize(self.words[index], scale, -self.affine_offset * scale, group_size=self.group_size, bits=self.bits, mode='affine', dtype=mx.float32).astype(dtype) @property def nbytes(self): return self.words.nbytes + self.scales.nbytes def packed_record(stream, record): rows, columns = record['shape'] precision = record['precision'] require(precision in ('q3_full8', 'q3', 'q4'), 'unsupported packed precision') bits, zero, maximum = {'q3_full8': (3, 4, 7), 'q3': (3, 3, 6), 'q4': (4, 7, 14)}[precision] if precision == 'q3_full8': require(record.get('encoding') == 'uniform_lsb_full8' and record.get('grid') == dict( bits=3, zero_point=4, max_code=7, padding_code=4, signed_min=-4, signed_max=3, inference_permutation_required=False), 'full8 version/grid metadata') else: require(record.get('encoding') == 'uniform_lsb' and 'grid' not in record, 'legacy grid metadata') require(type(rows) is int and type(columns) is int and 0 < rows <= 262144 and 0 < columns <= 32768 and columns % 64 == 0, 'matrix dimensions') row_bytes = columns * bits // 8 expected = {'group_size': 64, 'padded_cols': columns, 'row_bytes': row_bytes, 'codes_bytes': rows * row_bytes, 'scales_offset': record['offset'] + rows * row_bytes, 'scales_bytes': rows * (columns // 64) * 2, 'scales_dtype': 'F16', 'scales_shape': [rows, columns // 64], 'bytes': rows * row_bytes + rows * (columns // 64) * 2} require(all(record.get(k) == v for k, v in expected.items()), 'packed storage metadata') # One packed host matrix plus bounded 64-row decode scratch. No dense weight copy. stream.seek(record['offset']) data = bytearray(stream.read(expected['codes_bytes'])) require(len(data) == expected['codes_bytes'], 'truncated packed codes') packed = np.frombuffer(data, dtype=np.uint8).reshape(rows, row_bytes) scale_bytes = stream.read(expected['scales_bytes']) require(len(scale_bytes) == expected['scales_bytes'], 'truncated packed scales') scales = np.frombuffer(scale_bytes, dtype=' 0).all(), 'invalid packed scales') for start in range(0, rows, 64): block = packed[start:start + 64] if maximum != (1 << bits) - 1 or precision == 'q3': codes = unpack_lsb(block, bits, columns) require((codes <= maximum).all(), 'reserved packed code') if precision == 'q3': block[:] = pack_lsb(codes + np.uint8(1), 3) if precision == 'q4': block += np.uint8(0x11) # no nibble carry: original codes <=14 result = Packed(mx.array(packed.view('