{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Import","metadata":{}},{"cell_type":"code","source":"!cp -r ../input/python-packages2 ./\n!cp -r ../input/beng-pip-packages ./\n!cp -r ../input/sctk-rover/SCTK ./\n\n!tar xvfz ./beng-pip-packages/indic-nlp.tgz\n!pip install ./indicnlp/indic_nlp_library-0.81-py3-none-any.whl -f ./ --no-index --find-links ./indicnlp/\n\n!tar xvfz ./python-packages2/jiwer.tgz\n!pip install ./jiwer/jiwer-2.3.0-py3-none-any.whl -f ./ --no-index\n!tar xvfz ./python-packages2/normalizer.tgz\n!pip install ./normalizer/bnunicodenormalizer-0.0.24.tar.gz -f ./ --no-index\n!tar xvfz ./python-packages2/pyctcdecode.tgz\n!pip install ./pyctcdecode/attrs-22.1.0-py2.py3-none-any.whl -f ./ --no-index --no-deps\n!pip install ./pyctcdecode/exceptiongroup-1.0.0rc9-py3-none-any.whl -f ./ --no-index --no-deps\n!pip install ./pyctcdecode/hypothesis-6.54.4-py3-none-any.whl -f ./ --no-index --no-deps\n!pip install ./pyctcdecode/numpy-1.21.6-cp37-cp37m-manylinux_2_12_x86_64.manylinux2010_x86_64.whl -f ./ --no-index --no-deps\n!pip install ./pyctcdecode/pygtrie-2.5.0.tar.gz -f ./ --no-index --no-deps\n!pip install ./pyctcdecode/sortedcontainers-2.4.0-py2.py3-none-any.whl -f ./ --no-index --no-deps\n!pip install ./pyctcdecode/pyctcdecode-0.4.0-py2.py3-none-any.whl -f ./ --no-index --no-deps\n\n!tar xvfz ./python-packages2/pypikenlm.tgz\n!pip install ./pypikenlm/pypi-kenlm-0.1.20220713.tar.gz -f ./ --no-index --no-deps\n\n!pip install datasets --no-index --find-links=file:///kaggle/input/hf-ds -U -q","metadata":{"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rm -r python-packages2 jiwer normalizer pyctcdecode pypikenlm beng-pip-packages indicnlp","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport sys\nimport shutil\n\nimport typing as tp\nfrom pathlib import Path\nimport functools\nimport ctypes\nfrom functools import partial\nfrom dataclasses import dataclass\n\nimport pandas as pd\nimport numpy as np\nimport pickle\nfrom tqdm.notebook import tqdm\n\nimport librosa\n\nimport pyctcdecode\nimport kenlm\nimport json\nimport torch\nimport gc\nimport re\nimport abc\nimport math\nimport subprocess\n\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.nn import CrossEntropyLoss\nfrom torch.utils.data.dataloader import DataLoader\nimport torchaudio\nfrom transformers import Wav2Vec2Processor, Wav2Vec2ProcessorWithLM, Wav2Vec2ForCTC\nfrom transformers import AutoTokenizer, AutoModelForTokenClassification\nfrom transformers import AutoConfig, AutoModel, PretrainedConfig\nfrom transformers.modeling_outputs import TokenClassifierOutput\nfrom transformers import AutomaticSpeechRecognitionPipeline\nfrom transformers import DataCollatorForTokenClassification\nfrom transformers.pipelines.automatic_speech_recognition import *\nfrom transformers.pipelines.pt_utils import KeyDataset\nfrom transformers.tokenization_utils import PreTrainedTokenizer\nfrom transformers.modelcard import ModelCard\nfrom transformers.pipelines.base import ArgumentHandler, ChunkPipeline, infer_framework_load_model\nfrom transformers.models.auto.modeling_auto import MODEL_FOR_CTC_MAPPING_NAMES, MODEL_FOR_SPEECH_SEQ_2_SEQ_MAPPING_NAMES\n\nfrom indicnlp.tokenize import indic_tokenize\nfrom bnunicodenormalizer import Normalizer\nfrom collections import defaultdict, Counter\n\nimport datasets\nfrom datasets import Dataset\nfrom datasets import Audio\nimport logging\n\nlogger = logging.getLogger(__name__)\nfrom typing import (\n    TYPE_CHECKING,\n    Any,\n    Collection,\n    Dict,\n    Iterable,\n    List,\n    Optional,\n    Sequence,\n    Tuple,\n    TypeVar,\n    Union,\n)\n\nimport multiprocessing ","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROOT = Path.cwd().parent\nINPUT = ROOT / \"input\"\nDATA = INPUT / \"bengaliai-speech\"\nTRAIN = DATA / \"train_mp3s\"\nTEST = DATA / \"test_mp3s\"\nensemble = False\n\nSAMPLING_RATE = 16_000\nMODEL_PATH = INPUT / \"wv-shru-v3-s6/\"\nDEFAULT_TRANSCRIPTION = \"ও\" \nkenlm_model_path = \"/kaggle/input/lm-trained/lm_v9.binary\"\nALPHA = 0.5\nBETA = 1.0\n\ndecoder_kwargs={'beam_width':1024}","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load audio paths","metadata":{"execution":{"iopub.status.busy":"2023-10-20T17:30:07.139288Z","iopub.execute_input":"2023-10-20T17:30:07.140103Z","iopub.status.idle":"2023-10-20T17:30:07.144707Z","shell.execute_reply.started":"2023-10-20T17:30:07.140066Z","shell.execute_reply":"2023-10-20T17:30:07.143679Z"}}},{"cell_type":"code","source":"test = pd.read_csv(DATA / \"sample_submission.csv\", dtype={\"id\": str})\ntest_audio_paths = [str(TEST / f\"{aid}.mp3\") for aid in test[\"id\"].values]\nprint(test.head())","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## ASR model + Language model","metadata":{}},{"cell_type":"code","source":"processor = Wav2Vec2Processor.from_pretrained(MODEL_PATH)\nvocab_dict = processor.tokenizer.get_vocab()\nsorted_vocab_dict = {k: v for k, v in sorted(vocab_dict.items(), key=lambda item: item[1])}\ntoken_list = list(sorted_vocab_dict.keys())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open(\"/kaggle/input/lm-trained/vocab-4500000_v9.txt\") as f:\n    unigram_list = [t for t in f.read().strip().split(\"\\n\")]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"decoder = pyctcdecode.build_ctcdecoder(\n    token_list,\n    kenlm_model_path,\n    unigram_list,\n    alpha = ALPHA,\n    beta = BETA,\n)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not torch.cuda.is_available():\n    device = torch.device(\"cpu\")\nelse:\n    device = torch.device(\"cuda\")\nprint(device)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Function to clean RAM & vRAM\ndef clean_memory():\n    gc.collect()\n    ctypes.CDLL(\"libc.so.6\").malloc_trim(0)\n    torch.cuda.empty_cache()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_sentence_list = []\nbnorm = Normalizer()\n\ndef normalize_sentence(sentence):\n    if len(sentence)==0:\n        return DEFAULT_TRANSCRIPTION\n    _words = [bnorm(word)['normalized'] for word in sentence.split()]\n    sentence = \" \".join([word for word in _words if word is not None])    \n    return sentence","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Wav2Vec2ForCTC.from_pretrained(MODEL_PATH)\nprocessor = Wav2Vec2ProcessorWithLM(\n    feature_extractor=processor.feature_extractor,\n    tokenizer=processor.tokenizer,\n    decoder=decoder\n)\npipe = AutomaticSpeechRecognitionPipeline(\n    model=model,\n    tokenizer=processor.tokenizer,\n    feature_extractor=processor.feature_extractor,\n    decoder=processor.decoder,\n    device=0,\n    decoder_kwargs=decoder_kwargs\n)\npipe.decoder = processor.decoder\npipe.type = \"ctc_with_lm\"\n\nwith torch.no_grad():\n    for audio_path in tqdm(test_audio_paths):\n        w = librosa.load(audio_path, sr=16_000, mono=False)[0]\n        w = np.trim_zeros(w, 'fb')\n\n        text = pipe(w, chunk_length_s=14, stride_length_s=(6, 3))[\"text\"]\n        pred_sentence_list.append(text)\n\ndel pipe, processor, model, decoder\nclean_memory()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Punctuation model","metadata":{}},{"cell_type":"code","source":"#Credits to Feedback 1 4th place solution: https://www.kaggle.com/competitions/feedback-prize-2021/discussion/313330\ndef create_ner_conditional_masks(id2label: Dict[int, str]) -> torch.Tensor:\n    \"\"\"Create a NER-conditional mask matrix which implies the relations between\n    before-tag and after-tag.\n\n    According to the rule of BIO-naming system, it is impossible that `I-Dog` cannot be\n    appeard after `B-Dog` or `I-Dog` tags. This function creates the calculable\n    relation-based conditional matrix to prevent from generating wrong tags.\n\n    Args:\n        id2label: A dictionary which maps class indices to their label names.\n\n    Returns:\n        A conditional mask tensor.\n    \"\"\"\n    conditional_masks = torch.zeros(len(id2label), len(id2label))\n    for i, before in id2label.items():\n        for j, after in id2label.items():\n            if after == \"O\" or after.startswith(\"B-\") or after == f\"I-{before[2:]}\":\n                conditional_masks[i, j] = 1.0\n    return conditional_masks\n\ndef ner_beam_search_decode(\n    log_probs: torch.Tensor, id2label: Dict[int, str], beam_size: int = 2\n) -> Tuple[torch.Tensor, torch.Tensor]:\n    \"\"\"Decode NER-tags from the predicted log-probabilities using beam-search.\n\n    This function decodes the predictions using beam-search algorithm. Because all tags\n    are predicted simultaneously while the tags have dependencies of their previous\n    tags, the greedy algorithm cannot decode the tags properly. With beam-search, it is\n    possible to prevent the below situation:\n\n        >>> sorted = probs[t].sort(dim=-1)\n        >>> print(\"\\t\".join([f\"{id2label[i]} {p}\" for p, i in zip()]))\n        I-Dog 0.54  B-Cat 0.44  ...\n        >>> sorted = probs[t + 1].sort(dim=-1)\n        >>> print(\"\\t\".join([f\"{id2label[i]} {p}\" for p, i in zip()]))\n        I-Cat 0.99  I-Dog 0.01  ...\n\n    The above shows that if the locally-highest tags are selected, then `I-Dog, I-Dog`\n    will be generated even the confidence of the second tag `I-Dog` is significantly\n    lower than `I-Cat`. It is more natural that `B-Cat, I-Cat` is generated rather than\n    `I-Dog, I-Dog`. The beam-search for NER-tagging task can solve this problem.\n\n    Args:\n        log_probs: The log-probabilities of the token predictions.\n        id2label: A dictionary which maps class indices to their label names.\n        beam_size: The number of candidates for each search step. Default is `2`.\n\n    Returns:\n        A tuple of beam-searched indices and their probability tensors.\n    \"\"\"\n    # Create the log-probability mask for the invalid predictions.\n    log_prob_masks = -10000.0 * (1 - create_ner_conditional_masks(id2label))\n    log_prob_masks = log_prob_masks.to(log_probs.device)\n\n    beam_search_shape = (log_probs.size(0), beam_size, log_probs.size(1))\n    searched_tokens = log_probs.new_zeros(beam_search_shape, dtype=torch.long)\n    searched_log_probs = log_probs.new_zeros(beam_search_shape)\n\n    searched_scores = log_probs.new_zeros(log_probs.size(0), beam_size)\n    searched_scores[:, 1:] = -10000.0\n\n    for i in range(log_probs.size(1)):\n        # Calculate the accumulated score (log-probabilities) with excluding invalid\n        # next-tag predictions.\n        scores = searched_scores.unsqueeze(2)\n        scores = scores + log_probs[:, i, :].unsqueeze(1)\n        scores = scores + (log_prob_masks[searched_tokens[:, :, i - 1]] if i > 0 else 0)\n\n        # Select the top-k (beam-search size) predictions.\n        best_scores, best_indices = scores.flatten(1).topk(beam_size)\n        best_tokens = best_indices % scores.size(2)\n        best_log_probs = log_probs[:, i, :].gather(dim=1, index=best_tokens)\n\n        best_buckets = best_indices.div(scores.size(2), rounding_mode=\"floor\")\n        best_buckets = best_buckets.unsqueeze(2).expand(-1, -1, log_probs.size(1))\n\n        # Gather the best buckets and their log-probabilities.\n        searched_tokens = searched_tokens.gather(dim=1, index=best_buckets)\n        searched_log_probs = searched_log_probs.gather(dim=1, index=best_buckets)\n\n        # Update the predictions by inserting to the corresponding timestep.\n        searched_scores = best_scores\n        searched_tokens[:, :, i] = best_tokens\n        searched_log_probs[:, :, i] = best_log_probs\n\n    # Return the best beam-searched sequence and its probabilities.\n    return searched_tokens[:, 0, :], searched_log_probs[:, 0, :].exp()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PunctuationDataset(torch.utils.data.Dataset):\n    def __init__(self, texts, tokenizer):\n        self.texts = texts\n        self.tokenizer = tokenizer\n\n    def __len__(self):\n        return len(self.texts)\n\n    def __getitem__(self, item):\n        sentence = self.texts[item].split()\n\n        tokenized_inputs = self.tokenizer(\n            sentence, \n            truncation=True, \n            max_length=512, \n            padding='max_length', \n            is_split_into_words=True\n        )\n        word_ids = tokenized_inputs.word_ids()  \n\n        tokenized_inputs = {key: torch.as_tensor(val) for key, val in tokenized_inputs.items()}\n\n        word_ids2 = [w if w is not None else -1 for w in word_ids]\n        tokenized_inputs['wids'] = torch.as_tensor(word_ids2)\n        return tokenized_inputs\n\ntag_values = ['O', \n        'B-END', \n        'I-END', \n        'B-COMMA', \n        'I-COMMA', \n        'B-QM', \n        'I-QM', \n        'B-EXCLM',\n        'I-EXCLM',\n        ]\nlabels_to_ids = {v:int(k) for k,v in enumerate(tag_values)}\nids_to_labels = {k:v for k,v in enumerate(tag_values)}\n\npunctuation_dict = {\n                            \"O\": \" \", \n                            \"B-QM\": \"? \",\n                            \"I-QM\": \"? \",\n                            \"B-COMMA\": \", \", \n                            \"I-COMMA\": \", \", \n                            \"B-END\": \"। \",\n                            \"I-END\": \"। \",\n                            \"B-EXCLM\": \"! \", \n                            \"I-EXCLM\": \"! \", \n                            \"B-HYPHEN\":\"-\",\n                            \"I-HYPHEN\":\"-\",\n                            \"B-DOT\":\".\",\n                            \"I-DOT\":\".\",\n                        }\n\n# LSTM Module credits to Feedback 1 3rd place solution: https://www.kaggle.com/competitions/feedback-prize-2021/discussion/313235\nclass ResidualLSTM(nn.Module):\n\n    def __init__(self, d_model,rnn):\n        super(ResidualLSTM, self).__init__()\n        self.downsample=nn.Linear(d_model,d_model//2)\n        if rnn=='GRU':\n            self.LSTM=nn.GRU(d_model//2, d_model//2, num_layers=2, bidirectional=False, dropout=0.2)\n        else:\n            self.LSTM=nn.LSTM(d_model//2, d_model//2, num_layers=2, bidirectional=False, dropout=0.2)\n        self.dropout1=nn.Dropout(0.2)\n        self.norm1= nn.LayerNorm(d_model//2)\n        self.linear1=nn.Linear(d_model//2, d_model*4)\n        self.linear2=nn.Linear(d_model*4, d_model)\n        self.dropout2=nn.Dropout(0.2)\n        self.norm2= nn.LayerNorm(d_model)\n\n    def forward(self, x):\n        res=x\n        x=self.downsample(x)\n        x, _ = self.LSTM(x)\n        x=self.dropout1(x)\n        x=self.norm1(x)\n        x=F.relu(self.linear1(x))\n        x=self.linear2(x)\n        x=self.dropout2(x)\n        x=res+x\n        return self.norm2(x)\n\nclass CustomModel(nn.Module):\n    def __init__(self, config_path, num_labels):\n        super(CustomModel, self).__init__()\n        self.config = AutoConfig.from_pretrained(config_path)\n        self.num_labels = num_labels\n\n        self.model = AutoModel.from_pretrained(\n            config_path, add_pooling_layer=False\n        )\n\n        self.lstm = ResidualLSTM(self.config.hidden_size,'LSTM')\n        self.classification_head = nn.Linear(self.config.hidden_size,self.num_labels)\n\n    def _init_weights(self, module):\n        if isinstance(module, nn.Linear):\n            module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)\n            if module.bias is not None:\n                module.bias.data.zero_()\n        elif isinstance(module, nn.Embedding):\n            module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)\n            if module.padding_idx is not None:\n                module.weight.data[module.padding_idx].zero_()\n        elif isinstance(module, nn.LayerNorm):\n            module.bias.data.zero_()\n            module.weight.data.fill_(1.0)\n\n    def forward(\n        self,\n        input_ids: Optional[torch.Tensor] = None,\n        attention_mask: Optional[torch.Tensor] = None,\n        token_type_ids: Optional[torch.Tensor] = None,\n        position_ids: Optional[torch.Tensor] = None,\n        head_mask: Optional[torch.Tensor] = None,\n        inputs_embeds: Optional[torch.Tensor] = None,\n        labels: Optional[torch.Tensor] = None,\n        output_attentions: Optional[bool] = None,\n        output_hidden_states: Optional[bool] = None,\n        return_dict: Optional[bool] = None,\n    ) -> Union[Tuple[torch.Tensor], TokenClassifierOutput]:\n        r\"\"\"\n        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):\n            Labels for computing the token classification loss. Indices should be in `[0, ..., config.num_labels - 1]`.\n        \"\"\"\n        return_dict = return_dict if return_dict is not None else self.config.use_return_dict\n\n        outputs = self.model(\n            input_ids,\n            attention_mask=attention_mask,\n            token_type_ids=token_type_ids,\n            position_ids=position_ids,\n            head_mask=head_mask,\n            inputs_embeds=inputs_embeds,\n            output_attentions=output_attentions,\n            output_hidden_states=output_hidden_states,\n            return_dict=return_dict,\n        )\n\n        sequence_output = outputs[0]\n\n        x = self.lstm(sequence_output.permute(1,0,2)).permute(1,0,2)\n        logits = self.classification_head(x)\n\n        loss = None\n        if labels is not None:\n            loss_fct = CrossEntropyLoss()\n            loss = loss_fct(logits.view(-1, self.num_labels), labels.view(-1))\n\n        if not return_dict:\n            output = (logits,) + outputs[2:]\n            return ((loss,) + output) if loss is not None else output\n\n        return TokenClassifierOutput(\n            loss=loss,\n            logits=logits,\n            hidden_states=outputs.hidden_states,\n            attentions=outputs.attentions,\n        )","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def postprocess(sentence):\n    period_set = set([\"?\", \"!\", \"।\"])\n    _words = [bnorm(word)['normalized'] for word in sentence.split()]\n    sentence = \" \".join([word for word in _words if word is not None])\n    \n    try:\n        if sentence[-1] not in period_set:\n            sentence+=\"।\"\n    except:\n        sentence = DEFAULT_TRANSCRIPTION\n    return sentence","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pp_pred_sentence_list = [\n    normalize_sentence(s) for s in tqdm(pred_sentence_list)]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pretrained_path = '/kaggle/input/indic-bert-v2'\n\ntokenizer = AutoTokenizer.from_pretrained(pretrained_path)\nmodel = CustomModel(pretrained_path, len(tag_values))\n\nhyphen_caches = []\ndot_caches = []\npp_pred_sentence_list_new = []\nfor sentence in pp_pred_sentence_list:\n    hyphen_cache = []\n    dot_cache = []\n    sentence = sentence.replace(\"-\", \"- \").replace(\".\", \". \")\n    sentence = sentence.split(\" \")\n    sentence = [w for w in sentence if len(w)>0]\n    new_sentence = []\n    for idx, s in enumerate(sentence):\n        if s[-1]==\"-\":\n            hyphen_cache.append(idx)\n            new_sentence.append(s[:-1])\n        elif s[-1]==\".\":\n            dot_cache.append(idx)\n            new_sentence.append(s[:-1])\n        else:\n            new_sentence.append(s)\n    hyphen_caches.append(hyphen_cache)\n    dot_caches.append(dot_cache)\n    pp_pred_sentence_list_new.append(' '.join(new_sentence))\n\ntest_dataset = PunctuationDataset(pp_pred_sentence_list_new, tokenizer)\n\nbatch_size = 32\ntest_params = {'batch_size': batch_size,\n                'shuffle': False,\n                'num_workers': 2,\n                'pin_memory':True\n                }\n\ntest_dataloader = DataLoader(test_dataset, **test_params)\n\nensemble_preds = np.zeros((len(test_dataloader.dataset), 512, len(tag_values)), dtype=np.float32)\nwids = np.full((len(test_dataloader.dataset), 512), -1)\nall_tokens = []\n\nmodel_ids = [\n    \"/kaggle/input/punct-3seed-ner/train_demo_new_v2plus_s200_ner\",\n    \"/kaggle/input/punct-3seed-ner/train_demo_new_v2plus_s300_ner\",\n    \"/kaggle/input/punct-3seed-ner/train_demo_new_v2plus_s400_ner\",\n]\n\nfor model_i, model_id in enumerate(model_ids):\n\n    checkpoint = torch.load(f\"{model_id}/pytorch_model.bin\", map_location=device)\n    model.load_state_dict(checkpoint)\n    model = model.to(device)\n    model.eval()\n\n    for batch_i, batch in tqdm(enumerate(test_dataloader)):\n        # print(batch.keys())\n        if model_i == 0: \n            wids[batch_i*batch_size:(batch_i+1)*batch_size,:batch['wids'].shape[1]] = batch['wids'].numpy()\n            inputs = batch['input_ids'].numpy()\n            for i in range(inputs.shape[0]):\n                tokens_temp = tokenizer.convert_ids_to_tokens(inputs[i])\n                tokens_temp = [t for t in tokens_temp if t !=\"[PAD]\"]\n                all_tokens.append(tokens_temp)\n        # MOVE BATCH TO GPU AND INFER\n        ids = batch[\"input_ids\"].to(device)\n        mask = batch[\"attention_mask\"].to(device)\n\n        with torch.no_grad():\n            #with amp.autocast():\n            outputs = model(ids, attention_mask=mask)\n        all_preds = torch.nn.functional.softmax(outputs[0], dim=2).cpu().detach().numpy() \n        ensemble_preds[batch_i*batch_size:(batch_i+1)*batch_size,:all_preds.shape[1]] += all_preds\n        \n        # all_preds = torch.argmax(outputs[0], axis=-1).cpu().numpy() \n        del ids\n        del mask\n        del outputs\n        del all_preds\n        \n    clean_memory()\n\nensemble_preds /= len(model_ids)\nensemble_preds = torch.from_numpy(ensemble_preds).log()\npreds, pred_probs = ner_beam_search_decode(ensemble_preds, ids_to_labels, 4)\nensemble_preds = preds.cpu().numpy()\n\npredictions = []\nfor text_i in range(ensemble_preds.shape[0]):\n    prediction = []\n\n    label_indices = ensemble_preds[text_i]\n    word_ids = wids[text_i]\n    tokens = all_tokens[text_i]\n    new_tokens = []\n    new_labels = []\n\n    for i in range(1, len(tokens) - 1):\n        if word_ids[i]!=word_ids[i-1]:\n            if tokens[i].startswith(\"▁\"):\n                current_word = tokens[i][1:]\n            else:\n                current_word = tokens[i]\n\n            new_labels.append(list(labels_to_ids.keys())[list(labels_to_ids.values()).index(label_indices[i])])\n            for j in range(i + 1, len(tokens) - 1):\n                if word_ids[j]==word_ids[j-1]:\n                    current_word = current_word + tokens[j]\n                if word_ids[j]!=word_ids[j-1]:\n                    break\n            new_tokens.append(current_word)\n    full_text = ''\n    tokenized_text = indic_tokenize.trivial_tokenize_indic(pp_pred_sentence_list_new[text_i])\n        \n    new_labels = ['blank' if x=='PAD' else x for x in new_labels] #fix for PAD predicted in outputs\n\n    if len(tokenized_text) == len(new_labels):\n        full_text_tokens = tokenized_text\n    else:\n        full_text_tokens = new_tokens\n\n    hyphen_cache = hyphen_caches[text_i]\n    dot_cache = dot_caches[text_i]\n    for word_idx, (word, punctuation) in enumerate(zip(full_text_tokens, new_labels)):\n        if punctuation==\"O\":\n            if word_idx in hyphen_cache:\n                punctuation = \"B-HYPHEN\"\n            elif word_idx in dot_cache:\n                punctuation = \"B-DOT\"\n        full_text = full_text + word + punctuation_dict[punctuation]\n    predictions.append(full_text.strip())\n    \ndel model, ensemble_preds, test_dataloader, test_dataset, tokenizer, hyphen_caches, dot_caches\nclean_memory()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pp_pred_sentence_list = [\n    postprocess(s) for s in tqdm(predictions)]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test[\"sentence\"] = pp_pred_sentence_list\ntest.to_csv(\"submission.csv\", index=False)\nprint(test.head())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## EOF","metadata":{}}]}