feat(npu): KV-cache greedy decode + decoder on the OpenVINO GPU EP (T3.4)

Merged-decoder (use_cache_branch/past_key_values) loop discovered from the
graph at load; step 1 runs the no-cache branch over the prompt, later steps
feed one token + cached KVs (decoder slots refreshed, encoder cross-attn KVs
kept from step 1). Decoder session lands on the Intel iGPU via the OpenVINO
GPU plugin (same staged runtime) with CPU-EP fallback and a
WA_ONNX_DECODER_DEVICE=cpu|gpu override; DirectML backends keep the CPU
decoder. Legacy decoder_model.onnx is removed once the merged graph loads.

Measured (Core Ultra 5 135U, base.en, 19s TTS clip, identical transcripts):
decoder on OpenVINO/GPU infer=1288ms vs CPU-EP infer=6098ms
This commit is contained in:
iamdoubz
2026-07-15 19:48:14 -05:00
parent 177f1f7d90
commit 402ffe4586
+226 -23
View File
@@ -7,21 +7,24 @@
//!
//! Split of work, validated by the T3.4 spike (encoder ~3.6× faster on NPU):
//! - **encoder** (fixed `[1,80,3000]` shape) runs on the **accelerator** (NPU or
//! GPU) — the expensive graph, and the fixed shape accelerators want;
//! - **decoder** (dynamic, autoregressive) runs greedy on the **CPU** EP — a
//! dynamic KV loop is a poor accelerator fit and the cheap half anyway.
//!
//! ponytail: greedy, no-KV-cache decode (re-feeds the full token prefix each
//! step). Correct and simple; windows are short so the token count is small.
//! Add `decoder_with_past` only if the decoder shows up in a profile.
//! DirectML GPU) — the expensive graph, and the fixed shape accelerators want;
//! - **decoder** (dynamic, autoregressive) decodes greedily **with a KV cache**
//! (Optimum's merged export: `use_cache_branch` + `past_key_values`), so each
//! step feeds one token instead of re-running the whole prefix. On the NPU
//! path the decoder session itself goes to the **OpenVINO GPU EP** (Intel
//! iGPU, same runtime bundle) so sustained CPU stays near-idle; it falls back
//! to the CPU EP when no iGPU is usable, and `WA_ONNX_DECODER_DEVICE=cpu|gpu`
//! overrides the choice. DirectML backends keep the CPU decoder — the DML EP
//! handles per-step-growing KV shapes poorly.
use super::{AudioWindow, SegmentSink, Transcriber, TrxError};
use crate::models::{BackendId, TranscriptSegment};
use ort::execution_providers::{
DirectMLExecutionProvider, ExecutionProviderDispatch, OpenVINOExecutionProvider,
};
use ort::session::Session;
use ort::value::Tensor;
use ort::session::{Session, SessionInputValue};
use ort::value::{Tensor, ValueType};
use std::borrow::Cow;
use std::collections::HashMap;
use std::path::Path;
use std::sync::atomic::{AtomicU64, Ordering};
@@ -30,6 +33,36 @@ use std::sync::Mutex;
/// Safety cap on generated tokens per 30 s window (Whisper's own max is 448).
const MAX_NEW_TOKENS: usize = 224;
/// One KV-cache tensor pair of the merged decoder, discovered from the graph
/// at load: the `past_key_values.*` input, its matching `present.*` output,
/// whether it's decoder self-attention (refreshed every step) or encoder
/// cross-attention (computed once on the first step, then passed through
/// untouched by the cache branch), and the dummy shape to feed before any
/// past exists (see [`dummy_past_dims`]).
struct KvSlot {
past: String,
present: String,
self_attention: bool,
dummy_dims: Vec<usize>,
}
/// `"past_key_values.3.decoder.key"` → `"present.3.decoder.key"` — Optimum's
/// fixed naming convention for the merged decoder's cache I/O.
fn present_name_for(past: &str) -> String {
format!("present.{}", past.trim_start_matches("past_key_values."))
}
/// The "no past yet" shape for a KV input: batch → 1, any other dynamic (1)
/// dim → **1** (a dummy length-1 past — ORT's raw-data tensor creation rejects
/// zero-length dims, and the no-cache branch never reads the values anyway;
/// same dummy Optimum's own runner feeds), static dims kept.
fn dummy_past_dims(dims: &[i64]) -> Vec<usize> {
dims.iter()
.enumerate()
.map(|(i, &d)| if i == 0 || d < 0 { 1 } else { d as usize })
.collect()
}
pub struct OnnxTranscriber {
encoder: Mutex<Session>,
decoder: Mutex<Session>,
@@ -38,6 +71,13 @@ pub struct OnnxTranscriber {
dec_ids: String,
dec_hidden: String,
dec_logits: String,
/// The merged decoder's KV-cache plumbing (see [`KvSlot`]) plus its
/// branch selector input.
kv_slots: Vec<KvSlot>,
dec_use_cache: String,
/// Which EP the decoder session actually landed on ("OpenVINO/GPU" or
/// "CPU") — surfaced so `load_transcriber`'s engine log tells the truth.
decoder_ep: &'static str,
detok: Detok,
/// `[decoder_start_token_id, ...forced_decoder_ids]` — the fixed prompt
/// prefix before generation begins.
@@ -120,29 +160,76 @@ impl OnnxTranscriber {
let hidden = Tensor::from_array((hidden_shape, hidden_data))
.map_err(|e| TrxError::Inference(e.to_string()))?;
// Greedy decode on the CPU.
// Greedy KV-cache decode. Step 1 runs the merged graph's no-cache
// branch over the full prompt and yields every present KV; later steps
// feed a single token plus the cache, refreshing only the decoder
// self-attention slots — the cache branch passes encoder
// cross-attention KVs through untouched, so the step-1 tensors stay
// authoritative (same contract as Optimum's own runner).
// ponytail: KVs round-trip host memory each step; wire ort IoBinding
// to pin them on-device if GPU-decoder profiling shows the copies matter.
let mut kv_cache: Vec<Tensor<f32>> = self
.kv_slots
.iter()
.map(|slot| {
let len = slot.dummy_dims.iter().product::<usize>();
Tensor::from_array((slot.dummy_dims.clone(), vec![0.0f32; len]))
})
.collect::<Result<_, _>>()
.map_err(|e| TrxError::Inference(e.to_string()))?;
let mut tokens = self.initial_tokens.clone();
let mut cached = false; // flips true once step 1 has filled kv_cache
for _ in 0..MAX_NEW_TOKENS {
let ids: Vec<i64> = tokens.clone();
let ids: Vec<i64> = if cached {
vec![*tokens.last().unwrap_or(&self.eot)]
} else {
tokens.clone()
};
let ids_tensor = Tensor::from_array(([1usize, ids.len()], ids))
.map_err(|e| TrxError::Inference(e.to_string()))?;
let branch = Tensor::from_array(([1usize], vec![cached]))
.map_err(|e| TrxError::Inference(e.to_string()))?;
let next = {
// Presents are extracted (copied) inside the outputs' borrow scope,
// then written back into kv_cache after it ends.
let (next, fresh_kvs) = {
let mut dec = self
.decoder
.lock()
.map_err(|_| TrxError::Inference("decoder mutex poisoned".into()))?;
let mut inputs: Vec<(Cow<'_, str>, SessionInputValue<'_>)> = vec![
(Cow::from(self.dec_ids.as_str()), ids_tensor.view().into()),
(Cow::from(self.dec_hidden.as_str()), hidden.view().into()),
(Cow::from(self.dec_use_cache.as_str()), branch.view().into()),
];
for (slot, kv) in self.kv_slots.iter().zip(kv_cache.iter()) {
inputs.push((Cow::from(slot.past.as_str()), kv.view().into()));
}
let outputs = dec
.run(ort::inputs![
self.dec_ids.as_str() => ids_tensor.view(),
self.dec_hidden.as_str() => hidden.view(),
])
.run(inputs)
.map_err(|e| TrxError::Inference(e.to_string()))?;
let (shape, data) = outputs[self.dec_logits.as_str()]
.try_extract_tensor::<f32>()
.map_err(|e| TrxError::Inference(e.to_string()))?;
self.argmax_last(shape, data)
let next = self.argmax_last(shape, data);
let mut fresh: Vec<(usize, Vec<usize>, Vec<f32>)> = Vec::new();
for (i, slot) in self.kv_slots.iter().enumerate() {
if slot.self_attention || !cached {
let (s, d) = outputs[slot.present.as_str()]
.try_extract_tensor::<f32>()
.map_err(|e| TrxError::Inference(e.to_string()))?;
fresh.push((i, s.iter().map(|d| *d as usize).collect(), d.to_vec()));
}
}
(next, fresh)
};
for (i, dims, data) in fresh_kvs {
kv_cache[i] = Tensor::from_array((dims, data))
.map_err(|e| TrxError::Inference(e.to_string()))?;
}
cached = true;
if next == self.eot {
break;
@@ -180,6 +267,11 @@ impl OnnxTranscriber {
best
}
/// Which EP the decoder session landed on — "OpenVINO/GPU" or "CPU".
pub fn decoder_ep(&self) -> &'static str {
self.decoder_ep
}
fn segment(&self, text: String, start_ms: u64, end_ms: u64) -> Option<TranscriptSegment> {
if text.is_empty() {
return None;
@@ -252,7 +344,7 @@ impl Transcriber for OnnxTranscriber {
Self::ensure_runtime_env(runtime_dll);
let encoder_path = model.join("encoder_model.onnx");
let decoder_path = model.join("decoder_model.onnx");
let decoder_path = model.join("decoder_model_merged.onnx");
let tokenizer_path = model.join("tokenizer.json");
let gen_config_path = model.join("generation_config.json");
@@ -263,11 +355,46 @@ impl Transcriber for OnnxTranscriber {
.commit_from_file(&encoder_path)
.map_err(|e| TrxError::Load(e.to_string()))?;
// Decoder on the default CPU EP (dynamic shapes; poor NPU fit).
let decoder = Session::builder()
.map_err(|e| TrxError::Load(e.to_string()))?
.commit_from_file(&decoder_path)
.map_err(|e| TrxError::Load(e.to_string()))?;
// Decoder session. The KV-cache loop is dynamic-shaped and
// autoregressive — a poor NPU fit — but it needn't burn CPU either: on
// the NPU path, try the same OpenVINO runtime's GPU plugin (Intel
// iGPU) first and fall back to the default CPU EP if no GPU is usable.
// `WA_ONNX_DECODER_DEVICE=cpu|gpu` overrides (debug escape hatch, like
// WA_NPU_RUNTIME_URL). DirectML backends keep the CPU decoder — the
// DML EP handles per-step-growing KV shapes poorly.
let want_gpu = match std::env::var("WA_ONNX_DECODER_DEVICE").as_deref() {
Ok("cpu") => false,
Ok("gpu") => true,
_ => matches!(backend, BackendId::Npu),
};
let mut decoder_ep = "CPU";
let mut decoder = None;
if want_gpu {
let gpu = Session::builder()
.and_then(|b| {
b.with_execution_providers([OpenVINOExecutionProvider::default()
.with_device_type("GPU")
.build()
.error_on_failure()])
})
.and_then(|b| b.commit_from_file(&decoder_path));
match gpu {
Ok(s) => {
decoder_ep = "OpenVINO/GPU";
decoder = Some(s);
}
Err(e) => tracing::warn!(
"OpenVINO GPU EP unavailable for the decoder ({e}); using the CPU EP"
),
}
}
let decoder = match decoder {
Some(s) => s,
None => Session::builder()
.map_err(|e| TrxError::Load(e.to_string()))?
.commit_from_file(&decoder_path)
.map_err(|e| TrxError::Load(e.to_string()))?,
};
let enc_input = Self::find_input(&encoder, "input_features")?;
let enc_output = Self::find_output(&encoder, "last_hidden_state")
@@ -275,6 +402,53 @@ impl Transcriber for OnnxTranscriber {
let dec_ids = Self::find_input(&decoder, "input_ids")?;
let dec_hidden = Self::find_input(&decoder, "encoder_hidden")?;
let dec_logits = Self::find_output(&decoder, "logits")?;
let dec_use_cache = Self::find_input(&decoder, "use_cache_branch")?;
// KV-cache plumbing, discovered rather than hardcoded so the layer
// count comes from the graph (whisper-base: 6 layers → 24 slots).
let mut kv_slots: Vec<KvSlot> = Vec::new();
for input in &decoder.inputs {
if !input.name.starts_with("past_key_values.") {
continue;
}
let dims: Vec<i64> = match &input.input_type {
ValueType::Tensor { shape, .. } => shape.to_vec(),
other => {
return Err(TrxError::Load(format!(
"KV input '{}' has non-tensor type {other:?}",
input.name
)))
}
};
let present = present_name_for(&input.name);
if !decoder.outputs.iter().any(|o| o.name == present) {
return Err(TrxError::Load(format!(
"decoder missing present output '{present}'"
)));
}
kv_slots.push(KvSlot {
self_attention: input.name.contains(".decoder."),
dummy_dims: dummy_past_dims(&dims),
past: input.name.clone(),
present,
});
}
if kv_slots.is_empty() {
return Err(TrxError::Load(
"decoder has no past_key_values inputs — expected the merged (KV-cache) export"
.into(),
));
}
// A pre-KV-cache install leaves the old plain decoder behind — 208 MB
// of dead weight once the merged graph loads. Best-effort removal.
let legacy = model.join("decoder_model.onnx");
if legacy.exists() {
match std::fs::remove_file(&legacy) {
Ok(()) => tracing::info!("removed superseded decoder_model.onnx"),
Err(e) => tracing::warn!("couldn't remove superseded decoder_model.onnx: {e}"),
}
}
let detok = Detok::from_tokenizer_json(&tokenizer_path)?;
let gen = GenConfig::from_file(&gen_config_path)?;
@@ -290,6 +464,9 @@ impl Transcriber for OnnxTranscriber {
dec_ids,
dec_hidden,
dec_logits,
kv_slots,
dec_use_cache,
decoder_ep,
detok,
initial_tokens,
special_floor: gen.decoder_start,
@@ -457,6 +634,28 @@ mod tests {
assert_eq!(d.get(&'A'), Some(&b'A'));
}
#[test]
fn present_name_follows_optimum_convention() {
assert_eq!(
present_name_for("past_key_values.3.decoder.key"),
"present.3.decoder.key"
);
assert_eq!(
present_name_for("past_key_values.0.encoder.value"),
"present.0.encoder.value"
);
}
#[test]
fn dummy_past_dims_pins_dynamic_axes_to_one_and_keeps_static_ones() {
// whisper-base decoder self-attention past: [batch, 8, past_seq, 64]
// with batch/past_seq dynamic — both pin to 1 (ORT rejects 0-length
// dims for raw-data tensors; the no-cache branch ignores the values).
assert_eq!(dummy_past_dims(&[-1, 8, -1, 64]), vec![1, 8, 1, 64]);
// A static trailing shape is preserved (batch still pins to 1).
assert_eq!(dummy_past_dims(&[-1, 8, 1500, 64]), vec![1, 8, 1500, 64]);
}
#[test]
fn byte_decoder_covers_all_256_bytes() {
let d = byte_decoder();
@@ -481,10 +680,14 @@ mod tests {
return;
};
let wav = std::env::var("WA_NPU_TEST_WAV").expect("set WA_NPU_TEST_WAV");
// Surface the engine's own EP-selection logs (GPU-decoder fallback
// warnings etc.) — the test harness has no subscriber otherwise.
let _ = tracing_subscriber::fmt().with_env_filter("info").try_init();
let t0 = std::time::Instant::now();
let t = OnnxTranscriber::load(Path::new(&model_dir), BackendId::Npu, None)
.expect("load NPU transcriber");
let load_ms = t0.elapsed().as_millis();
eprintln!("[spike] decoder EP: {}", t.decoder_ep());
let t1 = std::time::Instant::now();
let segs = t.transcribe_file(Path::new(&wav)).expect("transcribe");
let infer_ms = t1.elapsed().as_millis();