{"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":"code","source":"!cp -r ../input/python-packages2 ./\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! python -m pip install --no-index --find-links=../input/install-notebook -r ../input/install-notebook/requirements.txt\n\n!pip install /kaggle/input/transformers/transformers-4.33.1-py3-none-any.whl","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-10-20T07:02:37.978529Z","iopub.execute_input":"2023-10-20T07:02:37.978898Z","iopub.status.idle":"2023-10-20T07:06:05.831276Z","shell.execute_reply.started":"2023-10-20T07:02:37.978849Z","shell.execute_reply":"2023-10-20T07:06:05.829825Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Main libraries & misc\nimport resource\nimport os\nimport random\nimport re\nimport argparse\nimport gc\nimport pickle\nfrom tqdm import tqdm\n\n#Numeric\nimport pandas as pd\nimport numpy as np\n\n#Challenge specific\nfrom bnunicodenormalizer import Normalizer\nimport librosa\nimport pyctcdecode\nfrom pyctcdecode import BeamSearchDecoderCTC\n\n#Deep Learning\nimport torch\nimport torch.nn as nn\nimport pytorch_lightning\nfrom transformers import pipeline\nfrom transformers import Wav2Vec2CTCTokenizer, Wav2Vec2FeatureExtractor, Wav2Vec2Processor, Wav2Vec2Model, Wav2Vec2Config, Wav2Vec2ForCTC\nfrom transformers import XLMRobertaModel, XLMRobertaTokenizer\nfrom torch.nn import TransformerEncoder, TransformerEncoderLayer\n\nrandom.seed(42)","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:06:05.833952Z","iopub.execute_input":"2023-10-20T07:06:05.834340Z","iopub.status.idle":"2023-10-20T07:06:19.557606Z","shell.execute_reply.started":"2023-10-20T07:06:05.834304Z","shell.execute_reply":"2023-10-20T07:06:19.556942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"#Competition config\nclass CCFG:\n    #Data\n    train_data = \"/kaggle/input/bengaliai-speech/train_mp3s\"\n    test_data = \"/kaggle/input/bengaliai-speech/test_mp3s\"\n    \n    #Punctuation model\n    punc_base = \"/kaggle/input/xlm-roberta-large/xlm-roberta-large\"\n    punc_tokenizer = \"/kaggle/input/xlm-roberta-large/xlm-roberta-large\"\n    punc_weights = \"/kaggle/input/punct-correct-roberta-bn/xlm-roberta-large-bn.pt\"\n    \n    #ASR processor & models\n    processor = \"/kaggle/input/v20-processor/best_v20_processor\"\n    small_model = \"/kaggle/input/v35-210ksteps\"\n    large_model = \"/kaggle/input/v32-130k\"\n    ensemble_model = \"/kaggle/input/ensemble-v1/pytorch_model.bin\"\n    \n    #Decoder\n    decoder = \"/kaggle/input/llm-pruned-00011/new_model_bin_mixed\"\n    \n    #Neural rescoring [not used in the final submission due memory restrictions]\n    neural_rescoring = \"/kaggle/input/neural-rescoring\"","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:06:19.558706Z","iopub.execute_input":"2023-10-20T07:06:19.559040Z","iopub.status.idle":"2023-10-20T07:06:19.564601Z","shell.execute_reply.started":"2023-10-20T07:06:19.559010Z","shell.execute_reply":"2023-10-20T07:06:19.563694Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda:0')","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:06:19.566257Z","iopub.execute_input":"2023-10-20T07:06:19.566977Z","iopub.status.idle":"2023-10-20T07:06:19.579337Z","shell.execute_reply.started":"2023-10-20T07:06:19.566954Z","shell.execute_reply":"2023-10-20T07:06:19.578433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Functions and params","metadata":{}},{"cell_type":"code","source":"bnorm = Normalizer()\n\ndef normalize(sen):\n    \"\"\"\n    Normalize a sentence by applying the 'bnorm' Normalizer to each word in the sentence.\n\n    Args:\n        sen (str): The input sentence to be normalized.\n\n    Returns:\n        str: The normalized sentence where each word has been normalized using 'bnorm'.\n    \"\"\"\n    _words = [bnorm(word)['normalized'] for word in sen.split()]\n    return \" \".join([word for word in _words if word is not None])","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:06:19.580530Z","iopub.execute_input":"2023-10-20T07:06:19.581228Z","iopub.status.idle":"2023-10-20T07:06:19.588760Z","shell.execute_reply.started":"2023-10-20T07:06:19.581203Z","shell.execute_reply":"2023-10-20T07:06:19.587875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TOKEN_IDX = {\n    'bert': {\n        'START_SEQ': 101,\n        'PAD': 0,\n        'END_SEQ': 102,\n        'UNK': 100\n    },\n    'xlm': {\n        'START_SEQ': 0,\n        'PAD': 2,\n        'END_SEQ': 1,\n        'UNK': 3\n    },\n    'roberta': {\n        'START_SEQ': 0,\n        'PAD': 1,\n        'END_SEQ': 2,\n        'UNK': 3\n    },\n    'albert': {\n        'START_SEQ': 2,\n        'PAD': 0,\n        'END_SEQ': 3,\n        'UNK': 1\n    },\n}\n\npunctuation_dict = {'O': 0, 'COMMA': 1, 'PERIOD': 2, 'QUESTION': 3}\n\nMODELS = {\n    'xlm-roberta-base': (XLMRobertaModel, XLMRobertaTokenizer, 768, 'roberta'),\n    'xlm-roberta-large': (XLMRobertaModel, XLMRobertaTokenizer, 1024, 'roberta')\n}\n\nclass DeepPunctuation(nn.Module):\n    \"\"\"\n    Initialize a Bengali specific punctuation model.\n\n    Args:\n        pretrained_model (str): The name of the pretrained model to use.\n        freeze_bert (bool): Whether to freeze the parameters of the BERT layer.\n        lstm_dim (int): The dimension of the LSTM hidden state. Set to -1 to use the BERT dimension.\n\n    \"\"\"\n    def __init__(self, pretrained_model, freeze_bert=False, lstm_dim=-1):\n        super(DeepPunctuation, self).__init__()\n        self.output_dim = len(punctuation_dict)\n        self.bert_layer = MODELS[pretrained_model][0].from_pretrained(CCFG.punc_base)\n        if freeze_bert:\n            for p in self.bert_layer.parameters():\n                p.requires_grad = False\n        bert_dim = MODELS[pretrained_model][2]\n        if lstm_dim == -1:\n            hidden_size = bert_dim\n        else:\n            hidden_size = lstm_dim\n        self.lstm = nn.LSTM(input_size=bert_dim, hidden_size=hidden_size, num_layers=1, bidirectional=True)\n        self.linear = nn.Linear(in_features=hidden_size*2, out_features=len(punctuation_dict))\n\n    def forward(self, x, attn_masks):\n        if len(x.shape) == 1:\n            x = x.view(1, x.shape[0])\n        x = self.bert_layer(x, attention_mask=attn_masks)[0]\n        x = torch.transpose(x, 0, 1)\n        x, (_, _) = self.lstm(x)\n        x = torch.transpose(x, 0, 1)\n        x = self.linear(x)\n        return x\n\ndef inference_punc(text):\n\n    text = re.sub(r\"[,:\\-–.!;?]\", '', text)\n    words_original_case = text.split()\n    words = text.lower().split()\n\n    word_pos = 0\n    sequence_len = 256\n    result = \"\"\n    decode_idx = 0\n    punctuation_map = {0: '', 1: ',', 2: '.', 3: '?'}\n    punctuation_map[2] = '।'\n\n\n    while word_pos < len(words):\n        x = [TOKEN_IDX[token_style]['START_SEQ']]\n        y_mask = [0]\n\n        while len(x) < sequence_len and word_pos < len(words):\n            tokens = tokenizer.tokenize(words[word_pos])\n            if len(tokens) + len(x) >= sequence_len:\n                break\n            else:\n                for i in range(len(tokens) - 1):\n                    x.append(tokenizer.convert_tokens_to_ids(tokens[i]))\n                    y_mask.append(0)\n                x.append(tokenizer.convert_tokens_to_ids(tokens[-1]))\n                y_mask.append(1)\n                word_pos += 1\n        x.append(TOKEN_IDX[token_style]['END_SEQ'])\n        y_mask.append(0)\n        if len(x) < sequence_len:\n            x = x + [TOKEN_IDX[token_style]['PAD'] for _ in range(sequence_len - len(x))]\n            y_mask = y_mask + [0 for _ in range(sequence_len - len(y_mask))]\n        attn_mask = [1 if token != TOKEN_IDX[token_style]['PAD'] else 0 for token in x]\n\n        x = torch.tensor(x).reshape(1,-1)\n        y_mask = torch.tensor(y_mask)\n        attn_mask = torch.tensor(attn_mask).reshape(1,-1)\n        x, attn_mask, y_mask = x.to(device), attn_mask.to(device), y_mask.to(device)\n\n        with torch.no_grad():\n        \n            y_predict = deep_punctuation(x, attn_mask)\n\n            #Identify the last word and cut the logits (so just the logits for | and ? are taken into account)\n            #We will force the model to output | or ? as last sign\n            last_id = torch.where(y_mask != 0)[0][-1].item()\n            last_sign = torch.argmax(y_predict[0][last_id][2:4]).item() + 2 \n\n            y_predict = y_predict.view(-1, y_predict.shape[2])\n            y_predict = torch.argmax(y_predict, dim=1).view(-1)\n                \n        for i in range(y_mask.shape[0]):\n            if y_mask[i] == 1:\n                if i == last_id:\n                    result += words_original_case[decode_idx] + punctuation_map[last_sign] + ' '\n                else:\n                    result += words_original_case[decode_idx] + punctuation_map[y_predict[i].item()] + ' '\n                decode_idx += 1\n\n    return result","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:06:19.590016Z","iopub.execute_input":"2023-10-20T07:06:19.590264Z","iopub.status.idle":"2023-10-20T07:06:19.610809Z","shell.execute_reply.started":"2023-10-20T07:06:19.590243Z","shell.execute_reply":"2023-10-20T07:06:19.609935Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def punctuation_new(sentence):\n    \"\"\"\n    Punctuate a given sentence using a punctuation inference model.\n\n    This function takes an input sentence and attempts to punctuate it using an inference model.\n    If the sentence does not end with a Bengali full stop (।), it processes the sentence as is.\n    If the sentence ends with ।, it processes the sentence without the final । and then adds it back after punctuation.\n\n    Args:\n        sentence (str): The input sentence to be punctuated.\n\n    Returns:\n        str: The punctuated sentence.\n\n    \"\"\"\n    try:\n    \n        if sentence[-1]!=\"।\":\n            sentence = inference_punc(sentence).strip()\n        else:\n            sentence = inference_punc(sentence[:-1]).strip()\n    except:\n        print(\"error\")\n        pass\n    return sentence","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:06:19.611935Z","iopub.execute_input":"2023-10-20T07:06:19.612173Z","iopub.status.idle":"2023-10-20T07:06:19.624187Z","shell.execute_reply.started":"2023-10-20T07:06:19.612153Z","shell.execute_reply":"2023-10-20T07:06:19.623487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class W2v2Dataset(torch.utils.data.Dataset):\n    \"\"\"\n    Custom PyTorch dataset for speech data processing with optional denoising and debugging.\n\n    Args:\n        paths (list): List of file paths to the audio files.\n        denoising (bool): Whether to apply denoising to the audio. -> not used in the final submission\n        debug_style (bool): Whether to use debug-style data loading. -> \n\n    Attributes:\n        paths (list): List of file paths to the audio files.\n        denoising (bool): Flag indicating whether denoising is applied.\n        debug (bool): Flag indicating whether debug-style data loading is used.\n    \"\"\"\n    def __init__(self, paths, denoising=False, debug_style=False):\n        self.paths = paths\n        self.denoising = denoising\n        self.debug = debug_style\n\n    def __getitem__(self, idx):\n        apath = self.paths[idx]\n        \n        if self.debug:\n            waveform, sample_rate = librosa.load(f\"{CCFG.train_data}/{apath}\", sr=16000)\n        else:\n            waveform, sample_rate = librosa.load(f\"{CCFG.test_data}/{apath}\", sr=16000)\n            \n        audio = processor(waveform, sampling_rate=sample_rate).input_values[0]\n        \n        if self.denoising:\n            audio = denoise_infer(audio)\n        \n        id_name = self.paths[idx].replace('.mp3','')\n        \n        return audio, id_name\n\n    def __len__(self):\n        return len(self.paths)","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:06:19.625203Z","iopub.execute_input":"2023-10-20T07:06:19.625438Z","iopub.status.idle":"2023-10-20T07:06:19.635610Z","shell.execute_reply.started":"2023-10-20T07:06:19.625417Z","shell.execute_reply":"2023-10-20T07:06:19.634494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loading the testfiles","metadata":{}},{"cell_type":"code","source":"test_files = sorted(os.listdir(CCFG.test_data))","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:06:19.636838Z","iopub.execute_input":"2023-10-20T07:06:19.637220Z","iopub.status.idle":"2023-10-20T07:06:19.672520Z","shell.execute_reply.started":"2023-10-20T07:06:19.637191Z","shell.execute_reply":"2023-10-20T07:06:19.671869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"debug_mode = False\n\ntest_dataset = W2v2Dataset(test_files, denoising=False, debug_style=debug_mode)\ntest_loader =  torch.utils.data.DataLoader(test_dataset,\n                             batch_size=1,\n                             shuffle=False,\n                             num_workers=os.cpu_count())","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:06:19.675348Z","iopub.execute_input":"2023-10-20T07:06:19.675588Z","iopub.status.idle":"2023-10-20T07:06:19.680281Z","shell.execute_reply.started":"2023-10-20T07:06:19.675567Z","shell.execute_reply":"2023-10-20T07:06:19.679471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Creating, calculating and saving logits","metadata":{}},{"cell_type":"code","source":"processor = Wav2Vec2Processor.from_pretrained(CCFG.processor)","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:06:19.681252Z","iopub.execute_input":"2023-10-20T07:06:19.681564Z","iopub.status.idle":"2023-10-20T07:06:19.727215Z","shell.execute_reply.started":"2023-10-20T07:06:19.681517Z","shell.execute_reply":"2023-10-20T07:06:19.726483Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomModel(nn.Module):\n    \"\"\"\n    Custom model to combine features of finetuned models.\n    Inspired by this paper: https://arxiv.org/pdf/2206.05518.pdf\n\n    This model combines two pretrained Wav2Vec2 models, processes their features with a transformer encoder,\n    and produces output logits for transcription. It supports training with CTC loss.\n\n    Attributes:\n        model1 (nn.Module): The first pretrained Wav2Vec2 model.\n        model2 (nn.Module): The second pretrained Wav2Vec2 model.\n        encoder_layers (nn.TransformerEncoderLayer): The transformer encoder layer.\n        transformer_encoder (nn.TransformerEncoder): The transformer encoder.\n        lm_head (nn.Linear): The linear layer for output logits.\n        config: The configuration of the model.\n\n    Methods:\n        freeze_model(model): Freezes the model's parameters to prevent further training.\n        forward(input_values, labels=None, **kwargs): Forward pass of the model.\n\n    \"\"\"\n    def __init__(self):\n        super().__init__()\n        self.model1 = self.freeze_model(Wav2Vec2Model.from_pretrained(CCFG.small_model))\n        self.model2 = self.freeze_model(Wav2Vec2Model.from_pretrained(CCFG.large_model))\n        self.encoder_layers = nn.TransformerEncoderLayer(d_model=1024+1280, nhead=6) #hardcoded values, should be initialized\n        self.transformer_encoder = nn.TransformerEncoder(self.encoder_layers, num_layers=2)\n        self.lm_head = nn.Linear(1024+1280, len(processor.tokenizer)) #hardcoded values, should be initialized\n        self.config = self.model1.config\n\n    def freeze_model(self, model):\n        for param in model.parameters():\n            param.requires_grad = False\n        return model\n\n    def forward(self, input_values, labels=None, **kwargs):\n\n        with torch.no_grad():\n            feature1 = self.model1(input_values=input_values, output_hidden_states=True).last_hidden_state\n            feature2 = self.model2(input_values=input_values, output_hidden_states=True).last_hidden_state\n\n        concatenated_features = torch.cat((feature1, feature2), dim=-1)\n        \n        encoded_features = self.transformer_encoder(concatenated_features)\n        logits = self.lm_head(encoded_features)\n        \n        if labels is None:\n            \n            return {'logits': logits}\n        \n        else: \n            loss = None\n            attention_mask = torch.ones_like(input_values, dtype=torch.long)\n            input_lengths = self.model1._get_feat_extract_output_lengths(attention_mask.sum(-1)).to(torch.long)\n            labels_mask = labels >= 0\n            target_lengths = labels_mask.sum(-1)\n            flattened_targets = labels.masked_select(labels_mask)\n\n            log_probs = nn.functional.log_softmax(logits, dim=-1, dtype=torch.float32).transpose(0, 1)\n\n            loss = nn.functional.ctc_loss(\n                log_probs,\n                flattened_targets,\n                input_lengths,\n                target_lengths,\n                blank=62,\n                reduction='mean',\n                zero_infinity=True,\n            )\n\n            return {'loss': loss, 'logits': logits}","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:06:19.728146Z","iopub.execute_input":"2023-10-20T07:06:19.728379Z","iopub.status.idle":"2023-10-20T07:06:19.738089Z","shell.execute_reply.started":"2023-10-20T07:06:19.728360Z","shell.execute_reply":"2023-10-20T07:06:19.737091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_model = CustomModel()\nnew_model.load_state_dict(torch.load(CCFG.ensemble_model))\nnew_model.to(device)\nprint(\"\")","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:06:19.739176Z","iopub.execute_input":"2023-10-20T07:06:19.739559Z","iopub.status.idle":"2023-10-20T07:08:05.885616Z","shell.execute_reply.started":"2023-10-20T07:06:19.739529Z","shell.execute_reply":"2023-10-20T07:08:05.884602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:08:05.886756Z","iopub.execute_input":"2023-10-20T07:08:05.887045Z","iopub.status.idle":"2023-10-20T07:08:06.191969Z","shell.execute_reply.started":"2023-10-20T07:08:05.887021Z","shell.execute_reply":"2023-10-20T07:08:06.190953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.mkdir(\"/kaggle/working/logits\")","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:08:06.193100Z","iopub.execute_input":"2023-10-20T07:08:06.193343Z","iopub.status.idle":"2023-10-20T07:08:06.202145Z","shell.execute_reply.started":"2023-10-20T07:08:06.193323Z","shell.execute_reply":"2023-10-20T07:08:06.201263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with torch.no_grad():\n    for i, (aud, id_n) in enumerate(tqdm(test_loader)):\n\n        aud = aud.to(device)\n        logits = new_model(aud)['logits']\n        logits = logits.detach().cpu().numpy()\n        np.save(f'/kaggle/working/logits/{id_n[0]}', logits)","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:08:06.203036Z","iopub.execute_input":"2023-10-20T07:08:06.203274Z","iopub.status.idle":"2023-10-20T07:08:20.791280Z","shell.execute_reply.started":"2023-10-20T07:08:06.203253Z","shell.execute_reply":"2023-10-20T07:08:20.790226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del new_model, logits, test_dataset, test_loader, processor\ntorch.cuda.empty_cache()\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:08:20.792659Z","iopub.execute_input":"2023-10-20T07:08:20.793061Z","iopub.status.idle":"2023-10-20T07:08:21.400223Z","shell.execute_reply.started":"2023-10-20T07:08:20.793017Z","shell.execute_reply":"2023-10-20T07:08:21.399297Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:08:21.401531Z","iopub.execute_input":"2023-10-20T07:08:21.402229Z","iopub.status.idle":"2023-10-20T07:08:21.705617Z","shell.execute_reply.started":"2023-10-20T07:08:21.402196Z","shell.execute_reply":"2023-10-20T07:08:21.704596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Creating and decoding the logits","metadata":{"execution":{"iopub.status.busy":"2023-10-13T10:49:22.02342Z","iopub.execute_input":"2023-10-13T10:49:22.023808Z","iopub.status.idle":"2023-10-13T10:49:22.04612Z","shell.execute_reply.started":"2023-10-13T10:49:22.023778Z","shell.execute_reply":"2023-10-13T10:49:22.045233Z"}}},{"cell_type":"code","source":"decoder = BeamSearchDecoderCTC.load_from_dir(CCFG.decoder)","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:08:21.707026Z","iopub.execute_input":"2023-10-20T07:08:21.707332Z","iopub.status.idle":"2023-10-20T07:16:46.985403Z","shell.execute_reply.started":"2023-10-20T07:08:21.707307Z","shell.execute_reply":"2023-10-20T07:16:46.981998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_params = {'alpha': 0.5, 'beta': 1.5, 'beam_width': 100} #Standard params\ncustom_decoding = False\nneural_rescoring = False\n\nif neural_rescoring:\n    os.mkdir(\"/kaggle/working/neural_rescoring\")\n    \nif custom_decoding:\n    decoder.reset_params(\n            alpha=0.75\n        )\n\nids = []\npredictions = []\n\nfor index, file in enumerate(test_files):\n\n    logits = np.load(f'/kaggle/working/logits/{file.replace(\".mp3\", \".npy\")}')\n    for l in logits:\n        if neural_rescoring:\n            sentence = decoder.decode_beams(l, prune_history=False)[:5]\n            sentence = [(x[0],x[-1]) for x in sentence]\n            with open(f\"/kaggle/working/neural_rescoring/{file.replace('.mp3','.pkl')}\", \"wb\") as f:\n                pickle.dump(sentence, f)\n        else:\n            sentence = decoder.decode(l)\n            predictions.append(sentence)\n    ids.append(file.replace(\".mp3\", \"\"))","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:16:46.990001Z","iopub.execute_input":"2023-10-20T07:16:46.990288Z","iopub.status.idle":"2023-10-20T07:16:48.985691Z","shell.execute_reply.started":"2023-10-20T07:16:46.990264Z","shell.execute_reply":"2023-10-20T07:16:48.984888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"decoder.cleanup() #-> important step, prevents from out of memory -> https://github.com/kensho-technologies/pyctcdecode/pull/111\n\ndel decoder\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:16:48.986770Z","iopub.execute_input":"2023-10-20T07:16:48.987039Z","iopub.status.idle":"2023-10-20T07:16:59.332875Z","shell.execute_reply.started":"2023-10-20T07:16:48.987017Z","shell.execute_reply":"2023-10-20T07:16:59.331991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:16:59.334068Z","iopub.execute_input":"2023-10-20T07:16:59.334293Z","iopub.status.idle":"2023-10-20T07:16:59.635485Z","shell.execute_reply.started":"2023-10-20T07:16:59.334275Z","shell.execute_reply":"2023-10-20T07:16:59.634571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Rescoring hypothesis","metadata":{}},{"cell_type":"code","source":"def rescore_hypotheses(hypotheses, lm_weight=1.0):\n    rescored_hypotheses = []\n    for transcription, asr_score in hypotheses:\n\n        input_ids = neural_rescore.tokenizer.encode(transcription, truncation=True, max_length=128, return_tensors='pt')\n\n        with torch.no_grad():\n            outputs = neural_rescore.model(input_ids.to(device), labels=input_ids.to(device))\n            log_likelihood = outputs.loss.item()\n\n        transformer_score = lm_weight * (-log_likelihood)\n        combined_score = asr_score + transformer_score \n        rescored_hypotheses.append((transcription, combined_score))\n        \n    return sorted(rescored_hypotheses, key=lambda x: x[1], reverse=True)","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:16:59.636496Z","iopub.execute_input":"2023-10-20T07:16:59.636769Z","iopub.status.idle":"2023-10-20T07:16:59.645435Z","shell.execute_reply.started":"2023-10-20T07:16:59.636719Z","shell.execute_reply":"2023-10-20T07:16:59.644768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if neural_rescoring:\n    neural_rescore = pipeline('text-generation',model=\"/kaggle/input/neural-rescoring\", tokenizer='/kaggle/input/neural-rescoring', device=device)\n    \n    predictions = []\n    for index, file in enumerate(test_files):\n        with open(f\"/kaggle/working/neural_rescoring/{file.replace('.mp3','.pkl')}\", \"rb\") as f:\n            loaded_list = pickle.load(f)\n        predictions.append(rescore_hypotheses(loaded_list)[0][0])\n        \n    del neural_rescore\n    torch.cuda.empty_cache()\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:16:59.646579Z","iopub.execute_input":"2023-10-20T07:16:59.646889Z","iopub.status.idle":"2023-10-20T07:16:59.655772Z","shell.execute_reply.started":"2023-10-20T07:16:59.646866Z","shell.execute_reply":"2023-10-20T07:16:59.654957Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Adding punctuation","metadata":{}},{"cell_type":"code","source":"tokenizer = XLMRobertaTokenizer.from_pretrained(CCFG.punc_base)\ntoken_style = MODELS['xlm-roberta-large'][3]","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:16:59.656670Z","iopub.execute_input":"2023-10-20T07:16:59.656981Z","iopub.status.idle":"2023-10-20T07:17:00.306567Z","shell.execute_reply.started":"2023-10-20T07:16:59.656960Z","shell.execute_reply":"2023-10-20T07:17:00.305760Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"deep_punctuation = DeepPunctuation('xlm-roberta-large', freeze_bert=False, lstm_dim=-1)\ndeep_punctuation.to(device)\ndeep_punctuation.load_state_dict(torch.load(CCFG.punc_weights))\ndeep_punctuation.eval()","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:17:00.307762Z","iopub.execute_input":"2023-10-20T07:17:00.308467Z","iopub.status.idle":"2023-10-20T07:17:43.169257Z","shell.execute_reply.started":"2023-10-20T07:17:00.308435Z","shell.execute_reply":"2023-10-20T07:17:43.168367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = [normalize(punctuation_new(normalize(sentence))) for sentence in predictions]","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:17:43.170134Z","iopub.execute_input":"2023-10-20T07:17:43.170361Z","iopub.status.idle":"2023-10-20T07:17:43.607653Z","shell.execute_reply.started":"2023-10-20T07:17:43.170343Z","shell.execute_reply":"2023-10-20T07:17:43.606983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del deep_punctuation\ntorch.cuda.empty_cache()\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:17:43.611783Z","iopub.execute_input":"2023-10-20T07:17:43.612103Z","iopub.status.idle":"2023-10-20T07:17:43.979508Z","shell.execute_reply.started":"2023-10-20T07:17:43.612077Z","shell.execute_reply":"2023-10-20T07:17:43.978504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Cleaning and submission","metadata":{}},{"cell_type":"code","source":"!rm -r /kaggle/working/*","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:17:43.980775Z","iopub.execute_input":"2023-10-20T07:17:43.981141Z","iopub.status.idle":"2023-10-20T07:17:45.164381Z","shell.execute_reply.started":"2023-10-20T07:17:43.981110Z","shell.execute_reply":"2023-10-20T07:17:45.163090Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_df = pd.DataFrame({\"id\":ids,\"sentence\":predictions})\npred_df[\"sentence\"] = [x if len(x) > 0 else \"।\" for x in pred_df[\"sentence\"]]\npred_df = pred_df.sort_values(by='id')\npred_df.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:17:45.166185Z","iopub.execute_input":"2023-10-20T07:17:45.166485Z","iopub.status.idle":"2023-10-20T07:17:45.213781Z","shell.execute_reply.started":"2023-10-20T07:17:45.166459Z","shell.execute_reply":"2023-10-20T07:17:45.213082Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_df.head(3)","metadata":{"execution":{"iopub.status.busy":"2023-10-20T07:17:45.214792Z","iopub.execute_input":"2023-10-20T07:17:45.215052Z","iopub.status.idle":"2023-10-20T07:17:45.226687Z","shell.execute_reply.started":"2023-10-20T07:17:45.215029Z","shell.execute_reply":"2023-10-20T07:17:45.225791Z"},"trusted":true},"execution_count":null,"outputs":[]}]}