{"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":"import os\nIS_KAGGLE = bool(os.environ.get('KAGGLE_KERNEL_RUN_TYPE', ''))","metadata":{"tags":[],"execution":{"iopub.status.busy":"2023-10-18T12:35:00.358611Z","iopub.execute_input":"2023-10-18T12:35:00.359083Z","iopub.status.idle":"2023-10-18T12:35:00.369468Z","shell.execute_reply.started":"2023-10-18T12:35:00.359059Z","shell.execute_reply":"2023-10-18T12:35:00.368578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IS_KAGGLE:\n    !cp -r /kaggle/input/whisper /kaggle/working\n    !pip install --no-index --find-links=/kaggle/input/bengali-wheels /kaggle/working/whisper\n    \n    !cp -r /kaggle/input/python-packages2 /tmp\n    !tar xvfz /tmp/python-packages2/normalizer.tgz\n    !pip install ./normalizer/bnunicodenormalizer-0.0.24.tar.gz -f ./ --no-index\n\n\n    !tar xvfz /tmp/python-packages2/jiwer.tgz\n    !pip install ./jiwer/python-Levenshtein-0.12.2.tar.gz -f ./ --no-index\n    !pip install ./jiwer/jiwer-2.3.0-py3-none-any.whl -f ./ --no-index","metadata":{"tags":[],"execution":{"iopub.status.busy":"2023-10-18T12:35:00.370774Z","iopub.execute_input":"2023-10-18T12:35:00.371002Z","iopub.status.idle":"2023-10-18T12:36:16.305584Z","shell.execute_reply.started":"2023-10-18T12:35:00.370984Z","shell.execute_reply":"2023-10-18T12:36:16.304329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IS_INFERENCE = True\nINPUT_DIR = \"../input/bengaliai-speech/\"\nFILES_DIR = INPUT_DIR + f\"{'test' if IS_INFERENCE else 'train'}_mp3s/\"\n# FILES_DIR = INPUT_DIR + \"examples/\"\nBATCH_SIZE = 24\nBEAM_SIZE = 2\nDECODER_TYPE = \"beam\"\n# DECODER_TYPE = \"greedy\"","metadata":{"execution":{"iopub.status.busy":"2023-10-18T12:36:16.307463Z","iopub.execute_input":"2023-10-18T12:36:16.307772Z","iopub.status.idle":"2023-10-18T12:36:16.312582Z","shell.execute_reply.started":"2023-10-18T12:36:16.307745Z","shell.execute_reply":"2023-10-18T12:36:16.311838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd, numpy as np, matplotlib.pyplot as plt\nimport torch, torch.nn as nn, torch.nn.functional as F\nimport torchaudio\nimport transformers\nimport timm\nimport math, time, sys, random, gc, os\nfrom tqdm.auto import tqdm\nimport whisper   # Requires download\nimport soundfile \nfrom bnunicodenormalizer import Normalizer\nimport jiwer\nfrom glob import glob\n\nfrom torch.nn.parallel import DistributedDataParallel as DDP\nfrom torch.distributed import init_process_group, destroy_process_group\nimport torch.multiprocessing as mp\nfrom concurrent.futures import ThreadPoolExecutor\n\n# MULTIPROCESS\ndevice_0 = \"cuda:0\"\ndevice_1 = \"cuda:1\"\n\n# init_process_group(backend='nccl')\n# ddp_rank = int(os.environ['RANK'])\n# ddp_local_rank = int(os.environ['LOCAL_RANK'])\n# ddp_world_size = int(os.environ['WORLD_SIZE'])\n# torch.cuda.set_device(f\"cuda:{ddp_local_rank}\")\n# torch.cuda.empty_cache()\n\nassert torch.__version__[0] == \"2\", \"Torch v2 is required\"\n\ntorch.backends.cuda.matmul.allow_tf32 = True\ntorch.backends.cudnn.allow_tf32 = True\n\n\nif IS_INFERENCE:\n    files = sorted(list(map(str, glob(FILES_DIR + \"*.mp3\"))))\n#     files = sorted(list(map(str, glob(FILES_DIR + \"*.wav\"))))\nelse:\n#     df = pd.read_csv(INPUT_DIR + \"train.csv\")\n#     df = df[df.split == \"valid\"].sample(500, random_state=42)\n#     files = [FILES_DIR + i + \".mp3\" for i in df.id.values]\n    files = sorted(list(map(str, glob(FILES_DIR + \"*.wav\"))))\n\nlengths = [soundfile.info(path).duration for path in files]\n\nsubmission_df = pd.read_csv(INPUT_DIR + \"sample_submission.csv\", nrows = 0)\nsubmission_df[\"paths\"]   = files\nsubmission_df[\"lengths\"] = lengths\nsubmission_df[\"id\"]      = [p.split(\"/\")[-1][:-4] for p in files]\n# submission_df.sort_values(by=\"lengths\", ascending=True, inplace=True)\nsubmission_df.sort_values(by=\"lengths\", ascending=False, inplace=True)\n\n\ntokenizer = whisper.tokenizer.get_tokenizer(False, language=\"Bengali\", task=\"transcribe\")\nINITIAL_SEQUENCE = list(tokenizer.sot_sequence) + [tokenizer.timestamp_begin]\n# INITIAL_SEQUENCE = list(tokenizer.sot_sequence_including_notimestamps)\n","metadata":{"execution":{"iopub.status.busy":"2023-10-18T12:36:16.313994Z","iopub.execute_input":"2023-10-18T12:36:16.314564Z","iopub.status.idle":"2023-10-18T12:36:27.605299Z","shell.execute_reply.started":"2023-10-18T12:36:16.314541Z","shell.execute_reply":"2023-10-18T12:36:27.604482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#### DATA #####\n\nclass InferenceDataset(torch.utils.data.Dataset):\n    def __init__(self, df):\n        super().__init__()\n        self.df = df\n        \n    def __len__(self): return len(self.df)\n\n    def __getitem__(self, idx):\n        item = self.df.iloc[idx]\n        audio = torch.from_numpy(whisper.load_audio(item[\"paths\"]))\n        n_samples = 60 * 16_000\n        audio = F.pad(audio, (0, n_samples))\n        audio = whisper.log_mel_spectrogram(audio)\n        return audio, item[\"id\"]\n\ndef inference_collator(batch):\n    audio, ids = map(list, zip(*batch))\n    return audio, ids\n\n\n#### DECODERS #####\n\nclass GreedyDecoder:\n    def __init__(self):\n        self.eot = tokenizer.eot\n\n    def reset(self, cache):\n        pass\n        \n    def update(self, tokens, logits, sum_logprobs):\n        next_token = logits.argmax(dim=-1)\n        next_token[tokens[:, -1] == self.eot] = self.eot\n        \n        tokens = torch.cat([tokens, next_token[:, None]], dim=-1)\n        completed = (tokens[:, -1] == self.eot).all()\n        return tokens, completed\n    \n    def finalize(self, preceding_tokens, sum_logprobs):\n        return F.pad(preceding_tokens.cpu(), (0,1), value=self.eot), sum_logprobs.cpu()\n    \n\nclass BeamDecoder:\n    def __init__(self):\n        self.beam_size = BEAM_SIZE\n        self.patience  = 1.0\n        self.max_candidates = round(self.beam_size * self.patience)\n        self.finished_sequences = None\n        self.eot = tokenizer.eot\n        self.cache = None\n        \n    \n    def reset(self, cache): \n        self.cache = cache\n        self.finished_sequences = None\n    \n    def update(self, tokens, logits, sum_logprobs):\n        advance = self.beam_size if self.finished_sequences is not None else 1\n        n_audio = tokens.shape[0] // advance\n        \n        if self.finished_sequences is None:\n            self.finished_sequences = [{} for _ in range(n_audio)]\n            self.unfinished_indices = list(range(n_audio))\n        \n        \n        probs = F.log_softmax(logits.float(), dim=-1)\n        next_tokens, source_indices, finished_sequences = [], [], []\n        \n        base_index = 0\n        cpu_tokens = tokens.cpu()\n        for i in range(n_audio):\n                \n            scores, sources, finished = {}, {}, {}\n            \n            for j in range(self.beam_size):\n                idx = i * advance + j\n                base = base_index * advance + j\n                prefix = cpu_tokens[idx].tolist()\n                \n                # Each beam gets to put N (beams) topk scores/tokens\n                # Calculate all candidates with cumulative probs\n                for prob, token in zip(*probs[idx].topk(self.beam_size + 1)):\n                    new_prob = (sum_logprobs[idx] + prob).item() # Cumulative probs\n                    sequence = tuple(prefix + [token.item()]) # New Sequence\n                    scores[sequence] = new_prob\n                    sources[sequence] = base\n                \n                if advance == 1:\n                    break\n            \n            # Keep top beam size based on scores\n            saved = 0\n            for sequence in sorted(scores, key=scores.get, reverse=True):\n                if sequence[-1] == self.eot:\n                    finished[sequence] = scores[sequence]\n                else:\n                    sum_logprobs[len(next_tokens)] = scores[sequence]\n                    next_tokens.append(sequence)\n                    if i in self.unfinished_indices:\n                        source_indices.append(sources[sequence])\n                    \n                    saved += 1\n                    if saved == self.beam_size:\n                        break\n            \n            if i in self.unfinished_indices:\n                base_index += 1\n            \n            finished_sequences.append(finished)\n        \n        tokens = torch.tensor(next_tokens, device=tokens.device)\n        \n        assert self.cache is not None\n        for module, tensor in self.cache.items():\n            self.cache[module] = tensor[source_indices].detach()\n        \n        assert len(self.finished_sequences) == len(finished_sequences)\n        \n        for index, (prev, new) in enumerate(zip(self.finished_sequences, finished_sequences)):\n            for seq in sorted(new, key=new.get, reverse=True):\n                if len(prev) >= self.max_candidates:\n                    if index in self.unfinished_indices:\n                    \n                        # Remove from cache\n                        n = []\n                        base = 0\n                        for i in self.unfinished_indices:\n                            if i != index:\n                                n.extend([base * self.beam_size + j for j in range(self.beam_size)])\n                            \n                            base += 1\n                        \n                        self.unfinished_indices.remove(index)\n                        \n                        for module,tensor in self.cache.items():\n                            self.cache[module] = tensor[n]\n\n                    break # list is complete of candidates\n                    \n                prev[seq] = new[seq]\n        \n        completed = all(\n            len(sequences) >= self.max_candidates for sequences in self.finished_sequences\n        )\n        return tokens, completed\n    \n    def finalize(self, preceding_tokens, sum_logprobs):\n        sum_logprobs = sum_logprobs.cpu()\n        for i, sequence in enumerate(self.finished_sequences):\n            if len(sequence) < self.beam_size:\n                for j in list(np.argsort(sum_logprobs[i]))[::-1]:\n                    seq = preceding_tokens[i, j].tolist() + [self.eot]\n                    sequence[tuple(seq)] = sum_logprobs[i][j].item()\n                    if len(sequence) > self.beam_size:\n                        break\n        \n        tokens = [\n            [torch.tensor(seq) for seq in sequences.keys()] for sequences in self.finished_sequences\n        ]\n        \n        sum_logprobs = [\n            list(seq.values()) for seq in self.finished_sequences\n        ]\n        \n        return tokens, sum_logprobs\n        \ndef rank(tokens, sum_logprobs):\n    def scores(probs, lenghts):\n        result = []\n        for prob, length in zip(probs, lenghts):\n            penalty = length\n            result.append(prob / penalty)\n        return result\n    \n    lenghts = [[len(token) for token in sequence] for sequence in tokens]\n    return [np.argmax(scores(p,l)) for p,l in zip(sum_logprobs, lenghts)]\n","metadata":{"execution":{"iopub.status.busy":"2023-10-18T12:36:27.607282Z","iopub.execute_input":"2023-10-18T12:36:27.607557Z","iopub.status.idle":"2023-10-18T12:36:27.630681Z","shell.execute_reply.started":"2023-10-18T12:36:27.607534Z","shell.execute_reply":"2023-10-18T12:36:27.629726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#### INIT MODELS ####\n\ndims = whisper.ModelDimensions(n_mels=80, n_audio_ctx=1500, n_audio_state=1024, n_audio_head=16, n_audio_layer=24, n_vocab=51865, n_text_ctx=448, n_text_state=1024, n_text_head=16, n_text_layer=24)\n# dims = whisper.ModelDimensions(n_mels=80, n_audio_ctx=1500, n_audio_state=1280, n_audio_head=20, n_audio_layer=32, n_vocab=51865, n_text_ctx=448, n_text_state=1280, n_text_head=20, n_text_layer=32)\nmodel = whisper.Whisper(dims).cuda().eval()\nn_vocab = model.decoder.token_embedding.weight.shape[0]\n\nstate = torch.load(\"/kaggle/input/bengali-weights/whisper/check.pt\")\n# state = torch.load(\"/kaggle/input/bengali-weights/whisper/check-Copy1.pt\")\nmodel.load_state_dict(state)\ndel state\ngc.collect()\ntorch.cuda.empty_cache()\n\nimport copy\nmodel  = torch.compile(model)\nmodel2 = copy.deepcopy(model).to(device_1)","metadata":{"execution":{"iopub.status.busy":"2023-10-18T12:36:27.632625Z","iopub.execute_input":"2023-10-18T12:36:27.633133Z","iopub.status.idle":"2023-10-18T12:37:04.781521Z","shell.execute_reply.started":"2023-10-18T12:36:27.633108Z","shell.execute_reply":"2023-10-18T12:37:04.780682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = InferenceDataset(submission_df)\ndl = torch.utils.data.DataLoader(ds, BATCH_SIZE, collate_fn=inference_collator, num_workers = 2)\nnorm = Normalizer()\n\ndef postprocess(sentence):\n    _words = [norm(word)['normalized']  for word in sentence.split()]\n    sentence = \" \".join([word for word in _words if word is not None])\n    return sentence\n\n\ndef get_filters(initial_sequence_len):\n    whisper_logit_filters = [\n        whisper.decoding.SuppressBlank(tokenizer, initial_sequence_len),\n        whisper.decoding.SuppressTokens(\n            list(tokenizer.non_speech_tokens) +\n            [\n                tokenizer.transcribe,\n                tokenizer.translate,\n                tokenizer.sot,\n                tokenizer.sot_prev,\n                tokenizer.sot_lm,\n            ]\n        )\n    ]\n    return whisper_logit_filters\n    ","metadata":{"execution":{"iopub.status.busy":"2023-10-18T12:37:04.782853Z","iopub.execute_input":"2023-10-18T12:37:04.783180Z","iopub.status.idle":"2023-10-18T12:37:04.791501Z","shell.execute_reply.started":"2023-10-18T12:37:04.783148Z","shell.execute_reply":"2023-10-18T12:37:04.790565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def infer_block(model, audio, initial_sequence, device):\n    with torch.cuda.device(int(device[-1])):\n        decoder = BeamDecoder() if DECODER_TYPE == \"beam\" else GreedyDecoder()\n        n_groups = BEAM_SIZE if isinstance(decoder, BeamDecoder) else 1\n        \n        audio = audio.to(device, non_blocking=True)\n        tokens = torch.tensor(initial_sequence, device=device)\n        if tokens.ndim == 1:\n            tokens = tokens[None]\n        \n        initial_len = tokens.shape[-1]\n\n        whisper_logit_filters = get_filters(initial_len)\n        \n        cache, hooks = model.install_kv_cache_hooks()\n        decoder.reset(cache)\n\n        with torch.no_grad():\n            with torch.cuda.amp.autocast():\n                features = model.encoder(audio)\n\n                while tokens.shape[-1] < 448:\n                    if tokens.shape[-1] > initial_len:\n                        if n_groups > 1:\n                            logits = torch.ones(tokens.shape[0], n_vocab, device=device)\n                            n = [n * n_groups + j for n in decoder.unfinished_indices for j in range(n_groups)]                    \n                            temp = tokens[n, -1:]\n                            temp_logits = model.decoder(temp, features[n, :], cache)[:, -1]\n                            logits[n] = temp_logits\n                        else:\n                            logits = model.decoder(tokens[:, -1:], features, cache)[:, -1]\n                    else:\n                        logits = model.decoder(tokens, features, cache)[:, -1]\n        \n                        # Init beam \n                        features = features.repeat_interleave(n_groups, dim=0).to(device)\n                        sum_logprobs = torch.zeros(features.shape[0], device=device)\n\n                    for f in whisper_logit_filters:\n                        f.apply(logits, tokens)\n\n                    tokens, completed = decoder.update(tokens, logits, sum_logprobs)\n\n                    if completed:\n                        break\n\n            for h in hooks:\n                h.remove()\n\n            tokens = tokens.reshape(audio.shape[0], n_groups, -1)\n            sum_logprobs = sum_logprobs.reshape(audio.shape[0], n_groups)\n\n            tokens, sum_logprobs = decoder.finalize(tokens, sum_logprobs)\n            \n            tokens = [\n                [group_tokens[initial_len: (group_tokens == tokenizer.eot).nonzero()[0,0]] for group_tokens in batch_tokens]\n                for batch_tokens in tokens\n            ]\n\n            selected = rank(tokens, sum_logprobs)   \n            final_tokens = [t[i].tolist() for i,t in zip(selected, tokens)]\n        \n        return final_tokens\n\n\nn_frames = 3_000\nstride = 2\nhop_rate = 160 / 16_000\ntolerance = 1 * 1 / hop_rate # 1 seconds\ncontext_length = 224\n\n# NOTE: whisper uses stride 2 in second conv so, 3_000 timee frames correspond to 1500. Each timestamp has to be multiplied back with 2\n\n# TODO: Handle 30 > secs\nfor batch in tqdm(dl):\n    audio, ids = batch\n    unfinished = [True] * len(audio)\n    was_max_length = [False] * len(audio)\n    prev_sequence_size = [0] * len(audio)\n    \n    durations = [\n        (a.shape[-1] - 2 * n_frames) for a in audio\n    ]\n    \n    \n    mel_offsets = [0] * len(audio)\n    \n    all_results = [[] for _ in range(len(audio))]\n    initial_sequences = [INITIAL_SEQUENCE for _ in range(len(audio))]\n    \n    step = 0\n    while sum(unfinished):\n        audio_segments = torch.stack([\n            a[:, mel_offsets[idx] : mel_offsets[idx] + n_frames]  for idx, a in enumerate(audio) if unfinished[idx]\n        ])\n\n        sequences = [seq for idx, seq in enumerate(initial_sequences) if unfinished[idx]]\n        \n        if step > 0:\n            results = []\n            for idx in range(0, len(sequences), 2):\n                \n                if idx == len(sequences) - 1:\n                    res = infer_block(model, audio_segments[[idx]], sequences[idx], device_0)\n                    results.extend(res)\n                else:\n                    with ThreadPoolExecutor(max_workers=2) as executor:\n                        res = sum(list(executor.map(infer_block, [model,model2],\n                                                    [audio_segments[[idx]], audio_segments[[idx+1]]],\n                                                    [sequences[idx], sequences[idx+1]],\n                                                    [device_0, device_1])), [])\n                    for r in res:\n                        results.append(r)\n\n        else:\n#             results = infer_block(model, audio_segments, sequences, device_0)\n            audio_segments0, audio_segments1 = torch.chunk(audio_segments, 2)\n            sequences0, sequences1 = sequences[:len(audio_segments0)], sequences[len(audio_segments0): ]\n            \n            with ThreadPoolExecutor(max_workers=2) as executor:\n                results = sum(list(executor.map(infer_block, [model,model2],\n                                                [audio_segments0, audio_segments1],\n                                                [sequences0, sequences1],\n                                                [device_0, device_1])), [])\n                \n        map_ = [i for i, non_finished in enumerate(unfinished) if non_finished]\n        for idx in range(len(results)):\n            res = results[idx]\n            \n            map_idx = map_[idx]\n            all_results[map_idx].extend(res)\n            current_result = all_results[map_idx]\n            \n            duration = durations[map_idx]\n            timestamp_tokens = torch.tensor(res).ge(tokenizer.timestamp_begin)\n            \n            advance = None\n            prompt = None\n            next_tokens = None\n\n#             vis = 0\n#             if map_idx == vis:\n#                 print(len(res))\n#                 print(prev_sequence_size[map_idx])\n#                 print(len(res) + prev_sequence_size[map_idx])\n    \n            if timestamp_tokens.sum() == 0:\n                # First case, none timestamp tokens, advance all chunk\n                if len(res) == 445:\n                    # First iteration, we have 445 tokens of total response\n                    advance = n_frames // 2\n                    next_tokens = res[-context_length:]\n                elif (len(res) + prev_sequence_size[map_idx]) == 445:\n                    # In the middle of a iteration\n                    if len(res) < 100:\n                        advance = int(2 // hop_rate)\n                    else:\n                        advance = int(5 // hop_rate)\n                    \n                    next_tokens = res[-context_length:]\n                else:\n                    advance = n_frames\n                \n                prompt  = current_result\n\n            elif timestamp_tokens[-1]:\n                # Second case, there is a timestamp token at the end, check if it corresponds to the total audio len \n                was_max_length[map_idx] = False\n                \n                end_timestamp = res[-1] - tokenizer.timestamp_begin # TODO \n                end_timestamp = end_timestamp * 2\n                \n                    \n                if (end_timestamp + mel_offsets[map_idx]) < (duration - tolerance):\n                    # If it is smaller, start from here next iteration, otherwise we have FINISHED\n                    advance = end_timestamp\n                    prompt  = current_result\n                    current_result.append(220)\n\n            else:\n                # Third case, there is a segment in progress but not finished\n                \n                was_max_length[map_idx] = False\n                \n                last_timestamp_position = torch.where(timestamp_tokens)[0][-1]\n                end_timestamp = res[last_timestamp_position] - tokenizer.timestamp_begin\n                \n                advance     = end_timestamp * 2\n                next_tokens = res[last_timestamp_position + 2:]\n                prompt      = current_result[:-len(next_tokens)]\n\n            if advance is not None:\n                mel_offsets[map_idx] += advance\n                \n                new_tokens = []\n                \n                available_len = context_length\n                \n                if next_tokens is not None:\n                    available_len -= len(next_tokens)\n                \n                \n                if prompt is not None and available_len > 3:\n                    prompt = copy.deepcopy(prompt)\n                    removed = 0\n                    for i in range(len(prompt)):\n                        if prompt[i-removed] > tokenizer.timestamp_begin:\n                            del prompt[i-removed]\n                            removed += 1\n                            \n                            \n                    new_tokens = new_tokens + [tokenizer.sot_prev] + prompt[-available_len:]\n                    \n                new_tokens = new_tokens + INITIAL_SEQUENCE\n                    \n                if next_tokens is not None:\n                    new_tokens = new_tokens + next_tokens\n                 \n                prev_sequence_size[map_idx] = len(new_tokens) - 2\n                initial_sequences[map_idx] =  new_tokens\n            else:\n                # Can only access from second case\n                unfinished[map_idx] = False\n            \n            if mel_offsets[map_idx] > durations[map_idx]:\n                unfinished[map_idx] = False\n            \n        step += 1\n#         print(f\"STEP: {step}\")\n#         print(tokenizer.decode_with_timestamps(all_results[vis]))\n#         print(\"-----------------\")\n#         print(mel_offsets[vis], durations[vis])\n#         print(\"-----------------\")\n#         print(tokenizer.decode_with_timestamps(initial_sequences[vis]))\n\n    \n#     break\n    \n    for res, audio_id in zip(all_results, ids):\n        sentence   = tokenizer.decode(res)\n        sentence   = postprocess(sentence)\n        sentence   = sentence if len(sentence) > 0 else \"।\"\n        submission_df.loc[submission_df.id == audio_id, \"sentence\"] = sentence \n    \n        \nsubmission_df = submission_df.sort_index()\n    \nif IS_INFERENCE:\n    submission_df = submission_df[[\"id\", \"sentence\"]]\n    submission_df.to_csv(\"submission.csv\", index=False)\nelse:\n#     train_csv = pd.read_csv(\"../input/bengaliai-speech/train.csv\").rename(columns={\"sentence\" : \"ground_truth\"})\n#     submission_df = submission_df.merge(train_csv[[\"id\", \"ground_truth\"]], on=\"id\", how=\"left\")\n#     score = jiwer.wer(submission_df[\"ground_truth\"].to_list(), submission_df[\"sentence\"].to_list())\n#     print(score)\n\n    ann = pd.read_csv(\"/kaggle/input/ood-example-audios-hand-annotations/annoated.csv\", sep=\"\\t\")\n    ann.rename(columns={\"sentence\": \"ground\", \"file\": \"id\"}, inplace=True)\n    ann[\"id\"] = [p.replace(\".wav\", \"\") for p in ann.id.values]\n    ann = ann.set_index(\"id\")\n\n    submission_df = submission_df.merge(ann, on=\"id\", how=\"left\")\n    score = jiwer.wer(submission_df.ground.to_list(), submission_df.sentence.to_list())\n    print(score)\n    ","metadata":{"execution":{"iopub.status.busy":"2023-10-18T12:37:04.792660Z","iopub.execute_input":"2023-10-18T12:37:04.793034Z","iopub.status.idle":"2023-10-18T12:43:12.364147Z","shell.execute_reply.started":"2023-10-18T12:37:04.792984Z","shell.execute_reply":"2023-10-18T12:43:12.363041Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def remove_punc(string_a):\n#     import string\n#     return string_a.translate(str.maketrans(\"\",\"\",string.punctuation))","metadata":{"execution":{"iopub.status.busy":"2023-10-18T12:43:12.366146Z","iopub.execute_input":"2023-10-18T12:43:12.366428Z","iopub.status.idle":"2023-10-18T12:43:12.372875Z","shell.execute_reply.started":"2023-10-18T12:43:12.366405Z","shell.execute_reply":"2023-10-18T12:43:12.372050Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ground = [remove_punc(p) for p in submission_df.ground.to_list()]\n# sens = [remove_punc(p) for p in submission_df.sentence.to_list()]","metadata":{"execution":{"iopub.status.busy":"2023-10-18T12:43:12.373984Z","iopub.execute_input":"2023-10-18T12:43:12.374385Z","iopub.status.idle":"2023-10-18T12:43:12.386293Z","shell.execute_reply.started":"2023-10-18T12:43:12.374362Z","shell.execute_reply":"2023-10-18T12:43:12.385603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# jiwer.wer(ground, sens)","metadata":{"execution":{"iopub.status.busy":"2023-10-18T12:43:12.390536Z","iopub.execute_input":"2023-10-18T12:43:12.391132Z","iopub.status.idle":"2023-10-18T12:43:12.400028Z","shell.execute_reply.started":"2023-10-18T12:43:12.391096Z","shell.execute_reply":"2023-10-18T12:43:12.399040Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# wers = [\n#     jiwer.wer(g,s) for g,s in zip(ground, sens)\n# ]","metadata":{"execution":{"iopub.status.busy":"2023-10-18T12:43:12.401489Z","iopub.execute_input":"2023-10-18T12:43:12.401842Z","iopub.status.idle":"2023-10-18T12:43:12.412747Z","shell.execute_reply.started":"2023-10-18T12:43:12.401805Z","shell.execute_reply":"2023-10-18T12:43:12.411878Z"},"trusted":true},"execution_count":null,"outputs":[]}]}