"""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))