sumanthd/IndicTrans-MultilingualTranslation
6
1import time2 3import re4from math import floor, ceil5from fairseq import checkpoint_utils, distributed_utils, options, tasks, utils6# from nltk.tokenize import sent_tokenize7from flask import Flask, request, jsonify8from flask_cors import CORS, cross_origin9import webvtt10from io import StringIO11from mosestokenizer import MosesSentenceSplitter12 13from indicTrans.inference.engine import Model14from punctuate import RestorePuncts15from indicnlp.tokenize.sentence_tokenize import sentence_split16 17app = Flask(__name__)18cors = CORS(app)19app.config['CORS_HEADERS'] = 'Content-Type'20 21indic2en_model = Model(expdir='models/v3/indic-en')22en2indic_model = Model(expdir='models/v3/en-indic')23m2m_model = Model(expdir='models/m2m')24 25rpunct = RestorePuncts()26 27indic_language_dict = {28 'Assamese': 'as',29 'Hindi' : 'hi',30 'Marathi' : 'mr',31 'Tamil' : 'ta',32 'Bengali' : 'bn',33 'Kannada' : 'kn',34 'Oriya' : 'or',35 'Telugu' : 'te',36 'Gujarati' : 'gu',37 'Malayalam' : 'ml',38 'Punjabi' : 'pa',39}40 41splitter = MosesSentenceSplitter('en')42 43def get_inference_params():44 source_language = request.form['source_language']45 target_language = request.form['target_language']46 47 if source_language in indic_language_dict and target_language == 'English':48 model = indic2en_model49 source_lang = indic_language_dict[source_language]50 target_lang = 'en'51 elif source_language == 'English' and target_language in indic_language_dict:52 model = en2indic_model53 source_lang = 'en'54 target_lang = indic_language_dict[target_language]55 elif source_language in indic_language_dict and target_language in indic_language_dict:56 model = m2m_model57 source_lang = indic_language_dict[source_language]58 target_lang = indic_language_dict[target_language]59 60 return model, source_lang, target_lang61 62@app.route('/', methods=['GET'])63def main():64 return "IndicTrans API"65 66@app.route('/supported_languages', methods=['GET'])67@cross_origin()68def supported_languages():69 return jsonify(indic_language_dict)70 71@app.route("/translate", methods=['POST'])72@cross_origin()73def infer_indic_en():74 model, source_lang, target_lang = get_inference_params()75 source_text = request.form['text']76 77 start_time = time.time()78 target_text = model.translate_paragraph(source_text, source_lang, target_lang)79 end_time = time.time()80 return {'text':target_text, 'duration':round(end_time-start_time, 2)}81 82@app.route("/translate_vtt", methods=['POST'])83@cross_origin()84def infer_vtt_indic_en():85 start_time = time.time()86 model, source_lang, target_lang = get_inference_params()87 source_text = request.form['text']88 # vad_segments = request.form['vad_nochunk'] # Assuming it is an array of start & end timestamps89 90 vad = webvtt.read_buffer(StringIO(source_text))91 source_sentences = [v.text.replace('\r', '').replace('\n', ' ') for v in vad]92 93 ## SUMANTH LOGIC HERE ##94 95 # for each vad timestamp, do:96 large_sentence = ' '.join(source_sentences) # only sentences in that time range97 large_sentence = large_sentence.lower()98 # split_sents = sentence_split(large_sentence, 'en')99 # print(split_sents)100 101 large_sentence = re.sub(r'[^\w\s]', '', large_sentence)102 punctuated = rpunct.punctuate(large_sentence, batch_size=32)103 end_time = time.time()104 print("Time Taken for punctuation: {} s".format(end_time - start_time))105 start_time = time.time()106 split_sents = splitter([punctuated]) ### Please uncomment107 108 109 # print(split_sents)110 # output_sentence_punctuated = model.translate_paragraph(punctuated, source_lang, target_lang)111 output_sents = model.batch_translate(split_sents, source_lang, target_lang)112 # print(output_sents)113 # output_sents = split_sents114 # print(output_sents)115 # align this to those range of source_sentences in `captions`116 117 map_ = {split_sents[i] : output_sents[i] for i in range(len(split_sents))}118 # print(map_)119 punct_para = ' '.join(list(map_.keys()))120 nmt_para = ' '.join(list(map_.values()))121 nmt_words = nmt_para.split(' ')122 123 len_punct = len(punct_para.split(' '))124 len_nmt = len(nmt_para.split(' '))125 126 start = 0127 for i in range(len(vad)):128 if vad[i].text == '':129 continue130 131 len_caption = len(vad[i].text.split(' '))132 frac = (len_caption / len_punct)133 # frac = round(frac, 2)134 135 req_nmt_size = floor(frac * len_nmt)136 # print(frac, req_nmt_size)137 138 vad[i].text = ' '.join(nmt_words[start:start+req_nmt_size])139 # print(vad[i].text)140 # print(start, req_nmt_size)141 start += req_nmt_size142 143 end_time = time.time()144 145 print("Time Taken for translation: {} s".format(end_time - start_time))146 147 # vad.save('aligned.vtt')148 149 return {150 'text': vad.content,151 # 'duration':round(end_time-start_time, 2)152 }153 