Team Ai
Apppublic

sumanthd/IndicTrans-MultilingualTranslation

sourceHugging Facemitupdated 4y agoView on Hugging Face
6likes
api.py153 linesDownload Raw Back to api
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