moonshine-tiny-uk-onnx / validate_cached.py
vasyadeva's picture
Upload folder using huggingface_hub
66d1170 verified
Raw
History Blame Contribute Delete
2.81 kB
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]}')