File size: 4,717 Bytes
d976309
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
---
license: apache-2.0
language:
- en
pipeline_tag: automatic-speech-recognition
tags:
- litert
- tflite
- android
- gpu
- wav2vec2
- ctc
- speech-recognition
base_model: facebook/wav2vec2-base-960h
---

# wav2vec2-base-960h CTC β€” LiteRT (GPU)

English speech recognition with [wav2vec2-base-960h](https://huggingface.co/facebook/wav2vec2-base-960h)
running **fully on the LiteRT `CompiledModel` GPU** (ML Drift) β€” and with **zero FFT anywhere**:
the raw 16 kHz waveform goes straight into the 1D-conv feature extractor, so there is no mel/fbank
step even on the host. Character-level CTC (29 chars + specials), greedy decode, no language model.

![wav2vec2 CTC word onsets](assets/hero.png)
*Real model output: char-CTC word onsets for J.F. Kennedy's 1961 inaugural address
(U.S. National Archives recording, public domain).*

Ships as **two GPU graphs** β€” the fused graph exceeds the Mali whole-graph shader-compile limit
(a graph can be op-clean and still fail to compile when fused; each half compiles and runs fully
delegated):

| File | Size | Input | Output | API |
| ---- | ---- | ----- | ------ | --- |
| `w2v2_asr_frontend_fp16.tflite` | 9 MB | waveform `[1, 256000]` | features `[1, 799, 768]` | CompiledModel GPU |
| `w2v2_asr_head_fp16.tflite` | 180 MB | features `[1, 799, 768]` | CTC logits `[1, 799, 32]` | CompiledModel GPU |

## Pipeline

`16 kHz mono PCM in [-1, 1], zero-padded to the fixed 16 s window β†’ [GPU] conv frontend β†’
[GPU] 12-layer transformer + lm_head β†’ host greedy-CTC over the valid frames`

- Valid frames for `n` samples: run `L=(L-k)//s+1` over the conv stack
  `(10,5)(3,2)(3,2)(3,2)(3,2)(2,2)(2,2)` β€” 50 Hz frames (16 s β†’ 799).
- Blank id 0 (`<pad>`), `|` = word delimiter (`tokens.txt`, index-ordered).
- Greedy char-CTC without an LM has the model's known spelling quirks on hard words
  (e.g. GRAVED/GRAVE); add a beam+LM host-side if you need the last WER point.

## Minimal usage β€” Python

```python
import numpy as np, torch, torchaudio
from ai_edge_litert.interpreter import Interpreter

wave, sr = torchaudio.load("speech.wav")             # 16 kHz mono, [-1,1]
x = torch.zeros(1, 256000); n = min(wave.shape[1], 256000)
x[0, :n] = wave[0, :n]

def run(path, inp):
    it = Interpreter(model_path=path); it.allocate_tensors()
    d = it.get_input_details()[0]
    it.set_tensor(d["index"], inp.astype(np.float32)); it.invoke()
    return it.get_tensor(it.get_output_details()[0]["index"])

feat = run("w2v2_asr_frontend_fp16.tflite", x.numpy())
logits = run("w2v2_asr_head_fp16.tflite", feat)[0]    # [799, 32]

L = n
for k, s in [(10,5),(3,2),(3,2),(3,2),(3,2),(2,2),(2,2)]:
    L = (L - k) // s + 1
tokens = open("tokens.txt").read().splitlines()
out, prev = [], -1
for i in logits[:L].argmax(-1):
    if i != prev and i != 0: out.append(tokens[int(i)])
    prev = i
print("".join(out).replace("|", " ").strip())
```

## Minimal usage β€” Kotlin (Android)

```kotlin
val frontend = CompiledModel.create(frontendPath, CompiledModel.Options(Accelerator.GPU), null)
val head = CompiledModel.create(headPath, CompiledModel.Options(Accelerator.GPU), null)
val fIn = frontend.createInputBuffers(); val fOut = frontend.createOutputBuffers()
val hIn = head.createInputBuffers(); val hOut = head.createOutputBuffers()

fIn[0].writeFloat(pcm)                        // [-1,1] floats, zero-padded to 256000
frontend.run(fIn, fOut)
hIn[0].writeFloat(fOut[0].readFloat())        // features [1,799,768]
head.run(hIn, hOut)
val logits = hOut[0].readFloat()              // [799 * 32], readback syncs the GPU
// greedy CTC over the valid frames: argmax per frame, drop blanks (id 0) + repeats,
// map through tokens.txt, '|' -> space
```

## On-device performance (Pixel 8a, CompiledModel GPU)

- frontend 448 ms + head 391 ms per 16 s window (RTF β‰ˆ 0.05); GPU compile 0.7 s + 1.5 s.
- Device logits vs desktop float reference: corr 0.9928 (valid region), per-frame argmax
  agreement 97.0 %; transcript matches the desktop reference.

## Conversion notes

Converted with litert-torch, numerically exact (tflite vs PyTorch: corr 1.000000):
GELU β†’ tanh-GELU; frontend GroupNorm β†’ 4D-reshape group-norm (avoids GATHER_ND);
pos_conv weight-norm folded to a static weight; the all-valid bidirectional attention mask
removed (fixed window β†’ plain SDPA). The CTC head is a plain Linear β€” logits come out raw.

## Sources & license

- Model: [facebook/wav2vec2-base-960h](https://huggingface.co/facebook/wav2vec2-base-960h) β€” Apache-2.0.
- Paper: [wav2vec 2.0: A Framework for Self-Supervised Learning of Speech Representations](https://arxiv.org/abs/2006.11477).
- Hero audio: J.F. Kennedy inaugural address (1961), U.S. National Archives β€” public domain.