import argparse
import torch
from scipy.io import wavfile
from time import time
from models import model_dict
from utils.lpcnet_features import load_lpcnet_features
from utils import endoscopy
debug = False
if debug:
args = type('dummy', (object,),
{
'input' : 'testitems/all_0_orig.se',
'checkpoint' : 'testout/checkpoints/checkpoint_epoch_5.pth',
'output' : 'out.wav',
})()
else:
parser = argparse.ArgumentParser()
parser.add_argument('input', type=str, help='path to input features')
parser.add_argument('checkpoint', type=str, help='checkpoint file')
parser.add_argument('output', type=str, help='output file')
parser.add_argument('--debug', action='store_true', help='enables debug output')
args = parser.parse_args()
torch.set_num_threads(2)
input_folder = args.input
checkpoint_file = args.checkpoint
output_file = args.output
if not output_file.endswith('.wav'):
output_file += '.wav'
checkpoint = torch.load(checkpoint_file, map_location="cpu")
if not 'name' in checkpoint['setup']['model']:
print(f'warning: did not find model name entry in setup, using pitchpostfilter per default')
model_name = 'pitchpostfilter'
else:
model_name = checkpoint['setup']['model']['name']
model = model_dict[model_name](*checkpoint['setup']['model']['args'], **checkpoint['setup']['model']['kwargs'])
model.load_state_dict(checkpoint['state_dict'])
setup = checkpoint['setup']
testdata = load_lpcnet_features(input_folder)
features = testdata['features']
periods = testdata['periods']
if args.debug:
endoscopy.init()
start = time()
output = model.process(features, periods, debug=args.debug)
elapsed = time() - start
print(f"[timing] inference took {elapsed * 1000} ms")
wavfile.write(output_file, 16000, output.cpu().numpy())
if args.debug:
endoscopy.close()