moonshine-tiny-en-onnx / validate_vs_torch.py
vasyadeva's picture
Upload folder using huggingface_hub
9c1d5d9 verified
Raw
History Blame Contribute Delete
2.54 kB
import io, os, time, numpy as np, onnxruntime as ort, soundfile as sf, torch
from datasets import load_dataset, Audio
from transformers import AutoProcessor, AutoConfig, MoonshineForConditionalGeneration
MODEL = os.environ.get('MS_MODEL', 'moonshine-ai/moonshine-tiny')
OUT = os.environ.get('MS_OUT', '/work/out_en')
FLEURS = os.environ.get('MS_FLEURS', 'en_us')
L = 6
cfg, proc = AutoConfig.from_pretrained(MODEL), AutoProcessor.from_pretrained(MODEL)
torch_model = MoonshineForConditionalGeneration.from_pretrained(MODEL, attn_implementation='eager').eval()
S = lambda p: ort.InferenceSession(p, providers=['CPUExecutionProvider'])
enc, init, step = S(f'{OUT}/encoder.onnx'), S(f'{OUT}/decoder_init.onnx'), S(f'{OUT}/decoder_step.onnx')
names = [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
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())
t0 = time.time()
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(names[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())
return proc.batch_decode([ids], skip_special_tokens=True)[0], t_enc, time.time() - t0, steps
ds = load_dataset('google/fleurs', FLEURS, 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)
onnx_txt, t_enc, t_dec, steps = cached(wav)
ref = proc.batch_decode(torch_model.generate(**proc(wav, sampling_rate=16000, return_tensors='pt')),
skip_special_tokens=True)[0]
print(f'\n[{i}] {len(wav)/16000:.2f}s | encode {t_enc:.2f}s + decode {t_dec:.2f}s over {steps} steps')
print(' torch:', ref[:100])
print(' onnx :', onnx_txt[:100])
print(' match:', ref.strip() == onnx_txt.strip())