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())