sumanthd/IndicTrans-MultilingualTranslation
6
1# -*- coding: utf-8 -*-2# 💾⚙️🔮3 4# taken from https://github.com/Felflare/rpunct/blob/master/rpunct/punctuate.py5# modified to support batching during gpu inference6 7 8__author__ = "Daulet N."9__email__ = "daulet.nurmanbetov@gmail.com"10 11import time12import logging13import webvtt14import torch15from io import StringIO16from nltk.tokenize import sent_tokenize17#from langdetect import detect18from simpletransformers.ner import NERModel19 20 21class RestorePuncts:22 def __init__(self, wrds_per_pred=250):23 self.wrds_per_pred = wrds_per_pred24 self.overlap_wrds = 3025 self.valid_labels = ['OU', 'OO', '.O', '!O', ',O', '.U', '!U', ',U', ':O', ';O', ':U', "'O", '-O', '?O', '?U']26 self.model = NERModel("bert", "felflare/bert-restore-punctuation", labels=self.valid_labels,27 args={"silent": True, "max_seq_length": 512})28 # use_cuda isnt working and this hack seems to load the model correctly to the gpu29 self.model.device = torch.device("cuda:1")30 # dummy punctuate to load the model onto gpu31 self.punctuate("hello how are you")32 33 def punctuate(self, text: str, batch_size:int=32, lang:str=''):34 """35 Performs punctuation restoration on arbitrarily large text.36 Detects if input is not English, if non-English was detected terminates predictions.37 Overrride by supplying `lang='en'`38 39 Args:40 - text (str): Text to punctuate, can be few words to as large as you want.41 - lang (str): Explicit language of input text.42 """43 #if not lang and len(text) > 10:44 # lang = detect(text)45 #if lang != 'en':46 # raise Exception(F"""Non English text detected. Restore Punctuation works only for English.47 # If you are certain the input is English, pass argument lang='en' to this function.48 # Punctuate received: {text}""")49 50 def chunks(L, n):51 return [L[x : x + n] for x in range(0, len(L), n)]52 53 54 55 # plit up large text into bert digestable chunks56 splits = self.split_on_toks(text, self.wrds_per_pred, self.overlap_wrds)57 58 texts = [i["text"] for i in splits]59 batches = chunks(texts, batch_size)60 preds_lst = []61 62 63 for batch in batches:64 batch_preds, _ = self.model.predict(batch)65 preds_lst.extend(batch_preds)66 67 68 # predict slices69 # full_preds_lst contains tuple of labels and logits70 #full_preds_lst = [self.predict(i['text']) for i in splits]71 # extract predictions, and discard logits72 #preds_lst = [i[0][0] for i in full_preds_lst]73 # join text slices74 combined_preds = self.combine_results(text, preds_lst)75 # create punctuated prediction76 punct_text = self.punctuate_texts(combined_preds)77 return punct_text78 79 def predict(self, input_slice):80 """81 Passes the unpunctuated text to the model for punctuation.82 """83 predictions, raw_outputs = self.model.predict([input_slice])84 return predictions, raw_outputs85 86 @staticmethod87 def split_on_toks(text, length, overlap):88 """89 Splits text into predefined slices of overlapping text with indexes (offsets)90 that tie-back to original text.91 This is done to bypass 512 token limit on transformer models by sequentially92 feeding chunks of < 512 toks.93 Example output:94 [{...}, {"text": "...", 'start_idx': 31354, 'end_idx': 32648}, {...}]95 """96 wrds = text.replace('\n', ' ').split(" ")97 resp = []98 lst_chunk_idx = 099 i = 0100 101 while True:102 # words in the chunk and the overlapping portion103 wrds_len = wrds[(length * i):(length * (i + 1))]104 wrds_ovlp = wrds[(length * (i + 1)):((length * (i + 1)) + overlap)]105 wrds_split = wrds_len + wrds_ovlp106 107 # Break loop if no more words108 if not wrds_split:109 break110 111 wrds_str = " ".join(wrds_split)112 nxt_chunk_start_idx = len(" ".join(wrds_len))113 lst_char_idx = len(" ".join(wrds_split))114 115 resp_obj = {116 "text": wrds_str,117 "start_idx": lst_chunk_idx,118 "end_idx": lst_char_idx + lst_chunk_idx,119 }120 121 resp.append(resp_obj)122 lst_chunk_idx += nxt_chunk_start_idx + 1123 i += 1124 logging.info(f"Sliced transcript into {len(resp)} slices.")125 return resp126 127 @staticmethod128 def combine_results(full_text: str, text_slices):129 """130 Given a full text and predictions of each slice combines predictions into a single text again.131 Performs validataion wether text was combined correctly132 """133 split_full_text = full_text.replace('\n', ' ').split(" ")134 split_full_text = [i for i in split_full_text if i]135 split_full_text_len = len(split_full_text)136 output_text = []137 index = 0138 139 if len(text_slices[-1]) <= 3 and len(text_slices) > 1:140 text_slices = text_slices[:-1]141 142 for _slice in text_slices:143 slice_wrds = len(_slice)144 for ix, wrd in enumerate(_slice):145 # print(index, "|", str(list(wrd.keys())[0]), "|", split_full_text[index])146 if index == split_full_text_len:147 break148 149 if split_full_text[index] == str(list(wrd.keys())[0]) and \150 ix <= slice_wrds - 3 and text_slices[-1] != _slice:151 index += 1152 pred_item_tuple = list(wrd.items())[0]153 output_text.append(pred_item_tuple)154 elif split_full_text[index] == str(list(wrd.keys())[0]) and text_slices[-1] == _slice:155 index += 1156 pred_item_tuple = list(wrd.items())[0]157 output_text.append(pred_item_tuple)158 assert [i[0] for i in output_text] == split_full_text159 return output_text160 161 @staticmethod162 def punctuate_texts(full_pred: list):163 """164 Given a list of Predictions from the model, applies the predictions to text,165 thus punctuating it.166 """167 punct_resp = ""168 for i in full_pred:169 word, label = i170 if label[-1] == "U":171 punct_wrd = word.capitalize()172 else:173 punct_wrd = word174 175 if label[0] != "O":176 punct_wrd += label[0]177 178 punct_resp += punct_wrd + " "179 punct_resp = punct_resp.strip()180 # Append trailing period if doesnt exist.181 if punct_resp[-1].isalnum():182 punct_resp += "."183 return punct_resp184 185 186if __name__ == "__main__":187 188 start = time.time()189 punct_model = RestorePuncts()190 191 load_model = time.time()192 print(f'Time to load model: {load_model - start}')193 # read test file194 # with open('en_lower.txt', 'r') as fp:195 # # test_sample = fp.read()196 # lines = fp.readlines()197 198 with open('sample.vtt', 'r') as fp:199 source_text = fp.read()200 201 # captions = webvtt.read_buffer(StringIO(source_text))202 captions = webvtt.read('sample.vtt')203 source_sentences = [caption.text.replace('\r', '').replace('\n', ' ') for caption in captions]204 205 # print(source_sentences)206 207 sent = ' '.join(source_sentences)208 punctuated = punct_model.punctuate(sent)209 210 tokenised = sent_tokenize(punctuated)211 # print(tokenised)212 213 for i in range(len(tokenised)):214 captions[i].text = tokenised[i]215 # return captions.content216 captions.save('my_captions.vtt')217 218 end = time.time()219 print(f'Time for run: {end - load_model}')220 print(f'Total time: {end - start}')221 