| """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) |
|
|
| |
| 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) |
|
|
| |
| from peft.tuners.tuners_utils import BaseTunerLayer |
|
|
| 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) |
|
|
| |
| |
| 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) |
|
|
| |
| 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) |
|
|