# import main Flask class and request object
from flask import Flask, request #Flask for post requests
import xml.etree.ElementTree as ET #XML tree toolkit
import os.path, time, torch
from argparse import ArgumentParser
from nemo.collections.asr.metrics.wer import WER, word_error_rate
from nemo.collections.asr.models import EncDecCTCModel
from nemo.utils import logging
#Get Phrase from Library
def getPhrase(itemNum):
tree = ET.parse('phrase.xml')
root = tree.getroot()
text = root.find('.//phrase[@id="{value}"]'.format(value=itemNum)).text
print("{}. {}".format(itemNum, text))
return text
# Get Key of phrase in-time
def getKey(string):
#!echo $text | phonemize > key.txt
with open('key.txt','r') as file:
key = file.read()
return(key)
#Run ASR
def Grade(key, input):
parser = ArgumentParser()
parser.add_argument(
"--asr_model", type=str, default="QuartzNet15x5Base-En", required=True, help="Pass: 'QuartzNet15x5Base-En'",
)
parser.add_argument("--dataset", type=str, required=True, help="path to evaluation data")
parser.add_argument("--batch_size", type=int, default=4)
parser.add_argument("--wer_tolerance", type=float, default=1.0, help="used by test")
parser.add_argument(
"--normalize_text", default=True, type=bool, help="Normalize transcripts or not. Set to False for non-English."
)
args = parser.parse_args(["--dataset", "dataset.json", "--asr_model", "QuartzNet15x5Base-En"])
torch.set_grad_enabled(False)
if args.asr_model.endswith('.nemo'):
logging.info(f"Using local ASR model from {args.asr_model}")
asr_model = EncDecCTCModel.restore_from(restore_path=args.asr_model)
else:
logging.info(f"Using NGC cloud ASR model {args.asr_model}")
asr_model = EncDecCTCModel.from_pretrained(model_name=args.asr_model)
asr_model.setup_test_data(
test_data_config={
'sample_rate': 16000,
'manifest_filepath': args.dataset,
'labels': asr_model.decoder.vocabulary,
'batch_size': args.batch_size,
'normalize_transcripts': args.normalize_text,
}
)
if can_gpu:
asr_model = asr_model.cuda()
asr_model.eval()
labels_map = dict([(i, asr_model.decoder.vocabulary[i]) for i in range(len(asr_model.decoder.vocabulary))])
wer = WER(vocabulary=asr_model.decoder.vocabulary)
hypotheses = []
references = []
for test_batch in asr_model.test_dataloader():
if can_gpu:
test_batch = [x.cuda() for x in test_batch]
with autocast():
log_probs, encoded_len, greedy_predictions = asr_model(
input_signal=test_batch[0], input_signal_length=test_batch[1]
)
hypotheses += wer.ctc_decoder_predictions_tensor(greedy_predictions)
for batch_ind in range(greedy_predictions.shape[0]):
reference = ''.join([labels_map[c] for c in test_batch[2][batch_ind].cpu().detach().numpy()])
references.append(reference)
del test_batch
wer_value = word_error_rate(hypotheses=hypotheses, references=references)
for h, r in zip(hypotheses, references):
print("Recognized:\t{}\nReference:\t{}\n".format(h, r))
logging.info(f'Got WER of {wer_value}. Tolerance was {args.wer_tolerance}')
return wer_value
# create the Flask app
app = Flask(__name__)
@app.route('/hole', methods=['POST'])
def query():
# arguments args
if (action == 1): #Run the ASR
# take chunks, convert back to audio
key = getKey(text)
score = Grade(key, text)
return score
elif (action == 2):
#
elif (action == 3):
#
###
#chunks = request.args['audio']
#always check for chunks file -> call Grade()
print(test)
return 'Query String Example'
if __name__ == '__main__':
# run app in debug mode on port 5000
app.run(debug=True, port=5000)
Comments
0 B
|👍
/👎