granite-speech-3.3-2b-onnx / verify_native.py
Jarbas's picture
Upload verify_native.py with huggingface_hub
139d836 verified
Raw
History Blame Contribute Delete
5.54 kB
"""Native reference + LoRA-merge equivalence check for granite-speech-3.3-2b.
Runs three pipelines on the same clips:
1. stock transformers generate() with the audio LoRA adapter enabled
2. the same, but with the adapter *merged* into the base weights and the
language model rebuilt as a plain GraniteForCausalLM
3. a manual three-graph pipeline (encoder+projector -> embeds -> causal LM
with KV cache) on the merged LM: the torch twin of the ONNX runtime
"""
import json
import sys
import time
import soundfile as sf
import torch
from transformers import AutoModelForSpeechSeq2Seq, AutoProcessor, GraniteForCausalLM
MODEL = "ibm-granite/granite-speech-3.3-2b"
OUT = "/media/hdd16/onnx-asr-exports/granite-speech-3.3-2b"
CLIPS = ["en_1", "en_2", "pt_1", "pt_2"]
SYSTEM = (
"Knowledge Cutoff Date: April 2024.\nToday's Date: April 9, 2025.\n"
"You are Granite, developed by IBM. You are a helpful AI assistant"
)
USER = "<|audio|>can you transcribe the speech into a written format?"
torch.set_num_threads(8)
torch.set_grad_enabled(False)
processor = AutoProcessor.from_pretrained(MODEL)
tokenizer = processor.tokenizer
model = AutoModelForSpeechSeq2Seq.from_pretrained(MODEL, dtype=torch.float32, device_map="cpu")
model.eval()
print("peft loaded:", model._hf_peft_config_loaded, flush=True)
chat = [dict(role="system", content=SYSTEM), dict(role="user", content=USER)]
prompt = tokenizer.apply_chat_template(chat, tokenize=False, add_generation_prompt=True)
# ---- prompt ids around the audio placeholder -------------------------------
audio_id = model.config.audio_token_index
one = tokenizer(prompt.replace("<|audio|>", "<|audio|>"), return_tensors="pt")["input_ids"][0].tolist()
first = one.index(audio_id)
last = len(one) - 1 - one[::-1].index(audio_id)
prefix_ids = one[:first]
suffix_ids = one[last + 1 :]
print("prefix", len(prefix_ids), "suffix", len(suffix_ids), "audio placeholders", last - first + 1, flush=True)
def load(name):
wav, sr = sf.read(f"{OUT}/clips/{name}.wav", dtype="float32")
assert sr == 16000, sr
return torch.from_numpy(wav)[None]
def native(wav):
inputs = processor(prompt, wav, device="cpu", return_tensors="pt")
n = inputs["input_ids"].shape[-1]
out = model.generate(**inputs, max_new_tokens=256, do_sample=False, num_beams=1)
return tokenizer.decode(out[0, n:], skip_special_tokens=True).strip()
results = {"prefix_ids": prefix_ids, "suffix_ids": suffix_ids}
results["native_adapter"] = {}
for name in CLIPS:
t = time.perf_counter()
results["native_adapter"][name] = native(load(name))
print(f"[adapter] {name}: {results['native_adapter'][name]!r} ({time.perf_counter() - t:.1f}s)", flush=True)
# ---- merge the LoRA --------------------------------------------------------
from peft.tuners.tuners_utils import BaseTunerLayer # noqa: E402
n_merged = 0
for mod in model.language_model.modules():
if isinstance(mod, BaseTunerLayer):
mod.merge(safe_merge=True)
n_merged += 1
print("merged lora layers:", n_merged, flush=True)
# Unload: replace every peft LoraLayer with its (now merged) base Linear, so the
# language model is a plain GraniteForCausalLM again.
n_unloaded = 0
for name in [n for n, m in model.language_model.named_modules() if isinstance(m, BaseTunerLayer)]:
parent_name, _, attr = name.rpartition(".")
parent = model.language_model.get_submodule(parent_name)
setattr(parent, attr, getattr(parent, attr).base_layer)
n_unloaded += 1
print("unloaded lora layers:", n_unloaded, flush=True)
assert not [m for m in model.language_model.modules() if isinstance(m, BaseTunerLayer)]
model._hf_peft_config_loaded = False
merged_lm = model.language_model
assert isinstance(merged_lm, GraniteForCausalLM), type(merged_lm)
merged_lm.eval()
results["native_merged"] = {}
for name in CLIPS:
results["native_merged"][name] = native(load(name))
print(f"[merged ] {name}: {results['native_merged'][name]!r}", flush=True)
# ---- manual three-graph pipeline on the merged LM --------------------------
embed = merged_lm.get_input_embeddings()
def manual(wav):
feats = processor.audio_processor(wav, device="cpu")
audio_embeds = model.projector(model.encoder(feats["input_features"]))
audio_embeds = audio_embeds[:, : int(feats["input_features_mask"].sum())]
ids = torch.tensor([prefix_ids], dtype=torch.long)
embeds = torch.cat(
[embed(ids), audio_embeds, embed(torch.tensor([suffix_ids], dtype=torch.long))], dim=1
)
past = None
tokens = []
for _ in range(256):
out = merged_lm(inputs_embeds=embeds, past_key_values=past, use_cache=True)
past = out.past_key_values
tok = int(out.logits[0, -1].argmax())
if tok == 0:
break
tokens.append(tok)
embeds = embed(torch.tensor([[tok]], dtype=torch.long))
return tokenizer.decode(tokens, skip_special_tokens=True).strip(), audio_embeds.shape[1]
results["manual_merged"] = {}
results["audio_embed_sizes"] = {}
for name in CLIPS:
text, n = manual(load(name))
results["manual_merged"][name] = text
results["audio_embed_sizes"][name] = n
print(f"[manual ] {name}: {text!r} (audio tokens {n})", flush=True)
ok = all(
results["native_adapter"][n] == results["native_merged"][n] == results["manual_merged"][n] for n in CLIPS
)
results["equivalent"] = ok
print("EQUIVALENT:", ok, flush=True)
with open(f"{OUT}/native.json", "w") as f:
json.dump(results, f, ensure_ascii=False, indent=2)
sys.exit(0 if ok else 1)