File size: 15,874 Bytes
21fd722
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c49eca9
 
 
21fd722
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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 {})})