| 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]}') |
|
|