Views
No views yet
1import onnxruntime as ort
2
3encoder_session = ort.InferenceSession("encoder_model.onnx")
4decoder_session = ort.InferenceSession("decoder_model.onnx")
5
6encoder_inputs = {encoder_session.get_inputs()[0].name: dummy_encoder_inputs.numpy()}
7encoder_outputs = encoder_session.run(None, encoder_inputs)[0]
8
9decoder_inputs = {decoder_session.get_inputs()[0].name: encoder_outputs}
10decoder_outputs = decoder_session.run(None, decoder_inputs)[0]
11
12# Print the results
13print("Encoder Output Shape:", encoder_outputs.shape)
14print("Decoder Output Shape:", decoder_outputs.shape)1import torch
2import torch.nn as nn
3from transformers import MimiModel
4
5class MimiEncoder(nn.Module):
6 def __init__(self, model):
7 super(MimiEncoder, self).__init__()
8 self.model = model
9
10 def forward(self, input_values, padding_mask=None):
11 return self.model.encode(input_values, padding_mask=padding_mask).audio_codes
12
13class MimiDecoder(nn.Module):
14 def __init__(self, model):
15 super(MimiDecoder, self).__init__()
16 self.model = model
17
18 def forward(self, audio_codes, padding_mask=None):
19 return self.model.decode(audio_codes, padding_mask=padding_mask).audio_values
20
21model = MimiModel.from_pretrained("kyutai/mimi")
22encoder = MimiEncoder(model)
23decoder = MimiDecoder(model)
24
25dummy_encoder_inputs = torch.randn((5, 1, 82500))
26torch.onnx.export(
27 encoder,
28 dummy_encoder_inputs,
29 "encoder_model.onnx",
30 export_params=True,
31 opset_version=14,
32 do_constant_folding=True,
33 input_names=['input_values'],
34 output_names=['audio_codes'],
35 dynamic_axes={
36 'input_values': {0: 'batch_size', 1: 'num_channels', 2: 'sequence_length'},
37 'audio_codes': {0: 'batch_size', 2: 'codes_length'},
38 },
39)
40
41dummy_decoder_inputs = torch.randint(100, (4, model.config.num_quantizers, 91))
42torch.onnx.export(
43 decoder,
44 dummy_decoder_inputs,
45 "decoder_model.onnx",
46 export_params=True,
47 opset_version=14,
48 do_constant_folding=True,
49 input_names=['audio_codes'],
50 output_names=['audio_values'],
51 dynamic_axes={
52 'audio_codes': {0: 'batch_size', 2: 'codes_length'},
53 'audio_values': {0: 'batch_size', 1: 'num_channels', 2: 'sequence_length'},
54 },
55)