sayedM's picture
int8 CPU build + Arabic CTC draft head
71c20bd verified
Raw History Blame Contribute Delete
13.1 kB
"""Load the int8 safetensors export back into a working model. Ships inside the published repo.
from cpu_model_loader import load_cpu_model
model, processor = load_cpu_model("path-or-hub-id")
Nothing here unpickles anything. The skeleton is built from `config.json` on the `meta` device,
which allocates no storage at all, and real tensors are attached afterwards — so peak memory tracks
the finished 2.84 GB model instead of the 8.26 GB fp32 graph you would get by instantiating it
normally.
Rebuilding a quantized layer is the only fiddly part. `quantized::linear_dynamic` needs its weight
as a genuine quantized tensor, not an int8 one, so the stored `int_repr` and scale are recombined
with `_make_per_tensor_quantized_tensor` before being handed to the module.
Requires `transformers >= 5.4` (for `cohere_asr`) and a torch build whose quantized engine is
fbgemm, onednn or qnnpack — i.e. any normal CPU build.
"""
from __future__ import annotations
import os
import json
import torch
import torch.nn as nn
QLinear = torch.ao.nn.quantized.dynamic.Linear
def _resolve(path_or_id: str) -> str:
"""Local directory, or a Hub id to download."""
if os.path.isdir(path_or_id):
return path_or_id
from huggingface_hub import snapshot_download
return snapshot_download(path_or_id)
def _pick_engine():
for e in ("onednn", "fbgemm", "qnnpack", "x86"):
if e in torch.backends.quantized.supported_engines:
torch.backends.quantized.engine = e
return e
raise RuntimeError("no quantized engine available in this torch build")
def load_cpu_model(path_or_id: str, attn: str = "sdpa", verbose: bool = True):
"""Returns (model, processor), ready for `generate`."""
from safetensors.torch import load_file
from transformers import AutoConfig, AutoProcessor, CohereAsrForConditionalGeneration
d = _resolve(path_or_id)
engine = _pick_engine()
torch.set_grad_enabled(False)
meta_json = json.load(open(os.path.join(d, "quant_map.json"), encoding="utf-8"))
qmap = meta_json["quantized_modules"]
nonpersistent = meta_json.get("nonpersistent_buffers", [])
sd = load_file(os.path.join(d, "model.safetensors"))
cfg = AutoConfig.from_pretrained(d)
cfg._attn_implementation = attn
with torch.device("meta"): # no allocation
model = CohereAsrForConditionalGeneration(cfg)
# 1. Attach the plain tensors FIRST, while the Linears are still ordinary meta modules.
# Order matters: a quantized Linear's own _load_from_state_dict insists on `scale` and
# `zero_point` entries, so running load_state_dict after the swap raises KeyError on the
# first quantized layer. assign=True replaces the meta tensors rather than copying into
# them, which is what keeps this from materializing an fp32 model first.
rest = {k: v for k, v in sd.items()
if not (k.startswith("nonpersistent.") or k.endswith(".weight_int8")
or k.endswith(".weight_scale")
or (k.endswith(".bias") and k[: -len(".bias")] in qmap))}
missing, unexpected = model.load_state_dict(rest, strict=False, assign=True)
# buffers state_dict() never reported, and which a meta skeleton therefore cannot supply
by_name_all = dict(model.named_modules())
for full in nonpersistent:
owner, _, leaf = full.rpartition(".")
by_name_all[owner].register_buffer(leaf, sd["nonpersistent." + full], persistent=False)
# everything still missing must be a quantized layer's weight/bias, about to be supplied below
expected_missing = {f"{n}.weight" for n in qmap} | {
f"{n}.bias" for n, m in qmap.items() if m["has_bias"]}
stray = set(missing) - expected_missing
if stray:
raise RuntimeError(f"{len(stray)} tensors are missing from the checkpoint, "
f"first: {sorted(stray)[:3]}")
if unexpected:
raise RuntimeError(f"unexpected tensors in the checkpoint: {unexpected[:3]}")
# 2. swap every quantized Linear for a real one carrying the stored int8 weight
by_name = dict(model.named_modules())
for full, meta in qmap.items():
parent_name, _, leaf = full.rpartition(".")
parent = by_name[parent_name] if parent_name else model
q = QLinear(meta["in_features"], meta["out_features"],
bias_=meta["has_bias"], dtype=torch.qint8)
w = torch._make_per_tensor_quantized_tensor(
sd[full + ".weight_int8"], float(sd[full + ".weight_scale"]), meta["zero_point"])
q.set_weight_bias(w, sd.get(full + ".bias"))
setattr(parent, leaf, q)
still_meta = [n for n, p in list(model.named_parameters()) + list(model.named_buffers())
if p is not None and p.is_meta]
if still_meta:
raise RuntimeError(f"{len(still_meta)} tensors never got real storage, "
f"first: {still_meta[:3]}")
# Constructing the class directly derives a GenerationConfig from config.json and never reads
# generation_config.json, which is where decoder_start_token_id (13764) actually lives. Without
# this the model loads, transcribes, and quietly decodes differently from the original.
gc_path = os.path.join(d, "generation_config.json")
if os.path.exists(gc_path):
from transformers import GenerationConfig
model.generation_config = GenerationConfig.from_pretrained(d)
model.eval()
proc = AutoProcessor.from_pretrained(d)
if verbose:
gb = (sum(p.numel() * p.element_size() for p in model.parameters())
+ sum(v.numel() for k, v in sd.items() if v.dtype == torch.int8)) / 1e9
print(f" loaded int8 model on '{engine}' (~{gb:.2f} GB resident)", flush=True)
return model, proc
def load_draft_head(path_or_id: str):
"""The Arabic CTC draft head from the same repo, or (None, None) if it is not there."""
import sys
d = _resolve(path_or_id)
if not os.path.exists(os.path.join(d, "draft_head.safetensors")):
return None, None
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from draft_head import DraftHead
# config_dir=d: the head's adapter block is rebuilt from this repo's own config.json, so it
# works on a machine that has never seen the original model.
return DraftHead.from_checkpoint(d, config_dir=d)
def transcribe(model, processor, audio, language="ar", head=None, k=8, max_new_tokens=256):
"""Transcribe one clip, optionally with speculative decoding.
`audio` is a path or a 16 kHz float32 array. Pass `head` (from `load_draft_head`) to use the
draft head; without it this is plain greedy decoding.
Only single-chunk audio is drafted: `transformers`' assisted generation is batch-1 only, and
the head is trained on one language, so anything else falls through to greedy.
"""
import sys
import numpy as np
from transformers.audio_utils import load_audio
if isinstance(audio, str):
audio = load_audio(audio, sampling_rate=16000)
audio = np.asarray(audio, dtype="float32")
inp = processor(audio=audio, sampling_rate=16000, return_tensors="pt",
language=language, punctuation=True)
feats, mask = inp["input_features"].float(), inp["attention_mask"]
pid, aci = inp["decoder_input_ids"], inp.get("audio_chunk_index")
kw = dict(max_new_tokens=max_new_tokens, num_beams=1, do_sample=False, use_cache=True)
single = feats.shape[0] == 1
if head is not None and single:
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from draft_head import ctc_greedy
from spec_decode import SequenceDraftCandidateGenerator, spec_generate
enc = model.model.encoder(input_features=feats, attention_mask=mask)
flen = (enc.attention_mask.sum(-1) if enc.attention_mask is not None
else torch.tensor([enc.last_hidden_state.shape[1]]))
(seq, conf), = ctc_greedy(
head(enc.last_hidden_state.float(),
enc.attention_mask.bool() if enc.attention_mask is not None else None), flen)
if seq:
gen = SequenceDraftCandidateGenerator(
torch.tensor(seq, dtype=torch.long), torch.tensor(conf, dtype=torch.float),
k=k, eos_token_id=torch.tensor(
[model.generation_config.eos_token_id]).flatten(),
max_length=10 + max_new_tokens)
out = spec_generate(model, gen, trigger_k=k, input_features=feats,
attention_mask=mask, encoder_outputs=enc,
decoder_input_ids=pid, **kw)
return processor.tokenizer.decode(out[0], skip_special_tokens=True)
out = model.generate(input_features=feats, attention_mask=mask, encoder_outputs=enc,
decoder_input_ids=pid, **kw)
return processor.tokenizer.decode(out[0], skip_special_tokens=True)
out = model.generate(input_features=feats, attention_mask=mask,
decoder_input_ids=pid, **kw)
parts = [processor.tokenizer.decode(o, skip_special_tokens=True) for o in out]
if len(parts) == 1:
return parts[0]
return processor._reassemble_chunk_texts(parts, aci, " ")[0]
# --------------------------------------------------------------------- verify
def verify(export_dir: str, cache: str, n_layers: int = 8):
"""Does the reloaded export reproduce the source model exactly?
Compares the dequantized weights of a sample of layers, then runs both models over the same
audio and requires byte-identical token sequences. Weight equality alone is not enough: a
mis-wired module could still be attached to the wrong parent.
"""
import numpy as np
from transformers.audio_utils import load_audio
print("\nverifying the export against the source model", flush=True)
model, proc = load_cpu_model(export_dir)
src = torch.load(cache, map_location="cpu", weights_only=False)
srcmods = {n: m for n, m in src.named_modules() if isinstance(m, QLinear)}
newmods = {n: m for n, m in model.named_modules() if isinstance(m, QLinear)}
print(f" quantized modules: source {len(srcmods)}, reloaded {len(newmods)}")
assert set(srcmods) == set(newmods), "module sets differ"
names = list(srcmods)
sample = names[:: max(1, len(names) // n_layers)][:n_layers]
worst = 0.0
for n in sample:
a, b = srcmods[n].weight(), newmods[n].weight()
assert a.q_scale() == b.q_scale(), f"{n}: scale {a.q_scale()} != {b.q_scale()}"
assert torch.equal(a.int_repr(), b.int_repr()), f"{n}: int_repr differs"
worst = max(worst, float((a.dequantize() - b.dequantize()).abs().max()))
print(f" {len(sample)} sampled layers: scales and int_repr identical, "
f"max dequantized diff {worst:.1e}")
clip = os.path.join(os.path.dirname(cache), "..", "examples", "sample2.wav")
clip = os.path.normpath(clip)
y = np.asarray(load_audio(clip, sampling_rate=16000), dtype="float32")
inp = proc(audio=y, sampling_rate=16000, return_tensors="pt", language="ar", punctuation=True)
kw = dict(input_features=inp["input_features"].float(),
attention_mask=inp["attention_mask"],
decoder_input_ids=inp["decoder_input_ids"],
max_new_tokens=256, num_beams=1, do_sample=False, use_cache=True)
out_new = model.generate(**kw)
out_src = src.generate(**kw)
same = torch.equal(out_new, out_src)
print(f" generated tokens identical: {same} ({out_new.shape[1]} tokens)")
if not same:
t = proc.tokenizer
print(" source :", t.decode(out_src[0], skip_special_tokens=True)[:80])
print(" reloaded:", t.decode(out_new[0], skip_special_tokens=True)[:80])
raise SystemExit("EXPORT IS NOT FAITHFUL — do not publish it")
print(" PASS: the export is the deployed model")
return True
if __name__ == "__main__":
import argparse
ap = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("path", help="export directory or Hub id")
ap.add_argument("--audio", help="transcribe this file as a smoke test")
ap.add_argument("--language", default="ar", choices=["ar", "en"])
a = ap.parse_args()
m, p = load_cpu_model(a.path)
if a.audio:
import numpy as np
from transformers.audio_utils import load_audio
y = np.asarray(load_audio(a.audio, sampling_rate=16000), dtype="float32")
i = p(audio=y, sampling_rate=16000, return_tensors="pt",
language=a.language, punctuation=True)
o = m.generate(input_features=i["input_features"].float(),
attention_mask=i["attention_mask"],
decoder_input_ids=i["decoder_input_ids"],
max_new_tokens=256, num_beams=1, do_sample=False, use_cache=True)
print(p.tokenizer.decode(o[0], skip_special_tokens=True))