Audio8-ASR-Infinite-Compressed / mlx_runtime /native_model_format.py
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
11.6 kB
#!/usr/bin/env python3
"""Byte-preserving Audio8 bundle -> bounded-memory native model stream.
A8MOD001 deliberately covers only the pinned Audio8 architecture and group64
Q4 encoder/head + Q4 or full-eight Q3 decoder. It does not prove speech quality.
All integer fields are little endian. Footer is SHA256 of all preceding bytes.
"""
import argparse
import hashlib
import json
import math
import os
from pathlib import Path
import platform
import shutil
import struct
import sys
MAGIC = b"A8MOD001"
HEADER = struct.Struct("<8sIQ") # magic, records, payload bytes
RECORD = struct.Struct("<7IQ") # name bytes, kind, rank, group, bits, zero, max, payload
CHUNK = 64 * 1024
CONFIG_SHA256 = "744baed356a8c86a9791fc76cecb73219ed4aff795147c6cdc082337a13f019c"
ASSETS = ("config.json", "tokenizer.json", "tokenizer_config.json", "preprocessor_config.json",
"generation_config.json", "chat_template.jinja")
def sha256(path):
h = hashlib.sha256()
with Path(path).open("rb") as f:
for block in iter(lambda: f.read(CHUNK), b""):
h.update(block)
return h.hexdigest()
def schema():
"""Names map to (shape, kind). Kind 0 is BF16; 1 is packed uniform."""
s = {}
def put(name, shape, kind=0):
s[name] = (shape, kind)
for c, cols in ((1, 128), (2, 1280)):
put(f"audio_tower.embedder.conv{c}.weight", [1280, cols, 3])
put(f"audio_tower.embedder.conv{c}.bias", [1280])
suffixes = ("self_attn.q_proj", "self_attn.k_proj", "self_attn.v_proj", "self_attn.o_proj",
"mlp.gate_proj", "mlp.up_proj", "mlp.down_proj")
for tower, count, hidden, dims, biases, norms in (
("audio_tower", 32, 1280, ((2048,1280),)*3+((1280,2048),(5120,1280),(5120,1280),(1280,5120)),
(0,2,3,6), ("self_attn_layer_norm", "final_layer_norm")),
("language_model.model", 36, 2048, ((2048,2048),(256,2048),(256,2048),(2048,2048),
(11008,2048),(11008,2048),(2048,11008)),
(0,1,2), ("input_layernorm", "post_attention_layernorm"))):
for layer in range(count):
p = f"{tower}.layers.{layer}."
for i, (suffix, shape) in enumerate(zip(suffixes, dims)):
put(p+suffix+".weight", list(shape), 1)
if i in biases:
put(p+suffix+".bias", [shape[0]])
for norm in norms:
put(p+norm+".weight", [hidden])
if tower.startswith("language"):
put(p+"ada_rms_norm.linear1.weight", [32,2048])
put(p+"ada_rms_norm.linear2.weight", [2048,32])
put("audio_tower.norm.weight", [1280])
put("language_model.model.norm.weight", [2048])
put("frame_len_embedding.weight", [3,2048])
put("language_model.model.embed_tokens.weight", [151936,2048], 1)
put("multi_modal_projector.linear_1.weight", [2048,10240], 1)
put("multi_modal_projector.linear_2.weight", [2048,2048], 1)
for i in range(4):
put(f"semantic_vad_heads.{i}.weight", [8,2048])
put(f"semantic_vad_heads.{i}.bias", [8])
return dict(sorted(s.items()))
def record_metadata(record, expected):
shape, kind = expected
if record["shape"] != shape or record["source_dtype"] != "BF16":
raise ValueError(f"unexpected Audio8 shape/dtype: {record['name']}")
if kind == 0:
if record["precision"] != "original" or record["encoding"] != "original":
raise ValueError("expected original BF16 tensor")
grid = (0,0,0,0)
size = math.prod(shape)*2
else:
decoder = record["name"].startswith("language_model.model.layers.")
precision = record["precision"]
if precision == "q4":
grid = (64,4,7,14)
elif decoder and precision == "q3_full8":
grid = (64,3,4,7)
else:
raise ValueError("unsupported native model grid")
group,bits,_,_ = grid
groups = (shape[1]+group-1)//group
code_bytes = shape[0]*((groups*group*bits+7)//8)
if (record["group_size"] != group or record["codes_bytes"] != code_bytes
or record["scales_bytes"] != shape[0]*groups*2 or record["scales_dtype"] != "F16"
or record["scales_offset"] != record["offset"]+code_bytes):
raise ValueError("invalid packed tensor layout")
size = code_bytes+shape[0]*groups*2
if record["bytes"] != size:
raise ValueError("tensor payload size mismatch")
return kind,grid,size
def write_stream(source, records, output, expected_schema=None):
"""Bounded 64KiB copy; injectable tiny schema is only for low-level tests."""
expected = schema() if expected_schema is None else expected_schema
by_name = {r["name"]: r for r in records}
if len(by_name) != len(records) or set(by_name) != set(expected):
raise ValueError("missing, duplicate or unexpected model tensor")
metadata = {n: record_metadata(by_name[n], e) for n,e in expected.items()}
payload = sum(m[2] for m in metadata.values())
digest = hashlib.sha256()
tensor_hashes = {}
with Path(source).open("rb") as src, Path(output).open("xb") as dst:
def write(data):
dst.write(data)
digest.update(data)
write(HEADER.pack(MAGIC,len(expected),payload))
for name,(shape,_) in sorted(expected.items()):
record = by_name[name]
kind,grid,size = metadata[name]
encoded = name.encode("ascii")
write(RECORD.pack(len(encoded),kind,len(shape),*grid,size))
write(struct.pack("<"+"I"*len(shape),*shape))
write(encoded)
src.seek(record["offset"])
remaining = size
h = hashlib.sha256()
while remaining:
block = src.read(min(CHUNK,remaining))
if not block:
raise ValueError("source weights truncated during export")
h.update(block)
write(block)
remaining -= len(block)
tensor_hashes[name] = h.hexdigest()
dst.write(digest.digest())
dst.flush()
os.fsync(dst.fileno())
return {"payload_bytes":payload,"stream_body_sha256":digest.hexdigest(),
"bytes":Path(output).stat().st_size,"sha256":sha256(output),"tensor_sha256":tensor_hashes}
def verify_stream(path, expected_schema=None):
"""Independent streaming readback of all headers and footer, no reconstruction."""
expected = schema() if expected_schema is None else expected_schema
h = hashlib.sha256()
with Path(path).open("rb") as f:
def read(n):
data = f.read(n)
if len(data) != n:
raise ValueError("truncated native model")
h.update(data)
return data
magic,count,total = HEADER.unpack(read(HEADER.size))
if magic != MAGIC or count != len(expected):
raise ValueError("native model header mismatch")
actual = 0
for name,(shape,kind) in sorted(expected.items()):
n,k,rank,group,bits,zero,maximum,size = RECORD.unpack(read(RECORD.size))
if (n != len(name) or k != kind or rank != len(shape)
or list(struct.unpack("<"+"I"*rank,read(4*rank))) != shape
or read(n) != name.encode("ascii")):
raise ValueError("native tensor schema mismatch")
if kind == 0:
wanted = math.prod(shape)*2
valid = (group,bits,zero,maximum)==(0,0,0,0)
else:
allowed = [(64,4,7,14)]
if name.startswith("language_model.model.layers."):
allowed.append((64,3,4,7))
valid = (group,bits,zero,maximum) in allowed
wanted = shape[0]*(((shape[1]+63)//64)*64*bits//8+((shape[1]+63)//64)*2)
if not valid or size != wanted:
raise ValueError("native tensor grid/size mismatch")
remaining = size
while remaining:
block = read(min(CHUNK,remaining))
remaining -= len(block)
actual += size
if actual != total or f.read(32) != h.digest() or f.read(1):
raise ValueError("native payload/footer mismatch")
return {"records":count,"payload_bytes":total,"stream_body_sha256":h.hexdigest()}
def export_model(bundle, output, assets):
from mixed_bundle import _validate_manifest
bundle,output,assets = Path(bundle),Path(output),Path(assets)
manifest = _validate_manifest(bundle)
if (manifest["group_size"] != 64 or manifest["profile"] not in ("q4","e4_d3_full8_h4")
or manifest["aliases"] != {"language_model.lm_head.weight":"language_model.model.embed_tokens.weight"}
or manifest["tie_word_embeddings"] is not True):
raise ValueError("native export requires supported group64 profile and exactly tied head")
if sha256(assets/"config.json") != CONFIG_SHA256:
raise ValueError("native schema requires exact pinned Audio8 config")
asset_meta = manifest["external_assets_not_included"]
for name in ASSETS:
if sha256(assets/name) != asset_meta[name]["sha256"] or (assets/name).stat().st_size != asset_meta[name]["bytes"]:
raise ValueError(f"external asset changed: {name}")
output.mkdir(parents=True,exist_ok=False)
path = output/"model.a8m"
try:
result = write_stream(bundle/"weights.bin",manifest["tensors"],path)
verification = verify_stream(path)
for name in ASSETS:
shutil.copyfile(assets/name,output/name)
if sha256(output/name) != asset_meta[name]["sha256"]:
raise ValueError("copied asset mismatch")
# Full source rehash detects an accidental concurrent rewrite during copy.
if sha256(bundle/"weights.bin") != manifest["weights_sha256"]:
raise ValueError("source bundle changed during native export")
report = {"format":"A8MOD001","architecture":"pinned Audio8 32 encoder / 36 decoder",
"source_bundle":str(bundle.resolve()),"source_manifest_sha256":sha256(bundle/"manifest.json"),
"source_weights_sha256":manifest["weights_sha256"],"source_profile":manifest["profile"],
"source_payload_bytes":manifest["weights_bytes"],"stored_tensors":len(manifest["tensors"]),
"aliases":manifest["aliases"],"accuracy_validated":False,
"claim_limit":"native loading and tensor preservation only; no speech-quality claim",
"model_file":"model.a8m","model":result,"readback":verification,
"copied_assets":{n:asset_meta[n] for n in ASSETS},"copy_scratch_bytes":CHUNK,
"runtime":{"python":sys.version,"platform":platform.platform()},
"exporter_sha256":sha256(__file__)}
(output/"native-manifest.json").write_text(json.dumps(report,indent=2)+"\n")
return report
except BaseException:
(output/"INCOMPLETE").write_text("Export failed; do not consume this directory.\n")
raise
def main():
p=argparse.ArgumentParser(description=__doc__)
p.add_argument("bundle",type=Path)
p.add_argument("output",type=Path)
p.add_argument("--assets",type=Path,default=Path("models/original"))
a=p.parse_args()
report=export_model(a.bundle,a.output,a.assets)
print(json.dumps({k:report[k] for k in ("format","source_payload_bytes","stored_tensors","accuracy_validated")}))
if __name__=="__main__":
main()