{"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":"- This is a training demo, you can run this code locally, using better GPUs.\n- The inference part is here: [Bengali SR wav2vec_v1_bengali [Inference]](https://www.kaggle.com/takanashihumbert/bengali-sr-wav2vec-v1-bengali-inference), it scores **0.445** on the leaderboard.\n- Feel free to upvote, thanks!","metadata":{}},{"cell_type":"code","source":"%%capture\n!cp -r ../input/python-packages2 ./\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","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-10-17T12:51:01.628574Z","iopub.execute_input":"2023-10-17T12:51:01.628816Z","iopub.status.idle":"2023-10-17T12:52:06.680477Z","shell.execute_reply.started":"2023-10-17T12:51:01.628792Z","shell.execute_reply":"2023-10-17T12:52:06.679252Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%capture\n%mkdir torch-audiomentations\n%cd torch-audiomentations\n!pip wheel torch-audiomentations\n%cd ../","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:52:06.682473Z","iopub.execute_input":"2023-10-17T12:52:06.682740Z","iopub.status.idle":"2023-10-17T12:53:31.514661Z","shell.execute_reply.started":"2023-10-17T12:52:06.682715Z","shell.execute_reply":"2023-10-17T12:53:31.512925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%capture\n!pip install torch-audiomentations\n!pip install torch-time-stretch","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:53:31.516565Z","iopub.execute_input":"2023-10-17T12:53:31.516935Z","iopub.status.idle":"2023-10-17T12:53:51.881901Z","shell.execute_reply.started":"2023-10-17T12:53:31.516900Z","shell.execute_reply":"2023-10-17T12:53:51.880703Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install praat-parselmouth","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:53:51.884658Z","iopub.execute_input":"2023-10-17T12:53:51.885600Z","iopub.status.idle":"2023-10-17T12:54:00.872227Z","shell.execute_reply.started":"2023-10-17T12:53:51.885561Z","shell.execute_reply":"2023-10-17T12:54:00.870816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch \nimport torch.nn as nn\nimport torchaudio\nimport torchaudio.transforms as tat\nfrom datasets import load_dataset, load_metric, Audio\nfrom torch_audiomentations import Compose, Gain, PitchShift, Shift, AddColoredNoise, PolarityInversion\n \nimport os\n\nimport typing as tp\nfrom pathlib import Path\nfrom functools import partial\nfrom dataclasses import dataclass, field\nfrom typing import Any, Dict, List, Optional, Union\nfrom torch_time_stretch import *\n\nimport pandas as pd\nimport parselmouth\nimport pyctcdecode\nimport numpy as np\nfrom tqdm.notebook import tqdm\n\nimport librosa\nimport gc\nimport jiwer\nimport pyctcdecode\nimport kenlm\nimport torch\nfrom transformers import Wav2Vec2Processor, Wav2Vec2ProcessorWithLM, Wav2Vec2ForCTC\nfrom transformers import TrainingArguments, Trainer, EarlyStoppingCallback\nfrom bnunicodenormalizer import Normalizer\nimport warnings\nwarnings.filterwarnings('ignore')\ntorchaudio.set_audio_backend(\"soundfile\")","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:54:00.873888Z","iopub.execute_input":"2023-10-17T12:54:00.874974Z","iopub.status.idle":"2023-10-17T12:54:15.087016Z","shell.execute_reply.started":"2023-10-17T12:54:00.874936Z","shell.execute_reply":"2023-10-17T12:54:15.086135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### hyper-parameters\nSR = 16000\ntorch.backends.cudnn.benchmark = True\noutput_dir = \"./\"\nMODEL_PATH = \"/kaggle/input/ai4bharat-indicwav2vec-v1-bengali/indicwav2vec_v1_bengali\"\nLM_PATH = \"/kaggle/input/arijitx-full-model/wav2vec2-xls-r-300m-bengali/language_model\"","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:54:15.088351Z","iopub.execute_input":"2023-10-17T12:54:15.088657Z","iopub.status.idle":"2023-10-17T12:54:15.093274Z","shell.execute_reply.started":"2023-10-17T12:54:15.088626Z","shell.execute_reply":"2023-10-17T12:54:15.092284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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])}\n\ndecoder = pyctcdecode.build_ctcdecoder(\n    list(sorted_vocab_dict.keys()),\n    str(LM_PATH+\"/5gram.bin\"),\n    str(LM_PATH+\"/unigrams.txt\"),\n)\nprocessor_with_lm = Wav2Vec2ProcessorWithLM(\n    feature_extractor=processor.feature_extractor,\n    tokenizer=processor.tokenizer,\n    decoder=decoder\n)","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:54:15.094538Z","iopub.execute_input":"2023-10-17T12:54:15.095049Z","iopub.status.idle":"2023-10-17T12:54:57.715059Z","shell.execute_reply.started":"2023-10-17T12:54:15.095019Z","shell.execute_reply":"2023-10-17T12:54:57.713862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- From @mbmmurad's [Dataset overlaps with CommonVoice 11 bn](https://www.kaggle.com/code/mbmmurad/dataset-overlaps-with-commonvoice-11-bn), The competition dataset might contain the audios of the mozilla-foundation/common_voice_11_0 dataset. Here I just simply exclude them from the validation set.\n- Also, I use @UmongSain's normalized data [here](https://www.kaggle.com/code/umongsain/macro-normalization/notebook). Thanks to him!","metadata":{}},{"cell_type":"code","source":"sentences = pd.read_csv(\"/kaggle/input/macro-normalization/normalized.csv\")\nindexes = set(pd.read_csv(\"/kaggle/input/dataset-overlaps-with-commonvoice-11-bn/indexes.csv\")['id'])\nprint(len(sentences))\nsentences = sentences[~((sentences.index.isin(indexes))&(sentences['split']=='train'))].reset_index(drop=True)\nprint(len(sentences))","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:54:57.716596Z","iopub.execute_input":"2023-10-17T12:54:57.716907Z","iopub.status.idle":"2023-10-17T12:55:06.139025Z","shell.execute_reply.started":"2023-10-17T12:54:57.716877Z","shell.execute_reply":"2023-10-17T12:55:06.137980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"* sample 10% data from \"valid\" part into validation set, 90% into training set.\n* sample 5% data from \"train\" part, and additionally sample 8% from it into validation set, 92% into training set.\n* There will be **57776** train data, **5667** valid data.","metadata":{}},{"cell_type":"code","source":"data_0 = sentences.loc[sentences['split']=='valid'].reset_index(drop=True)\nvalid_0 = data_0.sample(frac=0.1, random_state=42)\ntrain_0 = data_0[~data_0.index.isin(valid_0.index)]\n\ndata_1 = sentences.loc[sentences['split']=='train'].reset_index(drop=True).sample(frac=0.05, random_state=42)\nvalid_1 = data_1.sample(frac=0.08, random_state=42)\ntrain_1 = data_1[~data_1.index.isin(valid_1.index)]\n\ntrain = pd.concat([train_0, train_1], axis=0).sample(frac=1, random_state=42).reset_index(drop=True)\nvalid = pd.concat([valid_0, valid_1], axis=0).sample(frac=1, random_state=42).reset_index(drop=True)\n\ndel data_0, data_1, valid_0, valid_1, train_0, train_1\nall_ids = sentences['id'].to_list()\ntrain_ids = train['id'].to_list()\nvalid_ids = valid['id'].to_list()\n\n# in kaggle notebook, validating is very time-consuming, so here I use a very small validation set, rather than 5667.\nvalid = valid.sample(n=500, random_state=42)\n\nprint(len(all_ids))\nprint(len(train_ids))\nprint(len(valid_ids))","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:55:06.143028Z","iopub.execute_input":"2023-10-17T12:55:06.143477Z","iopub.status.idle":"2023-10-17T12:55:06.520116Z","shell.execute_reply.started":"2023-10-17T12:55:06.143450Z","shell.execute_reply":"2023-10-17T12:55:06.519141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ParametricEqualizer(nn.Module):\n    \"\"\"Fast-parametric equalizer for approximation of Biquad IIR filter.\n    \"\"\"\n    def __init__(self, sr: int, windows: int):\n        \"\"\"Initializer.\n        Args:\n            sr: sample rate.\n            windows: size of the fft window.\n        \"\"\"\n        super().__init__()\n        self.sr = sr\n        self.windows = windows\n\n    def biquad(self, a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:\n        \"\"\"Construct frequency level biquad filter.\n        Args:\n            a: [torch.float32; [..., 3]], recursive filter, iir.\n            b: [torch.float32; [..., 3]], finite impulse filter.\n        Returns:\n            [torch.float32; [..., windows // 2 + 1]], biquad filter.\n        \"\"\"\n        iir = torch.fft.rfft(a, self.windows, dim=-1)\n        fir = torch.fft.rfft(b, self.windows, dim=-1)\n        return fir / iir\n\n    def low_shelving(self,\n                     cutoff: float,\n                     gain: torch.Tensor,\n                     q: torch.Tensor) -> torch.Tensor:\n        \"\"\"Frequency level low-shelving filter.\n        Args:\n            cutoff: cutoff frequency.\n            gain: [torch.float32; [...]], boost of attenutation in decibel.\n            q: [torch.float32; [...]], quality factor.\n        Returns:\n            [torch.float32; [..., windows // 2 + 1]], frequency filter.\n        \"\"\"\n        # ref: torchaudio.functional.lowpass_biquad\n        w0 = 2 * np.pi * cutoff / self.sr\n        cos_w0 = np.cos(w0)\n        # [B]\n        alpha = np.sin(w0) / 2 / q\n        cos_w0 = torch.full_like(alpha, np.cos(w0))\n        A = (gain / 40. * np.log(10)).exp()\n        # [...], fir\n        b0 = A * ((A + 1) - (A - 1) * cos_w0 + 2 * A.sqrt() * alpha)\n        b1 = 2 * A * ((A - 1) - (A + 1) * cos_w0)\n        b2 = A * ((A + 1) - (A - 1) * cos_w0 - 2 * A.sqrt() * alpha)\n        # [...], iir\n        a0 = (A + 1) + (A - 1) * cos_w0 + 2 * A.sqrt() * alpha\n        a1 = -2 * ((A - 1) + (A + 1) * cos_w0)\n        a2 = (A + 1) + (A - 1) * cos_w0 - 2 * A.sqrt() * alpha\n        # [..., windows // 2 + 1]\n        return self.biquad(\n            a=torch.stack([a0, a1, a2], dim=-1),\n            b=torch.stack([b0, b1, b2], dim=-1))\n\n    def high_shelving(self,\n                      cutoff: float,\n                      gain: torch.Tensor,\n                      q: torch.Tensor) -> torch.Tensor:\n        \"\"\"Frequency level high-shelving filter.\n        Args:\n            cutoff: cutoff frequency.\n            gain: [torch.float32; [...]], boost of attenutation in decibel.\n            q: [torch.float32; [...]], quality factor.\n        Returns:\n            [torch.float32; [..., windows // 2 + 1]], frequency filter.\n        \"\"\"\n        # ref: torchaudio.functional.highpass_biquad\n        w0 = 2 * np.pi * cutoff / self.sr\n        # [...]\n        alpha = np.sin(w0) / 2 / q\n        cos_w0 = torch.full_like(alpha, np.cos(w0))\n        A = (gain / 40. * np.log(10)).exp()\n        # [...], fir\n        b0 = A * ((A + 1) + (A - 1) * cos_w0 + 2 * A.sqrt() * alpha)\n        b1 = -2 * A * ((A - 1) + (A + 1) * cos_w0)\n        b2 = A * ((A + 1) + (A - 1) * cos_w0 - 2 * A.sqrt() * alpha)\n        # [...], iir\n        a0 = (A + 1) - (A - 1) * cos_w0 + 2 * A.sqrt() * alpha\n        a1 = 2 * ((A - 1) - (A + 1) * cos_w0)\n        a2 = (A + 1) - (A - 1) * cos_w0 - 2 * A.sqrt() * alpha\n        # [..., windows // 2 + 1]\n        return self.biquad(\n            a=torch.stack([a0, a1, a2], dim=-1),\n            b=torch.stack([b0, b1, b2], dim=-1))\n\n    def peaking_equalizer(self,\n                          center: torch.Tensor,\n                          gain: torch.Tensor,\n                          q: torch.Tensor) -> torch.Tensor:\n        \"\"\"Frequency level peaking equalizer.\n        Args:\n            center: [torch.float32; [...]], center frequency.\n            gain: [torch.float32; [...]], boost or attenuation in decibel.\n            q: [torch.float32; [...]], quality factor.\n        Returns:\n            [torch.float32; [..., windows // 2 + 1]], frequency filter.\n        \"\"\"\n        # ref: torchaudio.functional.highpass_biquad\n        # [...]\n        w0 = 2 * np.pi * center / self.sr\n        # [...]\n        alpha = torch.sin(w0) / 2 / q\n        cos_w0 = torch.cos(w0)\n        A = (gain / 40. * np.log(10)).exp()\n        # [..., windows // 2 + 1]\n        return self.biquad(\n            a=torch.stack([1 + alpha / A, -2 * cos_w0, 1 - alpha / A], dim=-1),\n            b=torch.stack([1 + alpha * A, -2 * cos_w0, 1 - alpha * A], dim=-1))","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:55:06.523913Z","iopub.execute_input":"2023-10-17T12:55:06.524526Z","iopub.status.idle":"2023-10-17T12:55:06.560993Z","shell.execute_reply.started":"2023-10-17T12:55:06.524501Z","shell.execute_reply":"2023-10-17T12:55:06.560147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PraatAugment:\n    \"\"\"Praat based augmentation.\n    \"\"\"\n    def __init__(self,\n                 pitch_steps: float = 0.01,\n                 pitch_floor: float = 75,\n                 pitch_ceil: float = 600):\n        \"\"\"Initializer.\n        Args:\n            config: configurations.\n            pitch_steps: pitch measurement intervals.\n            pitch_floor: minimum pitch.\n            pitch_ceil: maximum pitch.\n        \"\"\"\n        self.pitch_steps = pitch_steps\n        self.pitch_floor = pitch_floor\n        self.pitch_ceil = pitch_ceil\n\n    def augment(self,\n                snd: Union[parselmouth.Sound, np.ndarray],\n                formant_shift: float = 1.,\n                pitch_shift: float = 1.,\n                pitch_range: float = 1.,\n                duration_factor: float = 1.) -> np.ndarray:\n        \"\"\"Augment the sound signal with praat.\n        \"\"\"\n        if not isinstance(snd, parselmouth.Sound):\n            snd = parselmouth.Sound(snd, sampling_frequency=SR)\n        pitch = parselmouth.praat.call(\n            snd, 'To Pitch', self.pitch_steps, self.pitch_floor, self.pitch_ceil)\n        ndpit = pitch.selected_array['frequency']\n        # if all unvoiced\n        nonzero = ndpit > 1e-5\n        if nonzero.sum() == 0:\n            return snd.values[0]\n        # if voiced\n        median, minp = np.median(ndpit[nonzero]).item(), ndpit[nonzero].min().item()\n        # scale\n        updated = median * pitch_shift\n        scaled = updated + (minp * pitch_shift - updated) * pitch_range\n        # for preventing infinite loop of `Change gender`\n        # ref:https://github.com/praat/praat/issues/1926\n        if scaled < 0.:\n            pitch_range = 1.\n        out, = parselmouth.praat.call(\n            (snd, pitch), 'Change gender',\n            formant_shift,\n            median * pitch_shift,\n            pitch_range,\n            duration_factor).values\n        return out","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:55:06.562953Z","iopub.execute_input":"2023-10-17T12:55:06.563370Z","iopub.status.idle":"2023-10-17T12:55:06.581273Z","shell.execute_reply.started":"2023-10-17T12:55:06.563331Z","shell.execute_reply":"2023-10-17T12:55:06.580303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Augment(nn.Module):\n    \"\"\"Waveform augmentation.\n    \"\"\"\n    def __init__(self):\n        \"\"\"Initializer.\n        Args:\n            config: Nansy configurations.\n        \"\"\"\n        super().__init__()\n        self.praat = PraatAugment()\n        self.peq = ParametricEqualizer(16000, 1024)\n        self.register_buffer('window',torch.hann_window(1024),persistent=False)\n        f_min, f_max, peaks = 60, 10000, 8\n        # peaks except frequency min and max\n        self.register_buffer(\n            'peak_centers',\n            f_min * (f_max / f_min) ** (torch.arange(peaks + 2)[1:-1] / (peaks + 1)),\n            persistent=False)\n\n    def forward(self,\n                wavs: torch.Tensor,\n                pitch_shift: Optional[torch.Tensor] = None,\n                pitch_range: Optional[torch.Tensor] = None,\n                formant_shift: Optional[torch.Tensor] = None,\n                quality_power: Optional[torch.Tensor] = None,\n                gain: Optional[torch.Tensor] = None) -> torch.Tensor:\n        \"\"\"Augment the audio signal, random pitch, formant shift and PEQ.\n        Args:\n            wavs: [torch.float32; [B, T]], audio signal.\n            pitch_shift: [torch.float32; [B]], pitch shifts.\n            pitch_range: [torch.float32; [B]], pitch ranges.\n            formant_shift: [torch.float32; [B]], formant shifts.\n            quality_power: [torch.float32; [B, num_peak + 2]],\n                exponents of quality factor, for PEQ.\n            gain: [torch.float32; [B, num_peak + 2]], gain in decibel.\n        Returns:\n            [torch.float32; [B, T]], augmented.\n        \"\"\"\n        # B\n        bsize, _ = wavs.shape\n        # [B, F, T / S], complex64\n        fft = torch.stft(\n            wavs,\n            1024,\n            256,\n            1024,\n            self.window,\n            return_complex=True)\n        # PEQ\n        if quality_power is not None:\n            # alias\n            q_min, q_max = 2, 5\n            # [B, num_peak + 2]\n            q = q_min * (q_max / q_min) ** quality_power\n            if gain is None:\n                # [B, num_peak]\n                gain = torch.zeros_like(q[:, :-2])\n            # [B, num_peak]\n            center = self.peak_centers[None].repeat(bsize, 1)\n            # [B, F]\n            peaks = torch.prod(\n                self.peq.peaking_equalizer(center, gain[:, :-2], q[:, :-2]), dim=1)\n            # [B, F]\n            lowpass = self.peq.low_shelving(60, gain[:, -2], q[:, -2])\n            highpass = self.peq.high_shelving(10000, gain[:, -1], q[:, -1])\n            # [B, F]\n            filters = peaks * highpass * lowpass\n            # [B, F, T / S]\n            fft = fft * filters[..., None]\n        # [B, T]\n        out = torch.istft(\n            fft,\n            1024,\n            256,\n            1024,\n            self.window).clamp(-1., 1.)\n        # max value normalization\n        out = out / out.abs().amax(dim=-1, keepdim=True).clamp_min(1e-7)\n        if formant_shift is None and pitch_shift is None and pitch_range is None:\n            return out\n        # praat-based augmentation\n        if formant_shift is None:\n            formant_shift = torch.ones(bsize)\n        if pitch_shift is None:\n            pitch_shift = torch.ones(bsize)\n        if pitch_range is None:\n            pitch_range = torch.ones(bsize)\n        out = torch.tensor(\n            np.stack([\n                self.praat.augment(o, fs.item(), ps.item(), pr.item())\n                for o, fs, ps, pr in zip(\n                    out.cpu().numpy(),\n                    formant_shift.cpu().numpy(),\n                    pitch_shift.cpu().numpy(),\n                    pitch_range.cpu().numpy())], axis=0),\n            device=out.device, dtype=torch.float32)\n        return out","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:55:06.582666Z","iopub.execute_input":"2023-10-17T12:55:06.583245Z","iopub.status.idle":"2023-10-17T12:55:06.599269Z","shell.execute_reply.started":"2023-10-17T12:55:06.583215Z","shell.execute_reply":"2023-10-17T12:55:06.598269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Taken from Torch-Nansy\n\ndef sample_like(signal: torch.Tensor) -> List[torch.Tensor]:\n        \"\"\"Sample augmentation parameters.\n        Args:\n            signal: [torch.float32; [B, T]], speech signal.\n        Returns:\n            augmentation parameters.\n        \"\"\"\n        # [B]\n        bsize, _ = signal.shape\n        def sampler(ratio):\n            shifts = torch.rand(bsize, device=signal.device) * (ratio - 1.) + 1.\n            # flip\n            flip = torch.rand(bsize) < 0.5\n            shifts[flip] = shifts[flip] ** -1\n            return shifts\n        # sample shifts\n        fs = sampler(1.4)\n        ps = sampler(2.)\n        pr = sampler(1.5)\n        # parametric equalizer\n        peaks = 8\n        # quality factor\n        power = torch.rand(bsize, peaks + 2, device=signal.device)\n        # gains\n        g_min, g_max = -12, 12\n        gain = torch.rand(bsize, peaks + 2, device=signal.device) * (g_max - g_min) + g_min\n        return fs, ps, pr, power, gain","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:55:06.600645Z","iopub.execute_input":"2023-10-17T12:55:06.601228Z","iopub.status.idle":"2023-10-17T12:55:06.615386Z","shell.execute_reply.started":"2023-10-17T12:55:06.601170Z","shell.execute_reply":"2023-10-17T12:55:06.614409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def augment(signal: torch.Tensor, ps: bool = True) -> torch.Tensor:\n        \"\"\"Augment the speech.\n        Args:\n            signal: [torch.float32; [B, T]], segmented speech.\n            ps: whether use pitch shift.\n        Returns:\n            [torch.float32; [B, T]], speech signal.\n        \"\"\"\n        # B\n        bsize, _ = signal.shape\n        saves = None\n        while saves is None or len(saves) < bsize:\n            # [B] x 4\n            fshift, pshift, prange, power, gain = sample_like(signal)\n            if not ps:\n                pshift = None\n            # [B, T]\n            aug = Augment()\n            out = aug.forward(wavs=signal, pitch_shift=pshift, pitch_range=prange, formant_shift=fshift, quality_power=power, gain=gain)\n            # for covering unexpected NaN\n            nan = out.isnan().any(dim=-1)\n            if not nan.all():\n                # save the outputs for not-nan inputs\n                if saves is None:\n                    saves = out[~nan]\n                else:\n                    saves = torch.cat([saves, out[~nan]], dim=0)\n        # [B, T]\n        return saves[:bsize]","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:55:06.616791Z","iopub.execute_input":"2023-10-17T12:55:06.617389Z","iopub.status.idle":"2023-10-17T12:55:06.632437Z","shell.execute_reply.started":"2023-10-17T12:55:06.617358Z","shell.execute_reply":"2023-10-17T12:55:06.631640Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class W2v2Dataset(torch.utils.data.Dataset):\n    def __init__(self, df):\n        self.df = df\n        self.pathes = df['id'].values\n        self.sentences = df['normalized'].values\n        self.resampler = tat.Resample(32000, SR)\n        self.augment = Compose(\n            transforms=[\n                AddColoredNoise(min_snr_in_db = 3.0, max_snr_in_db = 15, p = 0.65),\n                PitchShift(min_transpose_semitones=-4, max_transpose_semitones=4, sample_rate = SR, p=0.65),\n                Gain(min_gain_in_db=-6, max_gain_in_db=6, p=0.65),\n                PolarityInversion(p=0.5)\n            ]\n        )\n                \n        \n    def __getitem__(self, idx):\n        apath = f'/kaggle/input/bengaliai-speech/train_mp3s/{self.pathes[idx]}.mp3'\n        waveform, sample_rate = torchaudio.load(apath, format=\"mp3\")\n        waveform = self.resampler(waveform)\n        \n        # Timbre perturbation\n        waveform = augment(signal=waveform).unsqueeze(0)\n        \n        # Noise + Pitch Shift + Gain + Polarity Inversion\n        waveform = self.augment(waveform, sample_rate=SR)\n        \n        # Time Stretch\n        waveform = time_stretch(waveform, np.random.uniform(0.75, 1.05), SR).squeeze(0)\n        \n\n        batch = dict()\n        y = processor(waveform.reshape(-1), sampling_rate=SR).input_values[0] \n        batch[\"input_values\"] = y\n        with processor.as_target_processor():\n            batch[\"labels\"] = processor(self.sentences[idx]).input_ids       \n        \n        return batch\n\n    def __len__(self):\n        return len(self.df)\n\ntrain_dataset = W2v2Dataset(train)\nvalid_dataset = W2v2Dataset(valid)","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:55:06.633804Z","iopub.execute_input":"2023-10-17T12:55:06.634327Z","iopub.status.idle":"2023-10-17T12:55:06.767819Z","shell.execute_reply.started":"2023-10-17T12:55:06.634297Z","shell.execute_reply":"2023-10-17T12:55:06.766994Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@dataclass\nclass DataCollatorCTCWithPadding:\n    \"\"\"\n    Data collator that will dynamically pad the inputs received.\n    Args:\n        processor (:class:`~transformers.Wav2Vec2Processor`)\n            The processor used for proccessing the data.\n        padding (:obj:`bool`, :obj:`str` or :class:`~transformers.tokenization_utils_base.PaddingStrategy`, `optional`, defaults to :obj:`True`):\n            Select a strategy to pad the returned sequences (according to the model's padding side and padding index)\n            among:\n            * :obj:`True` or :obj:`'longest'`: Pad to the longest sequence in the batch (or no padding if only a single\n              sequence if provided).\n            * :obj:`'max_length'`: Pad to a maximum length specified with the argument :obj:`max_length` or to the\n              maximum acceptable input length for the model if that argument is not provided.\n            * :obj:`False` or :obj:`'do_not_pad'` (default): No padding (i.e., can output a batch with sequences of\n              different lengths).\n        max_length (:obj:`int`, `optional`):\n            Maximum length of the ``input_values`` of the returned list and optionally padding length (see above).\n        max_length_labels (:obj:`int`, `optional`):\n            Maximum length of the ``labels`` returned list and optionally padding length (see above).\n        pad_to_multiple_of (:obj:`int`, `optional`):\n            If set will pad the sequence to a multiple of the provided value.\n            This is especially useful to enable the use of Tensor Cores on NVIDIA hardware with compute capability >=\n            7.5 (Volta).\n    \"\"\"\n\n    processor: Wav2Vec2Processor\n    padding: Union[bool, str] = True\n    max_length: Optional[int] = None\n    max_length_labels: Optional[int] = None\n    pad_to_multiple_of: Optional[int] = None\n    pad_to_multiple_of_labels: Optional[int] = None\n\n    def __call__(self, features: List[Dict[str, Union[List[int], torch.Tensor]]]) -> Dict[str, torch.Tensor]:\n        # split inputs and labels since they have to be of different lenghts and need\n        # different padding methods\n        input_features = [{\"input_values\": feature[\"input_values\"]} for feature in features]\n        label_features = [{\"input_ids\": feature[\"labels\"]} for feature in features]\n        batch = self.processor.pad(\n            input_features,\n            padding=self.padding,\n            max_length=self.max_length,\n            pad_to_multiple_of=self.pad_to_multiple_of,\n            return_tensors=\"pt\",\n        )\n        with self.processor.as_target_processor():\n            labels_batch = self.processor.pad(\n                label_features,\n                padding=self.padding,\n                max_length=self.max_length_labels,\n                pad_to_multiple_of=self.pad_to_multiple_of_labels,\n                return_tensors=\"pt\",\n            )\n\n        # replace padding with -100 to ignore loss correctly\n        labels = labels_batch[\"input_ids\"].masked_fill(labels_batch.attention_mask.ne(1), -100)\n\n        batch[\"labels\"] = labels\n\n        return batch","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:55:06.769162Z","iopub.execute_input":"2023-10-17T12:55:06.769747Z","iopub.status.idle":"2023-10-17T12:55:06.779162Z","shell.execute_reply.started":"2023-10-17T12:55:06.769718Z","shell.execute_reply":"2023-10-17T12:55:06.778246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_collator = DataCollatorCTCWithPadding(processor=processor, padding=True)","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:55:06.780210Z","iopub.execute_input":"2023-10-17T12:55:06.780944Z","iopub.status.idle":"2023-10-17T12:55:06.795109Z","shell.execute_reply.started":"2023-10-17T12:55:06.780914Z","shell.execute_reply":"2023-10-17T12:55:06.794212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- In kaggle notebook, there is an error: **cannot import name 'compute_measures' from 'jiwer' (unknown location)**. But in my local notebook, there is no such error.","metadata":{}},{"cell_type":"code","source":"#wer_metric = load_metric(\"wer\")\n\ndef compute_metrics(pred):\n    pred_logits = pred.predictions\n    pred_ids = np.argmax(pred_logits, axis=-1)\n\n    pred.label_ids[pred.label_ids == -100] = processor.tokenizer.pad_token_id\n\n    pred_str = processor.batch_decode(pred_ids)\n    # we do not want to group tokens when computing the metrics\n    label_str = processor.batch_decode(pred.label_ids, group_tokens=False)\n\n    wer = wer_metric.compute(predictions=pred_str, references=label_str)\n\n    return {\"wer\": wer}","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:55:06.796532Z","iopub.execute_input":"2023-10-17T12:55:06.797188Z","iopub.status.idle":"2023-10-17T12:55:06.806402Z","shell.execute_reply.started":"2023-10-17T12:55:06.797158Z","shell.execute_reply":"2023-10-17T12:55:06.805489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"state_dict = torch.load('../input/w2v2-dataaug-30k/pytorch_model.bin', map_location='cpu')","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:55:06.807638Z","iopub.execute_input":"2023-10-17T12:55:06.808164Z","iopub.status.idle":"2023-10-17T12:55:15.106437Z","shell.execute_reply.started":"2023-10-17T12:55:06.808135Z","shell.execute_reply":"2023-10-17T12:55:15.105451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Wav2Vec2ForCTC.from_pretrained(\n    MODEL_PATH,\n    attention_dropout=0.1,\n    hidden_dropout=0.1,\n    feat_proj_dropout=0.0,\n    mask_time_prob=0.05,\n    layerdrop=0.1,\n    #gradient_checkpointing=True, \n    ctc_loss_reduction=\"mean\", \n    pad_token_id=processor.tokenizer.pad_token_id,\n    vocab_size=len(processor.tokenizer),\n    ctc_zero_infinity=True,\n    diversity_loss_weight=100 \n)","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:55:15.107915Z","iopub.execute_input":"2023-10-17T12:55:15.108255Z","iopub.status.idle":"2023-10-17T12:55:27.561818Z","shell.execute_reply.started":"2023-10-17T12:55:15.108224Z","shell.execute_reply":"2023-10-17T12:55:27.560668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# you can freeze some params\nmodel.freeze_feature_extractor()\nmodel.load_state_dict(state_dict)\n#model.freeze_feature_encoder()","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:55:27.563242Z","iopub.execute_input":"2023-10-17T12:55:27.563790Z","iopub.status.idle":"2023-10-17T12:55:27.809268Z","shell.execute_reply.started":"2023-10-17T12:55:27.563756Z","shell.execute_reply":"2023-10-17T12:55:27.808286Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- As a demo, \"**num_train_epochs**\", \"**eval_steps**\" and \"**early_stopping_patience**\" are set to very small values, you can make them larger.\n- If there is no error about jiwer, you can set **metric_for_best_model**=\"wer\", and remember to set **greater_is_better**=False and use **compute_metrics**.","metadata":{}},{"cell_type":"code","source":"training_args = TrainingArguments(\n    output_dir=output_dir,\n    overwrite_output_dir=True,\n    group_by_length=False,\n    lr_scheduler_type='cosine',\n    weight_decay=0.01,\n    per_device_train_batch_size=4,\n    per_device_eval_batch_size=16,\n    gradient_accumulation_steps=1,\n    evaluation_strategy=\"steps\",\n    save_strategy=\"steps\",\n    max_steps=100000, # you can change to \"num_train_epochs\"\n    fp16=True,\n    save_steps=5000,\n    eval_steps=5000,\n    logging_steps=20,\n    learning_rate=5e-5,\n    warmup_steps=600,\n    save_total_limit=1,\n    load_best_model_at_end=True,\n    #metric_for_best_model=\"wer\",\n    #greater_is_better=False,\n    prediction_loss_only=False,\n    auto_find_batch_size=True,\n    report_to=\"none\"\n)","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:55:27.810418Z","iopub.execute_input":"2023-10-17T12:55:27.810760Z","iopub.status.idle":"2023-10-17T12:55:27.898761Z","shell.execute_reply.started":"2023-10-17T12:55:27.810728Z","shell.execute_reply":"2023-10-17T12:55:27.897858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = Trainer(\n    model=model,\n    data_collator=data_collator,\n    args=training_args,\n    #compute_metrics=compute_metrics,\n    train_dataset=train_dataset,\n    eval_dataset=valid_dataset,\n    tokenizer=processor.feature_extractor,\n    #callbacks=[EarlyStoppingCallback(early_stopping_patience=1)],\n)","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:55:27.899880Z","iopub.execute_input":"2023-10-17T12:55:27.900430Z","iopub.status.idle":"2023-10-17T12:55:32.806279Z","shell.execute_reply.started":"2023-10-17T12:55:27.900396Z","shell.execute_reply":"2023-10-17T12:55:32.805330Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.train()\ntrainer.save_model(output_dir)","metadata":{"execution":{"iopub.status.busy":"2023-10-17T12:55:32.807496Z","iopub.execute_input":"2023-10-17T12:55:32.808083Z","iopub.status.idle":"2023-10-17T19:49:49.122252Z","shell.execute_reply.started":"2023-10-17T12:55:32.808033Z","shell.execute_reply":"2023-10-17T19:49:49.116680Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls","metadata":{"execution":{"iopub.status.busy":"2023-10-17T19:49:49.124276Z","iopub.status.idle":"2023-10-17T19:49:49.125268Z","shell.execute_reply.started":"2023-10-17T19:49:49.125009Z","shell.execute_reply":"2023-10-17T19:49:49.125036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- To improve scores you can: \n    * use different pretrained models\n    * alter the parameters\n    * choose more data\n    * filter data in another way.","metadata":{}}]}