{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader, Subset, random_split\nimport os\nimport torch.utils.data as data\nimport numpy as np\nimport pandas as pd\nimport torch.utils.data as torch_data\nimport csv \nimport time\nimport gc\nimport torch\nimport torchaudio\nimport torch.nn.functional as F\nimport re\nfrom tqdm import tqdm\n\ntorch.backends.cudnn.enabled = True\ntorch.backends.cudnn.benchmark = True\ntorch.backends.cudnn.deterministic = False\n\ntorch.set_float32_matmul_precision('high')","metadata":{"_uuid":"8e44f95c-0df4-4bfa-ac83-f46e66d0d225","_cell_guid":"311630a7-ae24-4fdd-8f2a-3922cffb8ff3","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-03-30T11:58:19.667535Z","iopub.execute_input":"2025-03-30T11:58:19.66786Z","iopub.status.idle":"2025-03-30T11:58:19.673064Z","shell.execute_reply.started":"2025-03-30T11:58:19.667831Z","shell.execute_reply":"2025-03-30T11:58:19.672281Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"comet_api_key = \"dummy_key\"  # Add your comet API key if you have one\nproject_name = \"speechrecognition\"\nexperiment_name = \"speechrecognition-sanskrit-1\"\nexperiment = Experiment(api_key=comet_api_key, project_name=project_name, parse_args=False)\nexperiment.set_name(experiment_name)\nexperiment.display()","metadata":{"_uuid":"15d6a5a0-41b1-4644-9d00-d77ddcee1359","_cell_guid":"4d342698-9cd5-439e-bf34-dffcb6282fad","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-03-29T20:39:27.264179Z","iopub.execute_input":"2025-03-29T20:39:27.264556Z","iopub.status.idle":"2025-03-29T20:39:27.284181Z","shell.execute_reply.started":"2025-03-29T20:39:27.264527Z","shell.execute_reply":"2025-03-29T20:39:27.2831Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Processing","metadata":{"_uuid":"2932f5c8-5b3b-4431-a6a7-55ade20919d9","_cell_guid":"5ca38923-7f39-4b27-82ac-35247d1fec94","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"input_path = \"/kaggle/input/birdclef-2025\"\ntrain_audio_path = os.path.join(input_path, \"/kaggle/input/birdclef-2025/train_audio\" )\ntest_soundscapes = os.path.join(input_path,\"/kaggle/input/birdclef-2025/test_soundscapes\" )\ntrain_soundscapes_path = os.path.join(input_path, \"/kaggle/input/birdclef-2025/train_soundscapes\")\n\ntaxonomy = pd.read_csv(os.path.join(input_path, \"/kaggle/input/birdclef-2025/taxonomy.csv\"))\nmeta_csv = pd.read_csv(os.path.join(input_path, \"/kaggle/input/birdclef-2025/train.csv\"))","metadata":{"_uuid":"9b22a1a7-d7f2-4b4a-aed7-16c797654996","_cell_guid":"57ec0a67-00f4-4ebd-9e38-9c974b47e2c5","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-03-30T11:58:24.702254Z","iopub.execute_input":"2025-03-30T11:58:24.702553Z","iopub.status.idle":"2025-03-30T11:58:24.82056Z","shell.execute_reply.started":"2025-03-30T11:58:24.702526Z","shell.execute_reply":"2025-03-30T11:58:24.819653Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data preprocessing","metadata":{"_uuid":"ee35535f-3912-463b-ba96-2e60ee65bf81","_cell_guid":"296ad4b1-256e-4318-90ce-1c3a19c328f1","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"def preprocess_metadata(df):\n    df['secondary_labels'] = df['secondary_labels'].apply(lambda x : re.findall(r\" '(\\w+)'\", x))\n    df['len_sec_labels'] = df['secondary_labels'].map(len)\n    df['file_path'] = df.apply(lambda row : os.path.join(train_audio_path, row['filename']), axis = 1)\n    return df\n\nmeta_df = preprocess_metadata(meta_csv)\nprint(\"Train Meta Shape:\", meta_df.shape)","metadata":{"_uuid":"15214469-ee91-4eb4-b812-0f9689ae2bad","_cell_guid":"6b180d14-e239-4ac9-81ef-0228c5b361d9","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-03-30T11:58:25.56134Z","iopub.execute_input":"2025-03-30T11:58:25.561687Z","iopub.status.idle":"2025-03-30T11:58:25.78939Z","shell.execute_reply.started":"2025-03-30T11:58:25.561654Z","shell.execute_reply":"2025-03-30T11:58:25.788629Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"meta_df.head().style.background_gradient(cmap='YlOrBr')","metadata":{"_uuid":"839aa83b-f786-4217-b8d1-06ab457d51e8","_cell_guid":"7a5a38cc-b88c-44e9-adf8-edada20e602e","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-03-30T11:58:27.776597Z","iopub.execute_input":"2025-03-30T11:58:27.776923Z","iopub.status.idle":"2025-03-30T11:58:27.793515Z","shell.execute_reply.started":"2025-03-30T11:58:27.776896Z","shell.execute_reply":"2025-03-30T11:58:27.792648Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"taxonomy.head().style.background_gradient(cmap='plasma')","metadata":{"_uuid":"ee269bac-22ad-4f71-8c01-74b3ba949b93","_cell_guid":"17150196-6985-4335-8edc-d3942449fb1b","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-03-30T11:58:28.033949Z","iopub.execute_input":"2025-03-30T11:58:28.034166Z","iopub.status.idle":"2025-03-30T11:58:28.043444Z","shell.execute_reply.started":"2025-03-30T11:58:28.034148Z","shell.execute_reply":"2025-03-30T11:58:28.042608Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"meta_df.info()","metadata":{"_uuid":"65485e88-40d1-449d-a534-f4c82095fa71","_cell_guid":"74299878-085d-4f33-8526-d9be1fe975fe","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-03-30T11:58:28.268592Z","iopub.execute_input":"2025-03-30T11:58:28.268847Z","iopub.status.idle":"2025-03-30T11:58:28.291919Z","shell.execute_reply.started":"2025-03-30T11:58:28.268824Z","shell.execute_reply":"2025-03-30T11:58:28.290827Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"meta_df['primary_label'].unique().size","metadata":{"_uuid":"6d191334-8998-422f-b11b-666f49c345af","_cell_guid":"3a87bf4f-1e00-46ca-8162-d575ae9203b7","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-03-30T11:58:29.074937Z","iopub.execute_input":"2025-03-30T11:58:29.075217Z","iopub.status.idle":"2025-03-30T11:58:29.081436Z","shell.execute_reply.started":"2025-03-30T11:58:29.075193Z","shell.execute_reply":"2025-03-30T11:58:29.080635Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"    \nfrom sklearn.preprocessing import LabelEncoder\nimport torch\n\nclass LabelEncoderWrapper:\n    def __init__(self):\n        self.encoder = LabelEncoder()\n        self.classes_ = []\n        \n    def fit(self, all_labels):\n        self.encoder.fit(all_labels)\n        self.classes_ = self.encoder.classes_\n        \n    def transform(self, labels):\n        return torch.LongTensor(self.encoder.transform(labels))\n    \n    def inverse_transform(self, indices):\n        return self.encoder.inverse_transform(indices.cpu().numpy())\n\n\nclass CustomDataset(torch.utils.data.Dataset):\n    def __init__(self, meta_df, transform=None):\n        self.meta_df = meta_df\n        self.transform = transform\n        self.labels = meta_df.iloc[:, 0].tolist()\n        \n    def __len__(self):\n        return len(self.meta_df)\n    \n    def __getitem__(self, idx):\n        audio_path = str(self.meta_df.iloc[idx, -1])\n        label = self.meta_df.iloc[idx, 0]\n        \n        # Audio loading with format handling\n        try:\n            waveform, sr = torchaudio.load(audio_path, format=\"ogg\" if audio_path.endswith(\".ogg\") else None)\n        except:\n            waveform, sr = torchaudio.load(audio_path+'.wav')\n            \n        # Convert to mono if needed\n        if waveform.shape[0] > 1:\n            waveform = waveform.mean(dim=0, keepdim=True)\n            \n        return waveform, label, sr  # Changed order for better processing\n\n\ndef data_processing(batch, encoder, data_type=\"train\"):\n    waveforms, labels, sample_rates = zip(*batch) \n    # Convert labels first\n    encoded_labels = encoder.transform(labels)\n\n    \n    # feature extraction\n    mel_transform = train_audio_transforms if data_type == \"train\" else valid_audio_transforms\n    mfcc_transform = train_mfcc_transform if data_type == \"train\" else valid_mfcc_transform\n\n    fused_features = []\n    input_lengths = []\n    \n    for waveform in waveforms:\n        # Ensure correct shape: [channel, time]\n        if waveform.dim() == 1:\n            waveform = waveform.unsqueeze(0)\n            \n        # Process both features\n        \n        \n        # Format for processing\n\n        with torch.no_grad():\n            mel = mel_transform(waveform)\n            mfcc = mfcc_transform(waveform)\n            combined = torch.cat([mel, mfcc], dim=1)\n            combined = combined.squeeze(0).permute(1,0)   #[time, n_mels + n_mfcc]\n            \n        fused_features.append(combined)\n        input_lengths.append(combined.size(0))  # Both have same time dimension\n\n    # Dynamic padding for both features\n    padded = nn.utils.rnn.pad_sequence(fused_features, batch_first=True)    # [B, max_time, features]\n    \n    # Final reshape for CNN input\n    padded = padded.permute(0,2,1).unsqueeze(1)  # [batch, 1, n_mels, time]\n    \n    return padded, encoded_labels, torch.LongTensor(input_lengths)\n\n\n\n\n\n# Shared parameters for alignment\nCOMMON_FFT = 512\nCOMMON_HOP = 256\nCOMMON_NMELS = 80\n\ntrain_audio_transforms = nn.Sequential(\n    torchaudio.transforms.MelSpectrogram(sample_rate=16000,\n                                         n_fft=COMMON_FFT,\n                                         hop_length=COMMON_HOP,\n                                         n_mels=COMMON_NMELS),  # n_mel = 128 originally\n    torchaudio.transforms.FrequencyMasking(freq_mask_param=20),          # freq_mask_param originally 30\n    torchaudio.transforms.TimeMasking(time_mask_param=50)                # time_mask_param originally 100\n)\n\ntrain_mfcc_transform = nn.Sequential(\n    torchaudio.transforms.MFCC(sample_rate=16000, n_mfcc=40,\n                               melkwargs={\n                                   \"n_fft\": COMMON_FFT,\n                                   \"hop_length\": COMMON_HOP,\n                                   \"n_mels\":COMMON_NMELS\n                               }),                                       # n_mel = 128 originally\n    torchaudio.transforms.TimeMasking(time_mask_param=30)                # time_mask_param originally 100\n)\n\n# MFCC for training without masking \nvalid_audio_transforms = torchaudio.transforms.MelSpectrogram(sample_rate=16000,\n                                                              n_fft=COMMON_FFT,\n                                                              hop_length=COMMON_HOP,\n                                                              n_mels=COMMON_NMELS)\n\nvalid_mfcc_transform = torchaudio.transforms.MFCC(\n    sample_rate = 16000,\n    n_mfcc = 40,\n    melkwargs={\"n_fft\":COMMON_FFT, \"hop_length\":COMMON_HOP, \"n_mels\":COMMON_NMELS}\n)","metadata":{"_uuid":"a89df417-d94a-4ab7-9757-6cb4d8e35014","_cell_guid":"71f9779b-8979-45c2-a147-ff9c0697294e","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-03-30T11:58:31.788287Z","iopub.execute_input":"2025-03-30T11:58:31.788559Z","iopub.status.idle":"2025-03-30T11:58:31.805088Z","shell.execute_reply.started":"2025-03-30T11:58:31.788538Z","shell.execute_reply":"2025-03-30T11:58:31.804106Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass CNNLayerNorm(nn.Module):\n    def __init__(self, n_feats, norm_dim=\"channel\"):\n        super(CNNLayerNorm, self).__init__()\n        self.norm_dim = norm_dim\n        # nn.LayerNorm expects the normalized shape to match the last dimension\n        self.layer_norm = nn.LayerNorm(n_feats)\n        \n    def forward(self, x):\n        # x: [B, C, F, T]\n        if self.norm_dim == \"channel\":\n            # Bring channel dimension to last\n            x = x.permute(0, 2, 3, 1)  # [B, F, T, C]\n            x = self.layer_norm(x)\n            x = x.permute(0, 3, 1, 2).contiguous()\n        elif self.norm_dim == \"freq\":\n            # Bring frequency dimension (F) to last\n            # x: [B, C, F, T] -> [B, C, T, F]\n            x = x.permute(0, 1, 3, 2)\n            x = self.layer_norm(x)\n            x = x.permute(0, 1, 3, 2).contiguous()\n        else:\n            raise ValueError(\"norm_dim must be 'channel' or 'freq'\")\n        return x\n\n\nclass ResidualCNN(nn.Module):\n    def __init__(self, in_channels, out_channels, kernel, stride, dropout, n_feats):\n        super(ResidualCNN, self).__init__()\n        # Use a tuple stride so that only the time dimension is reduced.\n        self.cnn1 = nn.Conv2d(in_channels, out_channels, kernel, stride=(1, stride), padding=kernel//2)\n        self.cnn2 = nn.Conv2d(out_channels, out_channels, kernel, stride=(1, 1), padding=kernel//2)\n        self.dropout1 = nn.Dropout(dropout)\n        self.dropout2 = nn.Dropout(dropout)\n        self.layer_norm1 = CNNLayerNorm(n_feats, norm_dim=\"freq\")\n        self.layer_norm2 = CNNLayerNorm(n_feats, norm_dim=\"freq\")\n        \n        # If channel dimensions differ, adjust the residual connection.\n        self.residual_downsample = None\n        if in_channels != out_channels:\n            self.residual_downsample = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=(1, stride))\n    \n    def forward(self, x):\n        residual = x\n        if self.residual_downsample:\n            residual = self.residual_downsample(residual)\n        \n        x = self.layer_norm1(x)\n        x = F.gelu(x)\n        x = self.dropout1(x)\n        x = self.cnn1(x)\n        x = self.layer_norm2(x)\n        x = F.gelu(x)\n        x = self.dropout2(x)\n        x = self.cnn2(x)\n        x += residual\n        return x\n\n\nclass AudioModel(nn.Module):\n    def __init__(self, n_classes, n_feats = 120):     # n_feats = n_mels + n_mfcc\n        super().__init__()\n        self.base = nn.Sequential(\n            # Use a tuple stride so frequency dimension stays n_mels\n            nn.Conv2d(1, 32, kernel_size=3, stride=(1, 2), padding=1),  # [B, 32, n_mels, T//2]\n            CNNLayerNorm(32),\n            nn.GELU(),\n            nn.Dropout(0.1),\n            # Since frequency is preserved, we set n_feats=n_mels.\n            ResidualCNN(32, 32, kernel=3, stride=1, dropout=0.1, n_feats=n_feats),\n            ResidualCNN(32, 64, kernel=3, stride=2, dropout=0.1, n_feats=n_feats),\n            nn.AdaptiveAvgPool2d((1, 1))  # Output: [B, 64, 1, 1]\n        )\n        self.classifier = nn.Sequential(\n            nn.Linear(64, 256),\n            nn.GELU(),\n            nn.Dropout(0.1),\n            nn.Linear(256, n_classes)\n        )\n    \n    def forward(self, x):\n        x = self.base(x)  # [B, 64, 1, 1]\n        x = x.view(x.size(0), -1)  # Flatten to [B, 64]\n        return self.classifier(x)\n\n\n\n#(mel + mfcc)\nclass DualInputModel(nn.Module):\n    def __init__(self, n_mels=80, n_mfcc=40, n_classes=10):\n        super().__init__()\n        # Mel branch\n        self.mel_base = nn.Sequential(\n            nn.Conv2d(1, 32, kernel_size=3, stride=(1, 2), padding=1),\n            CNNLayerNorm(32),\n            nn.GELU(),\n            nn.Dropout(0.1),\n            ResidualCNN(32, 32, kernel=3, stride=1, dropout=0.1, n_feats=n_mels),\n            ResidualCNN(32, 64, kernel=3, stride=2, dropout=0.1, n_feats=n_mels),\n            nn.AdaptiveAvgPool2d((1, 1))\n        )\n        \n        # MFCC branch\n        self.mfcc_base = nn.Sequential(\n            nn.Conv2d(1, 32, kernel_size=3, stride=(1, 2), padding=1),\n            CNNLayerNorm(32),\n            nn.GELU(),\n            nn.Dropout(0.1),\n            ResidualCNN(32, 32, kernel=3, stride=1, dropout=0.1, n_feats=n_mfcc),\n            ResidualCNN(32, 64, kernel=3, stride=2, dropout=0.1, n_feats=n_mfcc),\n            nn.AdaptiveAvgPool2d((1, 1))\n        )\n        \n        # Combined classifier\n        self.classifier = nn.Sequential(\n            nn.Linear(128, 256),  # 64 (mel) + 64 (mfcc) = 128\n            nn.GELU(),\n            nn.Dropout(0.1),\n            nn.Linear(256, n_classes)\n        )\n\n    def forward(self, mel_input, mfcc_input):\n        # Process both inputs\n        mel_features = self.mel_base(mel_input).flatten(1)\n        mfcc_features = self.mfcc_base(mfcc_input).flatten(1)\n        \n        # Combine features\n        combined = torch.cat([mel_features, mfcc_features], dim=1)\n        return self.classifier(combined)\n\n\n\n\n\n\n\nclass SpeechRecognitionModelwithStacked(nn.Module):\n    def __init__(self, n_cnn_layers, n_rnn_layers, rnn_dim, n_class, n_feats, stride=2, dropout=0.1):\n        super(SpeechRecognitionModelwithStacked, self).__init__()\n        n_feats = n_feats // 2\n        self.cnn = nn.Conv2d(3, 32, 3, stride=stride, padding=3//2)\n        self.rescnn_layers = nn.Sequential(*[\n            ResidualCNN(32, 32, kernel=3, stride=1, dropout=dropout, n_feats=n_feats)\n            for _ in range(n_cnn_layers)\n        ])\n        self.fully_connected = nn.Linear(n_feats*32, rnn_dim)\n        self.birnn_layers = nn.Sequential(*[\n            BidirectionalGRU(rnn_dim=rnn_dim if i==0 else rnn_dim*2,\n                             hidden_size=rnn_dim, dropout=dropout, batch_first=i==0)\n            for i in range(n_rnn_layers)\n        ])\n        self.classifier = nn.Sequential(\n            nn.Linear(rnn_dim*2, rnn_dim),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(rnn_dim, n_class)\n        )\n        \n    def forward(self, x):\n        # print(f\"Input to model: {x.shape}\")\n        dx = torchaudio.functional.compute_deltas(x)\n        ddx = torchaudio.functional.compute_deltas(dx)\n        \n        x = torch.stack([x, dx, ddx], dim=1)   #dim : [batch_size, 3, 1, n_mels, time_steps]\n        # print(f\"After Stacking: {x.shape}\")\n        x = x.squeeze(2)  # Shape: [batch_size, 3, n_mels, time_steps]\n        x = self.cnn(x)\n        x = self.rescnn_layers(x)\n        sizes = x.size()\n        x = x.view(sizes[0], sizes[1]*sizes[2], sizes[3])\n        x = x.transpose(1, 2)\n        x = self.fully_connected(x)\n        x = self.birnn_layers(x)\n        x = self.classifier(x)\n        return x\n    \n    \n# Comet ML Experiment\n\n\n\n# experiment = Experiment(api_key=comet_api_key, project_name=project_name, parse_args=False)\n# experiment.set_name(experiment_name)\n# experiment.display()","metadata":{"_uuid":"cb6cd335-6289-48c6-bb9d-eb4fc2ca969e","_cell_guid":"12906ed6-64e6-489d-ba69-505ee96e24bf","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-03-30T11:58:32.263179Z","iopub.execute_input":"2025-03-30T11:58:32.263427Z","iopub.status.idle":"2025-03-30T11:58:32.281896Z","shell.execute_reply.started":"2025-03-30T11:58:32.263406Z","shell.execute_reply":"2025-03-30T11:58:32.280813Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize label encoder with all possible labels\nall_labels = meta_df.iloc[:, 0].tolist()\nlabel_encoder = LabelEncoderWrapper()\nlabel_encoder.fit(all_labels)\n\n# Create datasets\nfull_dataset = CustomDataset(meta_df)\ntrain_size = int(0.8 * len(full_dataset))\ntest_size = len(full_dataset) - train_size\ntrain_dataset, test_dataset = random_split(full_dataset, [train_size, test_size])\n\n# Create dataloaders with proper collation\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=4,\n    collate_fn=lambda x: data_processing(x, label_encoder, 'train'),\n    shuffle=True,\n    drop_last = True\n)\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=4,\n    collate_fn=lambda x: data_processing(x, label_encoder, 'valid'),\n    drop_last = True\n)\n\n# Model initialization\nmodel = DualInputModel(n_classes=len(label_encoder.classes_))","metadata":{"_uuid":"cd890779-0b8b-4935-ab09-eb983ba4ea64","_cell_guid":"907ad4d4-6589-444c-a745-a0ac288e5e09","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-03-30T11:58:32.493343Z","iopub.execute_input":"2025-03-30T11:58:32.493595Z","iopub.status.idle":"2025-03-30T11:58:32.515382Z","shell.execute_reply.started":"2025-03-30T11:58:32.493574Z","shell.execute_reply":"2025-03-30T11:58:32.514749Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"i = 0\nfor data in train_loader:\n    i = i + 1\n    mel_input, mfcc_input, labels = data\n    print(f'''\n              Shape of mfcc_input : {mfcc_input}\n              Shape of labels : {labels}\n              Shape of input_lengths: {input_lengths}''')\n\n    if i == 2:\n        break;","metadata":{"_uuid":"51e2afbe-1b26-47e1-bd61-99a8b559fcc4","_cell_guid":"818c4ea0-91c4-47b7-adf0-1519e295d21c","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-03-30T11:58:32.691258Z","iopub.execute_input":"2025-03-30T11:58:32.691568Z","iopub.status.idle":"2025-03-30T11:58:33.001441Z","shell.execute_reply.started":"2025-03-30T11:58:32.691543Z","shell.execute_reply":"2025-03-30T11:58:33.000316Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train(model, device, train_loader, criterion, optimizer, epoch, grad_clip=None):\n    model.train()\n    total_loss = 0\n    correct = 0\n    total = 0\n\n    pbar = tqdm(enumerate(train_loader), total=len(train_loader), desc=f\"Epoch {epoch+1}\")\n\n    for batch_idx, (padded, labels, _) in pbar:\n        padded, labels = padded.to(device), labels.to(device)\n\n        # if model works only with one input\n        # outputs = model(specs)\n        # elif model works with mfcc + mel i.e DualinputModel\n        outputs = model(padded)\n        loss = criterion(outputs, labels)\n\n        optimizer.zero_grad()\n        loss.backward()\n\n        if grad_clip:\n            nn.utils.clip_grad_norm_(model.parameters(), grad_clip)\n\n        optimizer.step()\n\n        total_loss += loss.item()\n        _, predicted = torch.max(outputs.data, 1)\n        total += labels.size(0)\n        correct += (predicted == labels).sum().item()\n        pbar.set_postfix({\n            'Loss :': f\"{loss.item():.4f}\",\n            'Accuracy :': f\"{correct/total:.2f}\"\n        })\n        \n        \n\n    avg_loss = total_loss/ len(train_loader)\n    avg_acc = 100 * correct/total\n    return avg_loss, avg_acc\n\n\ndef test(model, device, test_loader, criterion):\n    model.eval()\n    total_loss = 0\n    correct = 0\n    total = 0\n\n    with torch.no_grad():\n        for batch_idx, (padded, labels, _) in pbar:\n            padded, labels = padded.to(device), labels.to(device)            \n\n            outputs = model(padded)\n            loss = criterion(outputs, labels)\n        \n            total_loss += loss.item()\n            _, predicted = torch.max(outputs, 1)\n            correct += (predicted == labels).sum().item()\n            total += labels.size(0)\n    \n    avg_loss = total_loss / len(test_loader)\n    avg_acc = 100 * correct / total\n    return avg_loss, avg_acc\n\n\ndef main():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    num_epocs = 20\n    grad_clip = 3\n\n    model = AudioModel(n_classes=len(label_encoder.classes_)).to(device)\n    criterion = nn.CrossEntropyLoss(label_smoothing=0.1)\n    \n    optimizer = optim.Adam(model.parameters(), lr=5e-4, weight_decay=1e-5,betas=(.9, .999), amsgrad=True)\n    \n    # scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=3)\n    scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer,\n                                                    max_lr=5e-4,\n                                                    steps_per_epoch=len(train_loader),\n                                                    epochs=num_epocs,\n                                                    pct_start = .3,\n                                                    div_factor=25,\n                                                    final_div_factor=1e4\n    \n                                                   )\n    scaler = torch.cuda.amp.GradScaler() \n    \n    history = {\n        'train_loss': [],\n        'train_acc': [],\n        'test_loss': [],\n        'test_acc': []\n    }\n\n    best_acc = 0.0\n    for epoch in range(num_epocs):\n        train_loss, train_acc = train(model, device, train_loader, criterion, optimizer, epoch, grad_clip)\n        test_loss, test_acc = test(model, device, test_loader, criterion)\n        \n        scheduler.step(test_loss)\n        \n        # Update history\n        history['train_loss'].append(train_loss)\n        history['train_acc'].append(train_acc)\n        history['test_loss'].append(test_loss)\n        history['test_acc'].append(test_acc)\n        \n        # Print epoch summary\n        print(f\"\\nEpoch {epoch+1}/{num_epochs}\")\n        print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.2f}%\")\n        print(f\"Test Loss: {test_loss:.4f} | Test Acc: {test_acc:.2f}%\")\n        print(f\"LR: {optimizer.param_groups[0]['lr']:.2e}\")\n        \n        # Save best model\n        if test_acc > best_acc:\n            best_acc = test_acc\n            torch.save(model.state_dict(), 'best_model.pth')\n            print(\"Saved new best model!\")\n    \n    print(f\"\\nBest Test Accuracy: {best_acc:.2f}%\")\n    return history","metadata":{"_uuid":"1a8c0b88-783a-4552-9408-5e84dbb57075","_cell_guid":"fd940213-03c0-44e5-af1b-b3c0f56a09d4","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-03-30T11:58:41.716371Z","iopub.execute_input":"2025-03-30T11:58:41.716654Z","iopub.status.idle":"2025-03-30T11:58:41.728965Z","shell.execute_reply.started":"2025-03-30T11:58:41.716631Z","shell.execute_reply":"2025-03-30T11:58:41.727893Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"history = main()","metadata":{"_uuid":"d6b3f4c3-bd99-482b-a686-d0968f2b2c26","_cell_guid":"9b8c3f0f-7a2b-4aec-ac98-c9002238f0c8","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-03-30T11:58:46.00499Z","iopub.execute_input":"2025-03-30T11:58:46.005274Z","iopub.status.idle":"2025-03-30T12:19:51.846067Z","shell.execute_reply.started":"2025-03-30T11:58:46.005252Z","shell.execute_reply":"2025-03-30T12:19:51.844797Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Train/Validation Functions\nclass IterMeter(object):\n    def __init__(self):\n        self.val = 0\n    def step(self):\n        self.val += 1\n    def get(self):\n        return self.val\n\n\n\n\n# Initialize the accelerator\naccelerator = Accelerator()\n\ndef train(model, device, train_loader, criterion, optimizer, scheduler, epoch, iter_meter, experiment):\n    model.train()\n    data_len = len(train_loader.dataset)\n    with experiment.train():\n\n        watch(model)\n        for batch_idx, _data in enumerate(train_loader):\n            spectrograms, labels, input_lengths, label_lengths = _data\n            spectrograms, labels = spectrograms.to(device), labels.to(device)\n            \n            optimizer.zero_grad()\n            \n            output = model(spectrograms)\n            output = nn.functional.log_softmax(output, dim=2).transpose(0, 1)\n            # print(output.transpose(0,1).shape)\n            \n            input_lengths_tensor = torch.tensor(input_lengths, dtype=torch.long).to(device)\n            label_lengths_tensor = torch.tensor(label_lengths, dtype=torch.long).to(device)\n\n            # loss = criterion(output, labels, input_lengths, label_lengths)\n            loss = criterion(output, labels, input_lengths_tensor, label_lengths_tensor)\n            loss.backward()\n            # accelerator.backward(loss)\n            \n            optimizer.step()\n            scheduler.step()\n            iter_meter.step()\n            \n            if batch_idx % 10 == 0:\n                print(f'Train Epoch: {epoch} [{batch_idx * len(spectrograms)}/{len(train_loader.dataset)} '\n                    f'({100. * batch_idx / len(train_loader):.0f}%)]\\tLoss: {loss.item():.6f}')\n\n\n\n\n\ndef test(model, device, test_loader, criterion, epoch, iter_meter, experiment, name, tokens_path, lexicon_path, lm_path):\n    print('\\nEvaluating...')\n    model.eval()\n    test_loss = 0\n    test_cer, test_wer = [], []\n    greedy_test_cer, greedy_test_wer = [], []\n    kenlm_test_cer, kenlm_test_wer = [], []\n    beamsearch_cpu_cer, beamsearch_cpu_wer = [], []\n    \n    with experiment.test():\n        with torch.no_grad():\n            for _data in test_loader:\n                spectrograms, labels, input_lengths, label_lengths = _data\n                spectrograms, labels = spectrograms.to(device), labels.to(device)\n                \n                # Get emissions from the model\n                output = model(spectrograms)\n                output = nn.functional.log_softmax(output, dim=2).transpose(0, 1)\n                \n                input_lengths_tensor = torch.tensor(input_lengths, dtype=torch.long).to(device)\n                label_lengths_tensor = torch.tensor(label_lengths, dtype=torch.long).to(device)\n\n                # loss = criterion(output, labels, input_lengths, label_lengths)\n                loss = criterion(output, labels, input_lengths_tensor, label_lengths_tensor)\n                test_loss += loss.item()\n\n                # Use beam search decoder to get the transcript\n                # if epoch % 201 == 0:\n                #     beam_searchcpu_transcripts = beam_search_decoder(output.transpose(0, 1), tokens_path=\"/home/guest2/SanskritASRModel/Dataset/tokens_sanskrit.txt\", lexicon_path=\"/home/guest2/SanskritASRModel/Dataset/lexicon_sanskrit.txt\", lm_path=\"/home/guest2/lm_sanskrit.bin\")\n                \n                beam_search_transcripts = beam_search_decoder_gpu(output.transpose(0, 1), tokens_path=\"SanskritASRModel/Dataset/tokens_sanskrit.txt\")\n                kenlm_transcripts = beam_searchLM_decoder_gpu(output.transpose(0, 1), tokens_path=\"SanskritASRModel/Dataset/tokens_sanskrit.txt\", kenlm_model_path=kenlm_model_path)\n                \n                decoded_preds, decoded_targets = GreedyDecoder(output.transpose(0, 1), labels, label_lengths)\n                \n                # print(f\"labels : {decoded_targets[0]}\\n\")\n                # print(f\"Greedy Decoder Transcript: {decoded_preds[0]}\\n\")\n                # print(f\"Beam Search Transcript: {beam_search_transcripts[0]}\\n\")\n                # print(f\"KenLM Beam Search Transcript: {kenlm_transcripts[0]}\\n\")\n                # print(f\"Beam Search CPU Transcript: {beam_searchcpu_transcripts[0]}\\n\")\n               # decoded_preds = beam_search_transcript \n                 # Assuming labels are available for comparison\n\n                # Calculate CER and WER (functions assumed to be defined elsewhere)\n                for i in range(len(decoded_targets)):\n                    greedy_hypo = decoded_preds[i]\n                    target = decoded_targets[i] \n                    hypothesis_beam = beam_search_transcripts[i]  # Get the corresponding hypothesis\n                    kenlm_hypothesis = kenlm_transcripts[i]\n                    \n                    # if epoch % 201 == 0:\n                    #     hypothesis_beamcpu = beam_searchcpu_transcripts[i]\n                    # print(f\"Target: {target}, hypothetis: {hypothesis_beam}\\n\")\n                    \n                    cer_beam = cer(target, hypothesis_beam, ignore_case=True, remove_space=True)\n                    wer_beam = wer(target, hypothesis_beam, ignore_case=True)\n                    \n                    greedy_cer_beam = cer(target, greedy_hypo, ignore_case=True, remove_space=True)\n                    greedy_wer_beam = wer(target, greedy_hypo, ignore_case=True)\n                    \n                     \n                    kenlm_cer = cer(target, kenlm_hypothesis, ignore_case=True, remove_space=True)\n                    kenlm_wer = wer(target, kenlm_hypothesis, ignore_case=True)\n                    \n                    # if epoch % 201 == 0:\n                    #     beam_cpu_cer = cer(target, hypothesis_beamcpu, ignore_case=True, remove_space=True)\n                    #     beam_cpu_wer = wer(target, hypothesis_beamcpu, ignore_case=True)\n                    \n                    \n                    test_cer.append(cer_beam)\n                    test_wer.append(wer_beam)\n\n                    greedy_test_cer.append(greedy_cer_beam)\n                    greedy_test_wer.append(greedy_wer_beam)\n                    \n                    \n                    kenlm_test_cer.append(kenlm_cer)\n                    kenlm_test_wer.append(kenlm_wer)\n                    \n                    # if epoch % 201 == 0:\n                    #     beamsearch_cpu_cer.append(beam_cpu_cer)\n                    #     beamsearch_cpu_wer.append(beam_cpu_wer)\n    \n    \n                   \n    avg_cer = sum(test_cer) / len(test_cer)\n    avg_wer = sum(test_wer) / len(test_wer)\n    avg_loss = test_loss / len(test_loader)\n    \n    greedy_cer = sum(greedy_test_cer) / len(greedy_test_cer)\n    greedy_wer = sum(greedy_test_wer) / len(greedy_test_wer)\n    \n    \n    kenlm_avg_cer = sum(kenlm_test_cer) / len(kenlm_test_cer)\n    kenlm_avg_wer = sum(kenlm_test_wer) / len(kenlm_test_wer)\n    \n    # if epoch % 201 == 0:\n    #     beamsearch_cpu_avg_cer = sum(beamsearch_cpu_cer) / len(beamsearch_cpu_cer)\n    #     beamsearch_cpu_avg_wer = sum(beamsearch_cpu_wer) / len(beamsearch_cpu_wer)\n        \n    # experiment.log_metric(\"test_loss\", avg_loss, step=iter_meter.get())\n    # experiment.log_metric(\"cer\", avg_cer, step=iter_meter.get())\n    # experiment.log_metric(\"wer\", avg_wer, step=iter_meter.get())\n\n    # print(f'Test set for {name}: Average loss: {avg_loss:.4f}, Average beam search CER: {avg_cer:.4f}, Average beam search WER: {avg_wer:.4f}, Average kenlm WER: {kenlm_avg_cer:.4f},Average kenlm WER: {kenlm_avg_wer:.4f},  Average beam_search_cpu WER: {beamsearch_cpu_avg_cer:.4f},Average beam_search_cpu WER: {beamsearch_cpu_avg_wer:.4f}  \\n')\n    # print(f'Test set for {name}: Average loss: {avg_loss:.4f}, Average beam search CER: {avg_cer:.4f}, Average beam search WER: {avg_wer:.4f}, Average kenlm WER: {kenlm_avg_cer:.4f},Average kenlm WER: {kenlm_avg_wer:.4f} \\n')\n\n    # if epoch % 201 == 0:\n    #     print(f'Test set for {name}: Average loss: {avg_loss:.4f}, \"Averag greedy cer\" {greedy_cer:.4f},  \"Averag greedy wer\" {greedy_wer:.4f}, Average beam search CER: {avg_cer:.4f}, Average beam search WER: {avg_wer:.4f}, Average kenlm WER: {kenlm_avg_cer:.4f},Average kenlm WER: {kenlm_avg_wer:.4f}, Average beam_search_cpu WER: {beamsearch_cpu_avg_cer:.4f},Average beam_search_cpu WER: {beamsearch_cpu_avg_wer:.4f}\\n') \n    # else :\n    print(f'Test set for {name}: Average loss: {avg_loss:.4f}, \"Averag greedy cer\" {greedy_cer:.4f},  \"Averag greedy wer\" {greedy_wer:.4f}, Average beam search CER: {avg_cer:.4f}, Average beam search WER: {avg_wer:.4f}, Average kenlm WER: {kenlm_avg_cer:.4f},Average kenlm WER: {kenlm_avg_wer:.4f}\\n') \n\n    with open(csv_file_path, mode='a', newline='') as file:\n        writer = csv.writer(file)\n        writer.writerow([name])\n        # if epoch % 201 == 0:  \n        #     writer.writerow([epoch, avg_loss, greedy_cer, greedy_wer, avg_cer, avg_wer, kenlm_avg_cer, kenlm_avg_wer, beamsearch_cpu_avg_cer, beamsearch_cpu_avg_wer])\n        # else:\n        writer.writerow([epoch, avg_loss,greedy_cer, greedy_wer, avg_cer, avg_wer, kenlm_avg_cer, kenlm_avg_wer])\n\n\n\n\n\n\ndef main(learning_rate=5e-4, batch_size=10, total_epochs=10, chunk_size=10,\n         train_csv=\"/content/train_cleaned.csv\", test_csv=\"/content/test_cleaned.csv\",\n         train_audio_dir=\"/content/wav_data/train\", test_audio_dir=\"/content/wav_data/test\",\n         experiment=None, csv_file_path=\"ASR_Sanskrit_results.csv\", d1_eval_csv = None ,d1_eval_audio_dir = None, d2_eval_csv = None, d2_eval_audio_dir = None, d2t_eval_csv = None, d2t_eval_audio_dir = None, checkpoint_dir=None):\n\n    hparams = {\n        \"n_cnn_layers\": 3,\n        \"n_rnn_layers\": 5,\n        \"rnn_dim\": 512,\n        \"n_class\": 73,\n        \"n_feats\": 80,\n        \"stride\": 2,\n        \"dropout\": 0.1,\n        \"learning_rate\": learning_rate,\n        \"batch_size\": batch_size,\n        \"epochs\": total_epochs\n    }\n\n    with open(csv_file_path, mode='a', newline='') as file:\n        writer = csv.writer(file)\n        writer.writerow([\"training on D1 without stacked\"])\n        writer.writerow([\"Time\", \"CNNs\", \"GRUs\", \"rnn_dim\", \"n_class\", \"n_feats\", \"stride\", \"dropout\", \"learning_rate\", \"batch_size\", \"epochs\"])\n        writer.writerow([time.time(), hparams['n_cnn_layers'], hparams['n_rnn_layers'], hparams['rnn_dim'], hparams['n_class'],\n                         hparams['n_feats'], hparams['stride'], hparams['dropout'], hparams['learning_rate'], hparams['batch_size'], hparams['epochs']])\n        writer.writerow([\"Epoch\",\"Average Greedy CER\", \"Average Greedy WER\", \"Average beam search CER\", \"Average beam search WER\", \"Average kenlm WER\", \"Average kenlm WER\" , \" Average beam_search_cpu WER\", \"Average beam_search_cpu WER\"])\n\n    experiment.log_parameters(hparams) if experiment else None\n\n    use_cuda = torch.cuda.is_available()\n    torch.manual_seed(7)\n\n    # Prepare datasets and dataloaders\n    train_dataset = CustomSanskritDataset(train_csv, train_audio_dir, data_type='train')\n    test_dataset = CustomSanskritDataset(test_csv, test_audio_dir, data_type='test')\n\n    # Preparing datasets for Evaluation\n    \n    \n    \n    kwargs = {'num_workers': 1, 'pin_memory': True} if use_cuda else {}\n    \n    \n    # If want to use only MFCC features, use data_processing with 'mfcc_train' argument in collate_fn and 'SpeechRecognitionModel' in model\n    train_loader = torch.utils.data.DataLoader(dataset=train_dataset, batch_size=hparams['batch_size'], shuffle=True,\n                                               collate_fn=lambda x: data_processing(x, 'train'), **kwargs)\n    test_loader = torch.utils.data.DataLoader(dataset=test_dataset, batch_size=hparams['batch_size'], shuffle=False,\n                                              collate_fn=lambda x: data_processing(x, 'valid'), **kwargs)\n\n    \n    \n    # Initialize model\n    model = SpeechRecognitionModel(hparams['n_cnn_layers'], hparams['n_rnn_layers'], hparams['rnn_dim'],\n                                    hparams['n_class'], hparams['n_feats'], hparams['stride'], hparams['dropout'])\n    \n    model = torch.nn.DataParallel(model).to(accelerator.device)\n\n    print(model)\n    print('Num Model Parameters:', sum([param.nelement() for param in model.parameters()]))\n\n    criterion = nn.CTCLoss(blank=0).to(accelerator.device)\n    optimizer = optim.AdamW(model.parameters(), hparams['learning_rate'])\n    scheduler = optim.lr_scheduler.OneCycleLR(optimizer, max_lr=hparams['learning_rate'], steps_per_epoch=len(train_loader),\n                                              epochs=hparams['epochs'], anneal_strategy='linear')\n\n   \n\n    # model, optimizer, train_loader, test_loader, scheduler = accelerator.prepare(\n    #     model, optimizer, train_loader, test_loader, scheduler\n    # )\n\n     # Checkpoint Handling\n    CHECKPOINT_DIR = checkpoint_dir\n    os.makedirs(CHECKPOINT_DIR, exist_ok=True)\n    last_epoch = 0\n    checkpoint_path = os.path.join(CHECKPOINT_DIR, \"last_checkpoint.pth\")\n    if os.path.exists(checkpoint_path):\n        checkpoint = torch.load(checkpoint_path)\n        model.load_state_dict(checkpoint['model_state_dict'])\n        optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n        scheduler.load_state_dict(checkpoint['scheduler_state_dict'])\n        last_epoch = checkpoint['epoch']\n        print(f\"Resuming training from epoch {last_epoch + 1}\")\n    \n\n    iter_meter = IterMeter()\n    for start_epoch in range(last_epoch + 1, total_epochs + 1, chunk_size):\n        end_epoch = min(start_epoch + chunk_size - 1, total_epochs)\n\n        for epoch in range(start_epoch, end_epoch + 1):\n            experiment.log_current_epoch(epoch) if experiment else None\n            train(model, accelerator.device, train_loader, criterion, optimizer, scheduler, epoch, iter_meter, experiment)\n            test(model, accelerator.device, test_loader, criterion, epoch, iter_meter, experiment,name=\"D2_SA_BD\", tokens_path=\"SanskritASRModel/Dataset/tokens.txt\", lexicon_path=\"SanskritASRModel/Dataset/lexicon_sanskrit.txt\", lm_path=\"/home/guest2/kenlm/lm_sanskrit.bin\")\n            \n            if epoch%5==0 : \n                gc.collect()\n                torch.cuda.empty_cache()\n                \n        # Save checkpoint\n        checkpoint = {\n            'epoch': end_epoch,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'scheduler_state_dict': scheduler.state_dict()\n        }\n        torch.save(checkpoint, checkpoint_path)\n        print(f\"Checkpoint saved at epoch {end_epoch}\")\n        # del model\n        # del train_loader\n        # del test_loader\n        gc.collect()\n        torch.cuda.empty_cache()\n        # kill_gpu_process()\n    \n        \n        \n            \n            \n        \n        \n        \n        \n        \n    # # Testing on D1, D2 and D2t\n    Eval(model, accelerator.device, d1_eval_loader, criterion, iter_meter, experiment, name=\"D1 Eval\")\n    Eval(model, accelerator.device, d2_eval_loader, criterion, iter_meter, experiment, name=\"D2 Eval\")\n    kfold_eval(model, accelerator.device, d2_eval_dataset, criterion, hparams['batch_size'], 5, iter_meter, experiment, name=\"D2 Eval Kfold\")\n    kfold_eval(model, accelerator.device, d2t_eval_dataset, criterion, hparams['batch_size'], 5, iter_meter, experiment, name=\"D2t Eval kFold \")\n    \n    \n        \n\n    final_model_path = os.path.join(CHECKPOINT_DIR, \"final_model.pth\")\n    accelerator.save_state(final_model_path)\n    print(f\"Final model saved at: {final_model_path}\")\n\n# Parameters\nlearning_rate = 5e-4\nbatch_size = 10\nchunk_size = 5\ntotal_epochs = 200\n\ntrain_csv = \"SanskritASRModel/Dataset/Data_1/train.csv\"\ntest_csv = \"SanskritASRModel/Dataset/Data_1/test.csv\"\ntrain_audio_dir = \"SanskritASRModel/Dataset/sanskrit_dataset/audio\"\ntest_audio_dir = \"SanskritASRModel/Dataset/sanskrit_dataset/audio\"\n\nd1_eval_csv = \"SanskritASRModel/Dataset/Data_1/eval.csv\"\nd1_eval_audio_dir = \"SanskritASRModel/Dataset/sanskrit_dataset/audio\"\nd2_eval_csv = \"SanskritASRModel/Dataset/Dataset_2/eval.csv\"\nd2_eval_audio_dir = \"SanskritASRModel/Dataset/Dataset_2/audio/eval\"\nd2t_eval_csv = \"SanskritASRModel/Dataset/Dataset_2/train.csv\"\nd2t_eval_audio_dir = \"SanskritASRModel/Dataset/Dataset_2/audio/train\"\n\ncsv_file_path = \"SanskritASRModel/results/data_1_SA.csv\"\ncheckpoint_dir = \"SanskritASRModel/model_checkpoints/D1_SA\"\n\nif __name__ == \"__main__\":\n    main(learning_rate, batch_size, total_epochs, chunk_size, train_csv, test_csv, train_audio_dir, test_audio_dir, experiment=experiment, csv_file_path=csv_file_path, d1_eval_csv=\"SanskritASRModel/train_cleaned.csv\" ,d1_eval_audio_dir=\"SanskritASRModel/wav_data/train\", d2_eval_csv=\"SanskritASRModel/Dataset/Dataset_2/eval.csv\", d2_eval_audio_dir=\"SanskritASRModel/Dataset/Dataset_2/audio/eval\", d2t_eval_csv=\"SanskritASRModel/Dataset/Dataset_2/train.csv\", d2t_eval_audio_dir=\"SanskritASRModel/Dataset/Dataset_2/audio/train\" , checkpoint_dir=checkpoint_dir)","metadata":{"_uuid":"2cd880af-0898-4fad-9202-9b11bf6e4fd9","_cell_guid":"5635ae3e-2bba-4818-8020-ec0dc7b6248e","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null}]}