#!/usr/bin/env python3 # Copyright (C) 2024-2026 Intel Corporation # SPDX-License-Identifier: Apache-2.0 """ Cohere Transcribe — OpenVINO inference. Loads a converted model directory (stateless or stateful / KV-cache — auto detected from the metadata) and transcribes an audio file, printing the transcript and timing. Usage: python inference.py --model_dir . --audio ../../../Auto-Opt/output_001.wav python inference.py --model_dir /path/to/ir --audio sample.wav --device CPU """ import argparse import glob import json import time from pathlib import Path import numpy as np def _np(x, dtype): return np.asarray(x.cpu() if hasattr(x, "cpu") else x).astype(dtype) def load_meta(model_dir: Path) -> dict: """Find the converter metadata json (stateless or KV-cache) in the dir.""" for name in ("ov_cohere_transcribe_kvcache.json", "ov_cohere_transcribe.json"): p = model_dir / name if p.is_file(): return json.loads(p.read_text()) candidates = glob.glob(str(model_dir / "ov_cohere_transcribe*.json")) if candidates: return json.loads(Path(candidates[0]).read_text()) raise FileNotFoundError(f"No ov_cohere_transcribe*.json metadata found in {model_dir}") def load_audio(audio_path: str, sr: int) -> np.ndarray: import librosa audio, _ = librosa.load(audio_path, sr=sr, mono=True) return audio # --------------------------------------------------------------------------- # stateless decode # --------------------------------------------------------------------------- def transcribe_stateless(core, model_dir, meta, processor, feats, device, max_new_tokens): encoder = core.compile_model(model_dir / meta["encoder_ir"], device) decoder = core.compile_model(model_dir / meta["decoder_ir"], device) eos = meta["eos_token_id"] eos_set = set(eos) if isinstance(eos, (list, tuple)) else {eos} t0 = time.perf_counter() enc = encoder({"input_features": feats["input_features"], "attention_mask": feats["attention_mask"]}) ehs = enc["encoder_hidden_states"] emask = enc["encoder_attention_mask"] t_enc = time.perf_counter() gen = feats["decoder_input_ids"].copy() prompt_len = gen.shape[1] n_tokens = 0 for _ in range(max_new_tokens): logits = decoder( {"decoder_input_ids": gen, "encoder_hidden_states": ehs, "encoder_attention_mask": emask} )["logits"] next_id = int(logits[0, -1].argmax()) gen = np.concatenate([gen, np.array([[next_id]], dtype=np.int64)], axis=1) n_tokens += 1 if next_id in eos_set: break t_dec = time.perf_counter() text = processor.batch_decode(gen[:, prompt_len:], skip_special_tokens=True)[0].strip() return text, n_tokens, t_enc - t0, t_dec - t_enc, t_dec - t0 # --------------------------------------------------------------------------- # stateful (KV-cache) decode # --------------------------------------------------------------------------- def transcribe_stateful(core, model_dir, meta, processor, feats, device, max_new_tokens): num_layers = meta["num_layers"] encoder = core.compile_model(model_dir / meta["encoder_ir"], device) prefill = core.compile_model(model_dir / meta["decoder_ir"], device) decode = core.compile_model(model_dir / meta["decoder_with_past_ir"], device) eos = meta["eos_token_id"] eos_set = set(eos) if isinstance(eos, (list, tuple)) else {eos} t0 = time.perf_counter() enc = encoder({"input_features": feats["input_features"], "attention_mask": feats["attention_mask"]}) ehs = enc["encoder_hidden_states"] emask = enc["encoder_attention_mask"] t_enc = time.perf_counter() pf = prefill( { "decoder_input_ids": feats["decoder_input_ids"], "encoder_hidden_states": ehs, "encoder_attention_mask": emask, } ) pf = {k.get_any_name(): v for k, v in pf.items()} self_kv = [(pf[f"present.{i}.self.key"], pf[f"present.{i}.self.value"]) for i in range(num_layers)] cross_kv = [(pf[f"present.{i}.cross.key"], pf[f"present.{i}.cross.value"]) for i in range(num_layers)] next_id = int(pf["logits"][0, -1].argmax()) generated = [next_id] for _ in range(max_new_tokens): if next_id in eos_set: break seq_len = self_kv[0][0].shape[2] feed = { "decoder_input_ids": np.array([[next_id]], dtype=np.int64), "encoder_hidden_states": ehs, "encoder_attention_mask": emask, "self_attention_mask": np.ones((1, seq_len + 1), dtype=np.int64), } for i in range(num_layers): feed[f"past.{i}.self.key"] = self_kv[i][0] feed[f"past.{i}.self.value"] = self_kv[i][1] feed[f"past.{i}.cross.key"] = cross_kv[i][0] feed[f"past.{i}.cross.value"] = cross_kv[i][1] step = decode(feed) step = {k.get_any_name(): v for k, v in step.items()} self_kv = [(step[f"present.{i}.self.key"], step[f"present.{i}.self.value"]) for i in range(num_layers)] next_id = int(step["logits"][0, -1].argmax()) generated.append(next_id) t_dec = time.perf_counter() text = processor.batch_decode([generated], skip_special_tokens=True)[0].strip() return text, len(generated), t_enc - t0, t_dec - t_enc, t_dec - t0 def main(): ap = argparse.ArgumentParser(description="Transcribe audio with a converted Cohere Transcribe OpenVINO model") ap.add_argument("--model_dir", type=Path, default=Path(__file__).parent, help="Directory with the OpenVINO IR") ap.add_argument("--audio", required=True, help="Path to the input audio file (.wav/.flac/...)") ap.add_argument("--device", default="CPU") ap.add_argument("--max_new_tokens", type=int, default=256) args = ap.parse_args() import openvino as ov from transformers import AutoProcessor model_dir = args.model_dir.resolve() meta = load_meta(model_dir) stateful = "decoder_with_past_ir" in meta processor = AutoProcessor.from_pretrained(model_dir) audio = load_audio(args.audio, meta["sampling_rate"]) inputs = processor(audio, sampling_rate=meta["sampling_rate"], language="en", return_tensors="np") feats = { "input_features": _np(inputs["input_features"], np.float32), "attention_mask": _np(inputs["attention_mask"], bool), "decoder_input_ids": _np(inputs["decoder_input_ids"], np.int64), } core = ov.Core() fn = transcribe_stateful if stateful else transcribe_stateless text, n_tokens, enc_s, dec_s, total_s = fn( core, model_dir, meta, processor, feats, args.device, args.max_new_tokens ) print(f"\nModel : {model_dir} ({'stateful/KV-cache' if stateful else 'stateless'})") print(f"Audio : {args.audio}") print(f"Device : {args.device}") print("-" * 60) print(f"Transcript:\n {text}") print("-" * 60) print(f"Tokens : {n_tokens}") print(f"Encoder : {enc_s * 1e3:8.1f} ms") print(f"Decode : {dec_s * 1e3:8.1f} ms ({n_tokens / dec_s:.1f} tok/s)") print(f"Total : {total_s * 1e3:8.1f} ms") if __name__ == "__main__": main()