Views
No views yet
torch.amp.autocast(dtype=float16) — produces properly typed FP16 graphs that work on all onnxruntime execution providers..data files| Provider | RTF | Decode Speed | Hardware |
|---|---|---|---|
| CUDA EP | 0.04-0.10 | 9-33 ms/tok | RTX 4090 |
| CPU EP | 0.44-0.62 | 101-183 ms/tok | Ryzen 5 5600X |
| File | Size | Description |
|---|---|---|
encoder.onnx | 361 MB | Mel spectrogram → audio features |
decoder_init.onnx | 1.2 GB | Prefill: input embeds → logits + KV cache |
decoder_step.onnx | 1.2 GB | Autoregressive: token + KV cache → logits + updated KV cache |
embed_tokens.bin | 297 MB | Token embeddings, float16, shape [151936, 1024] |
tokenizer.json | 11 MB | HuggingFace tokenizers format |
present_keys: [num_layers=28, batch, kv_heads=8, seq_len, head_dim=128]present_values: [num_layers=28, batch, kv_heads=8, seq_len, head_dim=128][1, 128, T] → audio features [1, T/8, 1024]audio_pad placeholders)embed_tokens.bin, replace audio_pad positions with encoder featuresim_end or endoftext)float16. On CPU EP, onnxruntime auto-promotes to float32 internally.1import numpy as np
2import onnxruntime as ort
3
4# Load sessions
5encoder = ort.InferenceSession("encoder.onnx", providers=["CUDAExecutionProvider", "CPUExecutionProvider"])
6decoder_init = ort.InferenceSession("decoder_init.onnx", providers=["CUDAExecutionProvider", "CPUExecutionProvider"])
7decoder_step = ort.InferenceSession("decoder_step.onnx", providers=["CUDAExecutionProvider", "CPUExecutionProvider"])
8
9# Load embeddings
10embed_tokens = np.fromfile("embed_tokens.bin", dtype=np.float16).reshape(151936, 1024)
11
12# Compute mel spectrogram from audio (16kHz, 128 bins)
13mel = compute_log_mel_spectrogram(audio) # [1, 128, T], float16
14
15# Encode
16audio_features = encoder.run(None, {"mel": mel})[0] # [1, T/8, 1024]
17
18# Build prompt and embed (see inference pipeline above)
19prompt_embeds = build_and_embed_prompt(audio_features, embed_tokens) # [1, seq_len, 1024], float16
20
21# Prefill
22position_ids = np.arange(seq_len, dtype=np.int64).reshape(1, -1)
23logits, present_keys, present_values = decoder_init.run(None, {
24 "input_embeds": prompt_embeds,
25 "position_ids": position_ids,
26})
27
28# Greedy decode
29next_token = int(np.argmax(logits[0, -1, :]))
30generated = [next_token]
31cur_pos = seq_len
32
33while next_token not in (151645, 151643): # im_end, endoftext
34 token_embed = embed_tokens[next_token][np.newaxis, np.newaxis, :] # [1, 1, 1024]
35 logits, present_keys, present_values = decoder_step.run(None, {
36 "input_embeds": token_embed,
37 "position_ids": np.array([[cur_pos]], dtype=np.int64),
38 "past_keys": present_keys,
39 "past_values": present_values,
40 })
41 next_token = int(np.argmax(logits[0, -1, :]))
42 generated.append(next_token)
43 cur_pos += 1
44
45# Decode tokens
46text = tokenizer.decode(generated, skip_special_tokens=True)--dtype float16. The key modification wraps torch.onnx.export with torch.amp.autocast(device, dtype=torch.float16), which keeps LayerNorm in float32 while running matmuls in float16 — producing properly typed ONNX graphs.onnxconverter_common or onnxruntime.transformers.float16 produces ONNX graphs with type mismatches between nodes. These work on CPU EP (where FP16 is auto-promoted to FP32 anyway) but produce garbage on CUDA EP. Native export from torch with autocast avoids this entirely.