Audio8-ASR-Infinite-ONNX / scripts /audio8_infinite_graph.py
Masterx's picture
Upload scripts/audio8_infinite_graph.py
86e2529 verified
Raw History Blame Contribute Delete
24.2 kB
#!/usr/bin/env python3
"""Hand-built ONNX graphs for Audio8-ASR-Infinite (see audio8_infinite_ref.py for the IO contract).
The graphs are emitted node-by-node straight from the (memmapped, bf16) safetensors, one weight at
a time, with every initializer streamed into the external-data file as it is produced. Nothing ever
holds the whole fp32 model in RAM, which is what makes a 4 B-parameter export possible on a box
with single-digit free GB -- `torch.onnx.export` would need the full fp32 module plus its trace.
Precisions (`--precision`):
fp32 MatMul fp32 weights (parity reference; ~16 GB, not shipped)
fp16 whole graph fp16 (RMSNorm / softmax / mel in fp32), KV caches fp16, I/O otherwise fp32
int8 weight-only MatMulNBits (8-bit, asymmetric, block 32, accuracy_level 4), fp32 activations
int8dq (comparison only, not shipped) per-channel int8 weights + DynamicQuantizeLinear /
MatMulInteger, the quantize_dynamic pattern: ONE activation scale per tensor, so the
result depends on how many frames share a call and Qwen2.5's activation outliers flatten
the other channels -- measurably worse transcripts than either NBits tier
int4 weight-only MatMulNBits (4-bit, asymmetric, block 32, accuracy_level 4), fp32 activations
"""
from __future__ import annotations
import math
import os
from pathlib import Path
import numpy as np
import onnx
from onnx import TensorProto, helper
import audio8_infinite_ref as a8i
OPSET = 17
F32, F16 = TensorProto.FLOAT, TensorProto.FLOAT16
NP = {F32: np.float32, F16: np.float16}
class Graph:
def __init__(self, path: Path, act: int):
self.path = Path(path)
self.data_name = self.path.name + ".data"
self.data = open(self.path.with_name(self.data_name), "wb")
self.off = 0
self.act = act
self.nodes, self.inits, self.inputs, self.outputs = [], [], [], []
self.n = 0
self._consts: dict = {}
# ---- plumbing ----------------------------------------------------------------------------
def uid(self, p="t"):
self.n += 1
return f"{p}_{self.n}"
def init(self, name: str, arr: np.ndarray, *, onnx_type=None) -> str:
arr = np.ascontiguousarray(arr)
if arr.nbytes < 4096:
t = onnx.numpy_helper.from_array(arr, name)
self.inits.append(t)
return name
pad = (-self.off) % 4096
if pad:
self.data.write(b"\0" * pad)
self.off += pad
buf = arr.tobytes()
self.data.write(buf)
t = TensorProto()
t.name = name
t.data_type = onnx_type if onnx_type is not None else helper.np_dtype_to_tensor_dtype(arr.dtype)
t.dims.extend(arr.shape)
t.data_location = TensorProto.EXTERNAL
for k, v in (("location", self.data_name), ("offset", str(self.off)), ("length", str(len(buf)))):
e = t.external_data.add()
e.key, e.value = k, v
self.off += len(buf)
self.inits.append(t)
return name
def const(self, value, dtype=np.int64) -> str:
arr = np.asarray(value, dtype=dtype)
key = (arr.dtype.str, arr.shape, arr.tobytes())
if key not in self._consts:
self._consts[key] = self.init(self.uid("c"), arr)
return self._consts[key]
def fconst(self, value) -> str:
return self.const(value, NP[self.act])
def op(self, op_type, inputs, n_out=1, domain=None, **attrs):
outs = [self.uid(op_type.lower()) for _ in range(n_out)]
kw = {"domain": domain} if domain else {}
self.nodes.append(helper.make_node(op_type, list(inputs), outs, name=self.uid("n_" + op_type), **kw, **attrs))
return outs[0] if n_out == 1 else outs
def inp(self, name, dtype, shape):
self.inputs.append(helper.make_tensor_value_info(name, dtype, shape))
return name
def out(self, src, name, dtype, shape):
self.nodes.append(helper.make_node("Identity", [src], [name], name=self.uid("n_out")))
self.outputs.append(helper.make_tensor_value_info(name, dtype, shape))
def finish(self, meta: dict):
self.data.close()
g = helper.make_graph(self.nodes, self.path.stem, self.inputs, self.outputs, self.inits)
opsets = [helper.make_opsetid("", OPSET), helper.make_opsetid("com.microsoft", 1)]
m = helper.make_model(g, opset_imports=opsets, producer_name="winstt-audio8-infinite-export")
m.ir_version = 8
for k, v in meta.items():
e = m.metadata_props.add()
e.key, e.value = k, str(v)
onnx.save(m, str(self.path))
# ---- math helpers --------------------------------------------------------------------------
def cast(self, x, to):
return self.op("Cast", [x], to=to)
def rms(self, x, w: np.ndarray, eps: float, name: str):
"""x / sqrt(mean(x^2) + eps) * w, computed in fp32 (SimplifiedLayerNorm fusion pattern)."""
xf = self.cast(x, F32) if self.act != F32 else x
sq = self.op("Pow", [xf, self.const(2.0, np.float32)])
mean = self.op("ReduceMean", [sq], axes=[-1], keepdims=1)
den = self.op("Sqrt", [self.op("Add", [mean, self.const(eps, np.float32)])])
y = self.op("Div", [xf, den])
y = self.op("Mul", [y, self.init(name, w.astype(np.float32))])
return self.cast(y, self.act) if self.act != F32 else y
def gelu(self, x):
xf = x
e = self.op("Erf", [self.op("Div", [xf, self.fconst(math.sqrt(2.0))])])
return self.op("Mul", [self.op("Mul", [xf, self.op("Add", [e, self.fconst(1.0)])]), self.fconst(0.5)])
def silu(self, x):
return self.op("Mul", [x, self.op("Sigmoid", [x])])
def rope(self, x, cos, sin):
a, b = self.op("Split", [x], n_out=2, axis=-1)
rot = self.op("Concat", [self.op("Neg", [b]), a], axis=-1)
return self.op("Add", [self.op("Mul", [x, cos]), self.op("Mul", [rot, sin])])
def shape_dim(self, x, i):
return self.op("Gather", [self.op("Shape", [x]), self.const([i])], axis=0)
class Linear:
"""Emits y = x @ W.T (+ b) in the requested precision."""
def __init__(self, g: Graph, precision: str, block: int = 32, accuracy_level: int = 4):
self.g, self.prec, self.block, self.acc = g, precision, block, accuracy_level
self.count = 0
def __call__(self, x, W: np.ndarray, b: np.ndarray | None, name: str, quantize: bool = True):
g = self.g
N, K = W.shape
prec = self.prec if quantize else ("fp16" if self.prec == "fp16" else "fp32")
self.count += 1
if prec in ("fp32", "fp16"):
dt = np.float16 if prec == "fp16" else np.float32
y = g.op("MatMul", [x, g.init(name + ".weight_t", W.T.astype(dt))])
elif prec == "int8dq":
# Legacy quantize_dynamic pattern (per-tensor dynamic activations) — comparison only: its
# activation scale depends on how many frames share a call, and it loses accuracy.
Wt = W.T.astype(np.float32) # [K, N]
scale = np.maximum(np.abs(Wt).max(0), 1e-12) / 127.0
q = np.clip(np.rint(Wt / scale), -127, 127).astype(np.int8)
xq, xs, xz = g.op("DynamicQuantizeLinear", [x], n_out=3)
mm = g.op("MatMulInteger", [xq, g.init(name + ".weight_q8", q), xz,
g.init(name + ".weight_zp8", np.zeros(N, np.int8))])
s = g.op("Mul", [xs, g.init(name + ".weight_scale", scale.astype(np.float32))])
y = g.op("Mul", [g.op("Cast", [mm], to=F32), s])
elif prec == "int4":
from onnxruntime.capi._pybind_state import quantize_matmul_4bits
Wt = np.ascontiguousarray(W.T.astype(np.float32)) # [K, N]
kb = (K + self.block - 1) // self.block
assert K % self.block == 0, (name, K)
packed = np.zeros((N, kb, self.block // 2), np.uint8)
scales = np.zeros((N, kb), np.float32)
zp = np.zeros((N, (kb + 1) // 2), np.uint8)
quantize_matmul_4bits(packed, Wt, scales, zp, self.block, N, K, False)
y = g.op("MatMulNBits", [x, g.init(name + ".weight_q4", packed), g.init(name + ".weight_scales", scales),
g.init(name + ".weight_zp4", zp)], domain="com.microsoft",
K=K, N=N, bits=4, block_size=self.block, accuracy_level=self.acc)
elif prec == "int8":
# Weight-only 8-bit, block-wise (MatMulNBits bits=8). With accuracy_level 4 the CPU kernel
# quantizes the ACTIVATIONS per row-block too, so — unlike DynamicQuantizeLinear's single
# per-tensor scale — the result does not depend on how many frames share a call and
# Qwen2.5's activation outliers do not flatten every other channel.
from onnxruntime.capi._pybind_state import quantize_matmul_8bits
Wt = np.ascontiguousarray(W.T.astype(np.float32)) # [K, N]
kb = (K + self.block - 1) // self.block
assert K % self.block == 0, (name, K)
packed = np.zeros((N, kb, self.block), np.uint8)
scales = np.zeros((N, kb), np.float32)
zp = np.zeros((N, kb), np.uint8)
quantize_matmul_8bits(packed, Wt, scales, zp, self.block, N, K, False)
y = g.op("MatMulNBits", [x, g.init(name + ".weight_q8", packed), g.init(name + ".weight_scales", scales),
g.init(name + ".weight_zp8", zp)], domain="com.microsoft",
K=K, N=N, bits=8, block_size=self.block, accuracy_level=self.acc)
else:
raise ValueError(prec)
if b is not None:
y = g.op("Add", [y, g.init(name + ".bias", b.astype(NP[g.act]))])
return y
def attention(g: Graph, q, k_new, v_new, past_k, past_v, bias_f32, *, kv_heads: int, head_dim: int, groups: int):
"""Split-score attention: never concatenates the (large) past cache with the new keys.
q [1,H,T,D]; k_new/v_new [1,KV,T,D]; past_k/past_v [1,KV,C,D]; bias [T, C+T] (fp32).
Returns [1,KV,G*T,D] (== [1,H,T,D] after a reshape when G > 1).
"""
q = g.op("Mul", [q, g.fconst(head_dim ** -0.5)])
# [1,H,T,D] -> [1,KV,G*T,D] (heads h*G..h*G+G-1 share kv head h, like repeat_kv)
qg = g.op("Reshape", [q, g.const([1, kv_heads, -1, head_dim])]) if groups > 1 else q
sp = g.op("MatMul", [qg, g.op("Transpose", [past_k], perm=[0, 1, 3, 2])])
sn = g.op("MatMul", [qg, g.op("Transpose", [k_new], perm=[0, 1, 3, 2])])
s = g.op("Concat", [sp, sn], axis=-1)
if g.act != F32:
s = g.cast(s, F32)
if groups > 1:
# [1,KV,G*T,S] -> [1,KV,G,T,S] so the [T,S] bias broadcasts
s5 = g.op("Reshape", [s, g.op("Concat", [g.const([1, kv_heads, groups]), g.op("Shape", [bias_f32])], axis=0)])
p = g.op("Softmax", [g.op("Add", [s5, bias_f32])], axis=-1)
p = g.op("Reshape", [p, g.op("Shape", [s])])
else:
p = g.op("Softmax", [g.op("Add", [s, bias_f32])], axis=-1)
if g.act != F32:
p = g.cast(p, g.act)
C = g.shape_dim(past_k, 2)
pp = g.op("Slice", [p, g.const([0]), C, g.const([-1])])
pn = g.op("Slice", [p, C, g.const([1 << 40]), g.const([-1])])
return g.op("Add", [g.op("MatMul", [pp, past_v]), g.op("MatMul", [pn, v_new])])
def heads(g: Graph, x, n_heads: int, head_dim: int):
"""[1,T,H*D] -> [1,H,T,D]"""
r = g.op("Reshape", [x, g.const([1, -1, n_heads, head_dim])])
return g.op("Transpose", [r], perm=[0, 2, 1, 3])
def merge(g: Graph, x, width: int):
"""[1,H,T,D] -> [1,T,H*D]"""
t = g.op("Transpose", [x], perm=[0, 2, 1, 3])
return g.op("Reshape", [t, g.const([1, -1, width])])
# ================================================================================================
def build_encoder(W, dims: a8i.Dims, precision: str, path: Path, block: int = 32):
act = F16 if precision == "fp16" else F32
kv = act
g = Graph(path, act)
lin = Linear(g, precision, block)
d = dims
audio = g.inp("audio", F32, ["batch", "samples"])
frame_len = g.inp("frame_len", TensorProto.INT64, [1])
cos = g.inp("cos", F32, ["frames", d.a_head_dim])
sin = g.inp("sin", F32, ["frames", d.a_head_dim])
bias = g.inp("attn_bias", F32, ["frames", "past_plus_frames"])
past = []
for i in range(d.a_layers):
pk = g.inp(f"past_key_{i}", kv, [1, d.a_heads, "past", d.a_head_dim])
pv = g.inp(f"past_value_{i}", kv, [1, d.a_heads, "past", d.a_head_dim])
past.append((pk, pv))
if act != F32:
cos_a, sin_a = g.cast(cos, act), g.cast(sin, act)
else:
cos_a, sin_a = cos, sin
# ---- log-mel (fp32): reflect pad, DFT-as-conv, power, mel, log10 floor ----
a3 = g.op("Unsqueeze", [audio, g.const([1])]) # [B,1,S]
padded = g.op("Pad", [a3, g.const([0, 0, N_FFT_HALF, 0, 0, N_FFT_HALF])], mode="reflect")
spec = g.op("Conv", [padded, g.init("frontend.dft_basis", a8i.dft_basis())], strides=[a8i.HOP]) # [B,402,F]
F_ = g.shape_dim(spec, 2)
spec = g.op("Slice", [spec, g.const([0]), g.op("Sub", [F_, g.const([1])]), g.const([2])])
sq = g.op("Mul", [spec, spec])
re, im = g.op("Split", [sq], n_out=2, axis=1)
power = g.op("Add", [re, im]) # [B,201,F-1]
mel = g.op("MatMul", [g.init("frontend.mel_filters_t", a8i.mel_filters().T.copy()), power]) # [B,128,F-1]
lg = g.op("Div", [g.op("Log", [g.op("Max", [mel, g.const(1e-10, np.float32)])]), g.const(math.log(10.0), np.float32)])
lg = g.op("Max", [lg, g.const(a8i.LOG_MEL_FLOOR, np.float32)])
feats = g.op("Div", [g.op("Add", [lg, g.const(4.0, np.float32)]), g.const(4.0, np.float32)])
if act != F32:
feats = g.cast(feats, act)
dt = NP[act]
c1 = g.op("Conv", [feats, g.init("embedder.conv1.weight", W["audio_tower.embedder.conv1.weight"].astype(dt)),
g.init("embedder.conv1.bias", W["audio_tower.embedder.conv1.bias"].astype(dt))], pads=[2, 0])
c1 = g.gelu(c1)
c2 = g.op("Conv", [c1, g.init("embedder.conv2.weight", W["audio_tower.embedder.conv2.weight"].astype(dt)),
g.init("embedder.conv2.bias", W["audio_tower.embedder.conv2.bias"].astype(dt))],
pads=[1, 0], strides=[2])
c2 = g.gelu(c2) # [B,1280,Fc]
x = g.op("Transpose", [c2], perm=[0, 2, 1]) # [B,Fc,1280]
B = g.shape_dim(audio, 0)
T = g.shape_dim(cos, 0)
Kw = g.op("Div", [T, B])
x = g.op("Slice", [x, g.op("Neg", [Kw]), g.const([1 << 40]), g.const([1])])
x = g.op("Reshape", [x, g.const([1, -1, d.a_hidden])]) # [1,T,1280]
width = d.a_heads * d.a_head_dim
for i in range(d.a_layers):
p = f"audio_tower.layers.{i}."
h = g.rms(x, W[p + "self_attn_layer_norm.weight"], d.a_eps, f"enc.{i}.ln1")
q = heads(g, lin(h, W[p + "self_attn.q_proj.weight"], W[p + "self_attn.q_proj.bias"], f"enc.{i}.q"), d.a_heads, d.a_head_dim)
k = heads(g, lin(h, W[p + "self_attn.k_proj.weight"], None, f"enc.{i}.k"), d.a_heads, d.a_head_dim)
v = heads(g, lin(h, W[p + "self_attn.v_proj.weight"], W[p + "self_attn.v_proj.bias"], f"enc.{i}.v"), d.a_heads, d.a_head_dim)
q, k = g.rope(q, cos_a, sin_a), g.rope(k, cos_a, sin_a)
g.out(k, f"key_{i}", kv, [1, d.a_heads, "frames", d.a_head_dim])
g.out(v, f"value_{i}", kv, [1, d.a_heads, "frames", d.a_head_dim])
o = attention(g, q, k, v, past[i][0], past[i][1], bias, kv_heads=d.a_heads, head_dim=d.a_head_dim, groups=1)
o = merge(g, o, width)
x = g.op("Add", [x, lin(o, W[p + "self_attn.o_proj.weight"], W[p + "self_attn.o_proj.bias"], f"enc.{i}.o")])
h = g.rms(x, W[p + "final_layer_norm.weight"], d.a_eps, f"enc.{i}.ln2")
gt = lin(h, W[p + "mlp.gate_proj.weight"], None, f"enc.{i}.gate")
up = lin(h, W[p + "mlp.up_proj.weight"], None, f"enc.{i}.up")
x = g.op("Add", [x, lin(g.op("Mul", [g.silu(gt), up]), W[p + "mlp.down_proj.weight"], W[p + "mlp.down_proj.bias"], f"enc.{i}.down")])
x = g.rms(x, W["audio_tower.norm.weight"], d.a_eps, "enc.norm")
# group frame_len frames per token, zero-pad to max_frame_len slots, project
grp = g.op("Reshape", [x, g.op("Concat", [g.const([-1]), g.op("Mul", [frame_len, g.const([d.a_hidden])])], axis=0)])
padw = g.op("Sub", [g.const([d.projection]), g.op("Mul", [frame_len, g.const([d.a_hidden])])])
grp = g.op("Pad", [grp, g.op("Concat", [g.const([0, 0, 0]), padw], axis=0)])
e = lin(grp, W["multi_modal_projector.linear_1.weight"], None, "proj.1")
e = lin(g.gelu(e), W["multi_modal_projector.linear_2.weight"], None, "proj.2")
e = g.op("Unsqueeze", [e, g.const([0])])
if act != F32:
e = g.cast(e, F32)
g.out(e, "audio_embeds", F32, [1, "tokens", d.t_hidden])
g.finish({"winstt_audio8_infinite": "audio_encoder", "precision": precision,
"sliding_window": d.a_window, "rope_theta": d.a_theta, "head_dim": d.a_head_dim})
return lin.count
N_FFT_HALF = a8i.N_FFT // 2
def build_decoder(W, dims: a8i.Dims, precision: str, path: Path, block: int = 32, layers=None, head=True):
"""layers=(start, end) builds a shard (for fp32 parity on low-RAM boxes): a shard that does not
start at 0 takes `hidden` instead of `inputs_embeds`; one without `head` outputs `hidden`."""
act = F16 if precision == "fp16" else F32
kv = act
g = Graph(path, act)
lin = Linear(g, precision, block)
d = dims
l0, l1 = layers or (0, d.t_layers)
x_in = g.inp("inputs_embeds" if l0 == 0 else "hidden", F32, [1, "n", d.t_hidden])
ada = g.inp("ada_scale", F32, [d.t_layers, d.t_hidden])
cos = g.inp("cos", F32, ["n", d.t_head_dim])
sin = g.inp("sin", F32, ["n", d.t_head_dim])
bias = g.inp("attn_bias", F32, ["n", "past_plus_n"])
past = {}
for i in range(l0, l1):
past[i] = (g.inp(f"past_key_{i}", kv, [1, d.t_kv, "past", d.t_head_dim]),
g.inp(f"past_value_{i}", kv, [1, d.t_kv, "past", d.t_head_dim]))
x = g.cast(x_in, act) if act != F32 else x_in
cos_a, sin_a = (g.cast(cos, act), g.cast(sin, act)) if act != F32 else (cos, sin)
ada_a = g.cast(ada, act) if act != F32 else ada
groups = d.t_heads // d.t_kv
width = d.t_heads * d.t_head_dim
for i in range(l0, l1):
p = f"language_model.model.layers.{i}."
h = g.rms(x, W[p + "input_layernorm.weight"], d.t_eps, f"dec.{i}.ln1")
q = heads(g, lin(h, W[p + "self_attn.q_proj.weight"], W[p + "self_attn.q_proj.bias"], f"dec.{i}.q"), d.t_heads, d.t_head_dim)
k = heads(g, lin(h, W[p + "self_attn.k_proj.weight"], W[p + "self_attn.k_proj.bias"], f"dec.{i}.k"), d.t_kv, d.t_head_dim)
v = heads(g, lin(h, W[p + "self_attn.v_proj.weight"], W[p + "self_attn.v_proj.bias"], f"dec.{i}.v"), d.t_kv, d.t_head_dim)
q, k = g.rope(q, cos_a, sin_a), g.rope(k, cos_a, sin_a)
g.out(k, f"key_{i}", kv, [1, d.t_kv, "n", d.t_head_dim])
g.out(v, f"value_{i}", kv, [1, d.t_kv, "n", d.t_head_dim])
o = attention(g, q, k, v, past[i][0], past[i][1], bias, kv_heads=d.t_kv, head_dim=d.t_head_dim, groups=groups)
if groups > 1:
o = g.op("Reshape", [o, g.const([1, d.t_heads, -1, d.t_head_dim])])
o = merge(g, o, width)
x = g.op("Add", [x, lin(o, W[p + "self_attn.o_proj.weight"], None, f"dec.{i}.o")])
h = g.rms(x, W[p + "post_attention_layernorm.weight"], d.t_eps, f"dec.{i}.ln2")
row = g.op("Gather", [ada_a, g.const(i)], axis=0) # [2048]
h = g.op("Mul", [h, row])
gt = lin(h, W[p + "mlp.gate_proj.weight"], None, f"dec.{i}.gate")
up = lin(h, W[p + "mlp.up_proj.weight"], None, f"dec.{i}.up")
x = g.op("Add", [x, lin(g.op("Mul", [g.silu(gt), up]), W[p + "mlp.down_proj.weight"], None, f"dec.{i}.down")])
if head:
x = g.rms(x, W["language_model.model.norm.weight"], d.t_eps, "dec.norm")
n = g.shape_dim(x, 1)
last = g.op("Slice", [x, g.op("Sub", [n, g.const([1])]), n, g.const([1])]) # [1,1,2048]
last = g.op("Reshape", [last, g.const([1, d.t_hidden])])
emb = W["language_model.model.embed_tokens.weight"]
logits = lin(last, emb, None, "lm_head")
if act != F32:
logits = g.cast(logits, F32)
g.out(logits, "logits", F32, [1, d.vocab])
vw = np.concatenate([W[f"semantic_vad_heads.{j}.weight"] for j in range(d.vad_heads)], 0)
vb = np.concatenate([W[f"semantic_vad_heads.{j}.bias"] for j in range(d.vad_heads)], 0)
lastf = g.cast(last, F32) if act != F32 else last
gv = g.op("Add", [g.op("MatMul", [lastf, g.init("vad.weight_t", vw.T.astype(np.float32))]),
g.init("vad.bias", vb.astype(np.float32))])
g.out(g.op("Reshape", [gv, g.const([1, d.vad_heads, d.vad_classes])]), "vad_logits", F32,
[1, d.vad_heads, d.vad_classes])
else:
xo = g.cast(x, F32) if act != F32 else x
g.out(xo, "hidden", F32, [1, "n", d.t_hidden])
g.finish({"winstt_audio8_infinite": "decoder", "precision": precision, "layers": f"{l0}:{l1}",
"rope_theta": d.t_theta, "head_dim": d.t_head_dim})
return lin.count
# ------------------------------------------------------------------------------------------------
class OrtBackend:
"""Session backend over the exported graphs (numpy in/out)."""
def __init__(self, enc_path, dec_paths, providers=("CPUExecutionProvider",), threads=None, opt_dir=None):
import onnxruntime as ort
so = ort.SessionOptions()
if threads:
so.intra_op_num_threads = threads
so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
self.enc = ort.InferenceSession(str(enc_path), so, providers=list(providers)) if enc_path else None
self.decs = [ort.InferenceSession(str(p), so, providers=list(providers)) for p in (dec_paths or [])]
self.kv_dtype = {}
for s in ([self.enc] if self.enc else []) + self.decs:
for i in s.get_inputs():
if i.name.startswith("past_key_"):
self.kv_dtype[id(s)] = np.float16 if "float16" in i.type else np.float32
break
def encoder(self, audio, frame_len, cos, sin, bias, past_k, past_v):
dt = self.kv_dtype[id(self.enc)]
feed = {"audio": audio, "frame_len": np.array([frame_len], np.int64), "cos": cos, "sin": sin, "attn_bias": bias}
L = len(past_k)
for i in range(L):
feed[f"past_key_{i}"] = past_k[i].astype(dt, copy=False)
feed[f"past_value_{i}"] = past_v[i].astype(dt, copy=False)
names = ["audio_embeds"] + [f"key_{i}" for i in range(L)] + [f"value_{i}" for i in range(L)]
outs = self.enc.run(names, feed)
return outs[0], [o.astype(np.float32) for o in outs[1:1 + L]], [o.astype(np.float32) for o in outs[1 + L:]]
def decoder(self, embeds, ada, cos, sin, bias, past_k, past_v, all_logits=False):
assert not all_logits, "ORT decoder returns last-position logits only"
x = embeds.astype(np.float32)
nk, nv = [], []
logits = vad = None
for s in self.decs:
dt = self.kv_dtype[id(s)]
names_in = {i.name for i in s.get_inputs()}
feed = {("inputs_embeds" if "inputs_embeds" in names_in else "hidden"): x, "ada_scale": ada,
"cos": cos, "sin": sin, "attn_bias": bias}
layers = sorted(int(n.split("_")[-1]) for n in names_in if n.startswith("past_key_"))
for i in layers:
feed[f"past_key_{i}"] = past_k[i].astype(dt, copy=False)
feed[f"past_value_{i}"] = past_v[i].astype(dt, copy=False)
out_names = [o.name for o in s.get_outputs()]
res = dict(zip(out_names, s.run(out_names, feed)))
for i in layers:
nk.append(res[f"key_{i}"].astype(np.float32))
nv.append(res[f"value_{i}"].astype(np.float32))
if "hidden" in res:
x = res["hidden"]
else:
logits, vad = res["logits"], res["vad_logits"]
return logits, vad, nk, nv