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