Reza2kn's picture
Release complete mixed Q3/Q4 Audio8 with portable CPU, MLX and exact GGUF GPU consumers
c49eca9 verified
Raw History Blame Contribute Delete
15.9 kB
"""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='<f2').reshape(rows, columns // 64)
require(np.isfinite(scales).all() and (scales > 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('<u4').reshape(rows, -1)), mx.array(scales),
bits=bits, columns=columns)
mx.eval(result.words, result.scales)
return result
class Weights:
def __init__(self, tensors, provenance=None):
self.tensors = tensors
self.provenance = provenance or {}
self.fused = {}
def tensor(self, name):
item = self.tensors[name]
require(isinstance(item, Dense), f'expected original tensor: {name}')
return item.values
def linear(self, prefix, inputs):
result = self.tensors[prefix + '.weight'](inputs)
bias = self.tensors.get(prefix + '.bias')
return result if bias is None else result + bias.values.astype(inputs.dtype)
def fuse(self, prefixes):
"""Losslessly combine row-compatible packed matrices, replacing originals
with views of the one new allocation. Bias vectors stay original.
"""
key = tuple(prefixes)
if key in self.fused: return True
parts = [self.tensors[p + '.weight'] for p in key]
if not all(isinstance(p, Packed) for p in parts): return False
first = parts[0]
if any((p.bits, p.cols, p.group_size, p.scales.dtype) !=
(first.bits, first.cols, first.group_size, first.scales.dtype) for p in parts): return False
words = mx.concatenate([p.words for p in parts], axis=0)
scales = mx.concatenate([p.scales for p in parts], axis=0)
mx.eval(words, scales)
combined = Packed(words, scales, bits=first.bits, columns=first.cols, group_size=first.group_size)
offset = 0
for part in parts:
stop = offset + part.rows
part.words, part.scales = words[offset:stop], scales[offset:stop]
offset = stop
self.fused[key] = combined
return True
def linear_many(self, prefixes, inputs):
combined = self.fused.get(tuple(prefixes))
if combined is None: return [self.linear(p, inputs) for p in prefixes]
outputs, offset = [], 0
result = combined(inputs)
for prefix in prefixes:
rows = self.tensors[prefix + '.weight'].rows
part = result[..., offset:offset + rows]
bias = self.tensors.get(prefix + '.bias')
outputs.append(part if bias is None else part + bias.values.astype(inputs.dtype))
offset += rows
return outputs
def linear_plan(self, prefixes):
"""Pure operator plus an explicit array argument list (no copies).
The operator closes over integer shape/grid descriptors only. Compiled
callers pass arrays as arguments, keeping scale casts/affine bias
transient and avoiding hidden mutable/constant weight captures.
"""
combined = self.fused.get(tuple(prefixes))
matrices = [combined] if combined is not None else [self.tensors[p + '.weight'] for p in prefixes]
arrays, specs = [], []
for matrix in matrices:
specs.append((len(arrays), matrix.bits if isinstance(matrix, Packed) else 0,
matrix.group_size if isinstance(matrix, Packed) else 0))
arrays.extend([matrix.words, matrix.scales] if isinstance(matrix, Packed) else [matrix.values])
rows, biases = [], []
for prefix in prefixes:
rows.append(self.tensors[prefix + '.weight'].rows)
biases.append(len(arrays) if prefix + '.bias' in self.tensors else None)
if prefix + '.bias' in self.tensors: arrays.append(self.tensors[prefix + '.bias'].values)
is_fused = combined is not None
def operation(x, parameters):
values = []
for index, bits, group in specs:
if bits:
scale = parameters[index + 1].astype(x.dtype)
bias = -(4 if bits == 3 else 8) * scale
values.append(mx.quantized_matmul(x, parameters[index], scale, bias,
transpose=True, group_size=group, bits=bits, mode='affine'))
else: values.append(x @ parameters[index].astype(x.dtype).T)
outputs, offset = [], 0
for part, (width, bias_index) in enumerate(zip(rows, biases)):
value = values[0][..., offset:offset + width] if is_fused else values[part]
outputs.append(value if bias_index is None else value + parameters[bias_index].astype(x.dtype))
offset += width
return outputs
return operation, arrays
@property
def nbytes(self):
return sum(v.nbytes for v in {id(x): x for x in self.tensors.values()}.values())
@classmethod
def load(cls, directory, *, progress=lambda _: None):
directory = Path(directory)
manifest_path = directory / 'manifest.json'
manifest_sha = sha256(manifest_path)
manifest = json.loads(manifest_path.read_text())
if manifest.get('format') == 'A8MOD001':
from .native_weights import load_native
return load_native(cls, directory, progress=progress)
calibrated_q4 = manifest.get('format') == GPTQ_Q4_FORMAT
if calibrated_q4:
# Stdlib metadata/source/range hashes before any device allocations.
# Existing numeric row validation and loading remain unchanged.
manifest = verify_provenance(directory)
else:
require(manifest['format'] in ('audio8-mixed-bundle-v1', 'audio8-mixed-bundle-v2'), 'bundle format')
require(manifest['group_size'] == 64 and manifest['tie_word_embeddings'], 'bundle group/tied embedding')
require(Path(manifest['weights_file']).name == manifest['weights_file'], 'unsafe payload path')
path = directory / manifest['weights_file']
require(path.stat().st_size == manifest['weights_bytes'], 'weights file length')
require(sha256(path) == manifest['weights_sha256'], 'weights SHA mismatch')
records, offset, names = manifest['tensors'], 0, set()
require(1 <= len(records) <= 2000, 'tensor count bound')
for r in records:
require(r['name'] not in names and r['offset'] == offset and type(r['bytes']) is int
and r['bytes'] > 0, 'overlapping/noncontiguous/duplicate record')
names.add(r['name']); offset += r['bytes']
require(offset <= manifest['weights_bytes'], 'record outside payload')
require(manifest['format'] == 'audio8-mixed-bundle-v2' or r['precision'] != 'q3_full8', 'full8 requires v2')
require(offset == manifest['weights_bytes'], 'unreferenced payload bytes')
tensors = {}
with path.open('rb') as stream:
for r in records:
if r['precision'] == 'original':
shape, dtype = r['shape'], r['source_dtype']
require(r['encoding'] == 'original' and dtype in ('BF16', 'F16', 'F32'), 'dense format')
require(shape and all(type(x) is int and x > 0 for x in shape), 'dense shape')
count = math.prod(shape); itemsize = 4 if dtype == 'F32' else 2
require(count * itemsize == r['bytes'] and r['bytes'] <= 64 * 1024 * 1024, 'dense tensor bound')
stream.seek(r['offset']); data = stream.read(r['bytes'])
require(len(data) == r['bytes'], 'truncated dense tensor')
if dtype == 'BF16':
values = (np.frombuffer(data, '<u2').astype(np.uint32) << 16).view(np.float32)
require(np.isfinite(values).all(), 'nonfinite BF16 source')
array = mx.array(values.reshape(shape)).astype(mx.bfloat16)
else:
values = np.frombuffer(data, '<f4' if dtype == 'F32' else '<f2').reshape(shape)
require(np.isfinite(values).all(), 'nonfinite dense source')
array = mx.array(values)
mx.eval(array); tensors[r['name']] = Dense(array)
else:
tensors[r['name']] = packed_record(stream, r)
progress({'name': r['name'], 'resident_tensor_bytes': tensors[r['name']].nbytes})
require(sha256(path) == manifest['weights_sha256'] and sha256(manifest_path) == manifest_sha,
'bundle changed while loading')
aliases = manifest['aliases']
for alias, target in aliases.items():
require(alias not in tensors and target in tensors, 'invalid tied alias')
tensors[alias] = tensors[target]
require(tensors.get('language_model.lm_head.weight') is tensors.get('language_model.model.embed_tokens.weight')
and 'language_model.model.embed_tokens.weight' in tensors, 'missing tied head')
return cls(tensors, {'manifest_sha256': manifest_sha, 'weights_sha256': manifest['weights_sha256'],
'source_revision': manifest['source_revision'], 'profile': manifest['profile'],
'source_weight_bytes': manifest['weights_bytes'],
'expected_config_sha256': manifest.get('external_assets_not_included', {}).get('config.json', {}).get('sha256'),
'layout': 'mlx_affine_power_of_two_offset', 'affine_bias_storage': 'transient_derived_from_scale',
'refit': False, 'source_payload_sha_verified': True,
**({'format': GPTQ_Q4_FORMAT, 'calibration_provenance_verified': True,
'overlay_identity_sha256': manifest['provenance']['overlay_identity_sha256'],
'copied_tensor_range_hashes_verified': True,
'static_parent_scale_equality_rechecked': False}
if calibrated_q4 else {})})