Download scripts/audio8_infinite_graph.py from Masterx/Audio8-ASR-Infinite-ONNX: direct link, hf CLI and curl.
- Browser
- Download file 24.2 kB
-
https://huggingface.co/Masterx/Audio8-ASR-Infinite-ONNX/resolve/main/scripts/audio8_infinite_graph.py
- Command line
-
hf download hf://Masterx/Audio8-ASR-Infinite-ONNX/scripts/audio8_infinite_graph.py
-
curl -L -o audio8_infinite_graph.py https://huggingface.co/Masterx/Audio8-ASR-Infinite-ONNX/resolve/main/scripts/audio8_infinite_graph.py
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 | |