import os
import argparse
import sys
sys.path.append(os.path.join(os.path.dirname(__file__), '../weight-exchange'))
parser = argparse.ArgumentParser()
parser.add_argument('checkpoint', type=str, help='model checkpoint')
parser.add_argument('output_dir', type=str, help='output folder')
args = parser.parse_args()
import torch
import numpy as np
import plc
from wexchange.torch import dump_torch_weights
from wexchange.c_export import CWriter, print_vector
def c_export(args, model):
message = f"Auto generated from checkpoint {os.path.basename(args.checkpoint)}"
writer = CWriter(os.path.join(args.output_dir, "plc_data"), message=message, model_struct_name='PLCModel')
writer.header.write(
f"""
#include "opus_types.h"
"""
)
dense_layers = [
('dense_in', "plc_dense_in"),
('dense_out', "plc_dense_out")
]
for name, export_name in dense_layers:
layer = model.get_submodule(name)
dump_torch_weights(writer, layer, name=export_name, verbose=True, quantize=False, scale=None)
gru_layers = [
("gru1", "plc_gru1"),
("gru2", "plc_gru2"),
]
max_rnn_units = max([dump_torch_weights(writer, model.get_submodule(name), export_name, verbose=True, input_sparse=False, quantize=True, scale=None, recurrent_scale=None)
for name, export_name in gru_layers])
writer.header.write(
f"""
#define PLC_MAX_RNN_UNITS {max_rnn_units}
"""
)
writer.close()
if __name__ == "__main__":
os.makedirs(args.output_dir, exist_ok=True)
checkpoint = torch.load(args.checkpoint, map_location='cpu')
model = plc.PLC(*checkpoint['model_args'], **checkpoint['model_kwargs'])
model.load_state_dict(checkpoint['state_dict'], strict=False)
c_export(args, model)