moonshine-tiny-en-onnx / export_encoder.py
vasyadeva's picture
Upload folder using huggingface_hub
9c1d5d9 verified
Raw
History Blame Contribute Delete
1.47 kB
import os, json, torch
from torch.export import Dim
from transformers import MoonshineForConditionalGeneration, AutoProcessor, AutoConfig
MODEL = os.environ.get('MS_MODEL', 'moonshine-ai/moonshine-tiny-uk')
OUT = os.environ.get('MS_OUT', '/work/out')
m = MoonshineForConditionalGeneration.from_pretrained(MODEL, attn_implementation='eager').eval()
class Encoder(torch.nn.Module):
def __init__(self, m): super().__init__(); self.enc = m.model.encoder
def forward(self, input_values): return self.enc(input_values).last_hidden_state
audio = torch.zeros(1, 16000 * 3)
enc = Encoder(m).eval()
torch.onnx.export(enc, (audio,), f'{OUT}/encoder.onnx', dynamo=True,
input_names=['input_values'], output_names=['encoder_hidden_states'],
dynamic_shapes={'input_values': {1: Dim('samples', min=2000, max=16000*30)}})
print('encoder exported')
proc = AutoProcessor.from_pretrained(MODEL)
vocab = proc.tokenizer.get_vocab()
with open(f'{OUT}/tokens.txt', 'w', encoding='utf-8') as f:
for piece, idx in sorted(vocab.items(), key=lambda kv: kv[1]):
f.write(f'{idx}\t{piece}\n')
cfg = AutoConfig.from_pretrained(MODEL)
json.dump({'decoder_start_token_id': cfg.decoder_start_token_id, 'eos_token_id': cfg.eos_token_id,
'bos_token_id': cfg.bos_token_id, 'vocab_size': cfg.vocab_size, 'sampling_rate': 16000},
open(f'{OUT}/decode_config.json', 'w'), indent=2)
print('tokens.txt:', len(vocab), 'entries')