Team Ai
Apppublic

sumanthd/IndicTrans-MultilingualTranslation

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