import io, time, numpy as np, onnxruntime as ort, soundfile as sf from datasets import load_dataset, Audio from transformers import AutoProcessor, AutoConfig OUT='/work/out'; L=6 cfg = AutoConfig.from_pretrained('moonshine-ai/moonshine-tiny-uk') proc = AutoProcessor.from_pretrained('moonshine-ai/moonshine-tiny-uk') S = lambda p: ort.InferenceSession(p, providers=['CPUExecutionProvider']) enc, init, step, plain = S(f'{OUT}/encoder.onnx'), S(f'{OUT}/decoder_init.onnx'), S(f'{OUT}/decoder_step.onnx'), S(f'{OUT}/decoder.onnx') step_inputs = [i.name for i in step.get_inputs()] def cached(wav, max_new=96): t0=time.time(); hs = enc.run(None, {'input_values': wav[None,:]})[0]; t_enc=time.time()-t0 t0=time.time() out = init.run(None, {'input_ids': np.array([[cfg.decoder_start_token_id]], dtype=np.int64), 'encoder_hidden_states': hs}) logits, kv = out[0], out[1:] self_kv, cross_kv = list(kv[:2*L]), list(kv[2*L:]) ids, steps = [], 0 nxt = int(logits[0,-1].argmax()) while nxt != cfg.eos_token_id and steps < max_new: ids.append(nxt); steps += 1 feed = {'input_ids': np.array([[nxt]], dtype=np.int64), 'cache_position': np.array([steps], dtype=np.int64), 'encoder_hidden_states': hs} for name, arr in zip(step_inputs[3:], self_kv + cross_kv): feed[name] = arr out = step.run(None, feed) logits, self_kv = out[0], list(out[1:1+2*L]) nxt = int(logits[0,-1].argmax()) t_dec=time.time()-t0 return proc.batch_decode([ids], skip_special_tokens=True)[0], t_enc, t_dec, steps def uncached(wav, max_new=96): hs = enc.run(None, {'input_values': wav[None,:]})[0] ids = np.array([[cfg.decoder_start_token_id]], dtype=np.int64) t0=time.time() for _ in range(max_new): logits = plain.run(None, {'input_ids': ids, 'encoder_hidden_states': hs})[0] nxt = int(logits[0,-1].argmax()) if nxt == cfg.eos_token_id: break ids = np.concatenate([ids, [[nxt]]], axis=1) return proc.batch_decode(ids[:,1:], skip_special_tokens=True)[0], time.time()-t0 ds = load_dataset('google/fleurs','uk_ua',split='test',streaming=True).cast_column('audio', Audio(decode=False)) for i, s in enumerate(ds): if i >= 3: break wav, sr = sf.read(io.BytesIO(s['audio']['bytes']), dtype='float32'); wav = wav.astype(np.float32) txt_c, t_enc, t_dec, steps = cached(wav) txt_u, t_u = uncached(wav) print(f'\n[{i}] {len(wav)/16000:.2f}s audio') print(f' cached : encode {t_enc:.2f}s + decode {t_dec:.2f}s over {steps} steps ({t_dec/max(steps,1)*1000:.0f} ms/token)') print(f' uncached: decode {t_u:.2f}s') print(f' same text: {txt_c.strip() == txt_u.strip()}') print(f' text: {txt_c[:90]}')