{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.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":129276,"databundleVersionId":15506988,"sourceType":"competition"}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install torchaudio jiwer\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-05T03:35:06.05533Z","iopub.execute_input":"2026-02-05T03:35:06.056083Z","iopub.status.idle":"2026-02-05T03:35:11.43065Z","shell.execute_reply.started":"2026-02-05T03:35:06.05605Z","shell.execute_reply":"2026-02-05T03:35:11.429581Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob, torch, torchaudio, pandas as pd\nimport numpy as np\n\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn.utils.rnn import pad_sequence\nfrom torch.amp import autocast, GradScaler\nfrom sklearn.model_selection import train_test_split\nfrom jiwer import wer\ntorch.backends.cudnn.enabled = False\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"expandable_segments:True\"\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nBASE_PATH = \"/kaggle/input/dl-sprint-4-0-bengali-long-form-speech-recognition/transcription/transcription\"\nTRAIN_AUDIO = f\"{BASE_PATH}/train/audio\"\nTRAIN_TEXT  = f\"{BASE_PATH}/train/annotation\"\nTEST_AUDIO  = f\"{BASE_PATH}/test/audio\"\n\n\ndef load_transcript(path):\n    with open(path, \"r\", encoding=\"utf-8\") as f:\n        return f.read().strip().replace(\"।\", \"\").replace(\",\", \"\").replace(\"?\", \"\")\n\nchars = set()\nfor t in glob.glob(f\"{TRAIN_TEXT}/*.txt\"):\n    chars.update(load_transcript(t))\n\nchars = sorted(list(chars))\nvocab = [\"<blank>\"] + chars\n\nchar2idx = {c: i for i, c in enumerate(vocab)}\nidx2char = {i: c for c, i in char2idx.items()}\n\n\ndef chunk_audio(waveform, sr, chunk_sec=8, overlap_sec=1, max_chunks=3):\n    chunk_size = int(chunk_sec * sr)\n    overlap = int(overlap_sec * sr)\n\n    chunks = []\n    start = 0\n    while start + chunk_size < waveform.shape[-1] and len(chunks) < max_chunks:\n        chunks.append(waveform[:, start:start+chunk_size])\n        start += chunk_size - overlap\n\n    if len(chunks) < max_chunks:\n        chunks.append(waveform[:, start:])\n\n    return chunks[:max_chunks]\n\n\n\nclass BanglaASRDataset(Dataset):\n    def __init__(self, audio_files):\n        self.audio_files = audio_files\n        self.mel = torchaudio.transforms.MelSpectrogram(\n            sample_rate=16000, n_mels=80\n        )\n\n    def __len__(self):\n        return len(self.audio_files)\n\n    def __getitem__(self, idx):\n        wav = self.audio_files[idx]\n        name = os.path.basename(wav).replace(\".wav\", \"\")\n        txt = f\"{TRAIN_TEXT}/{name}.txt\"\n\n        try:\n            waveform, sr = torchaudio.load(wav)\n        except:\n            return None\n\n        waveform = torchaudio.functional.resample(waveform, sr, 16000)\n        chunks = [waveform]   # NO chunking during training\n        features = []\n        for c in chunks:\n            m = self.mel(c)\n            m = torch.log(m + 1e-9)\n            features.append(m.squeeze(0).transpose(0, 1))\n\n        targets = torch.tensor(\n            [char2idx[c] for c in load_transcript(txt)],\n            dtype=torch.long\n        )\n\n        return features, targets\n\n\ndef collate_fn(batch):\n    batch = [b for b in batch if b is not None]\n    if len(batch) == 0:\n        return None\n\n    feats, targets = [], []\n    feat_lens, target_lens = [], []\n\n    for f_list, t in batch:\n        for f in f_list:\n            feats.append(f)\n            feat_lens.append(f.shape[0])\n            targets.append(t)\n            target_lens.append(len(t))\n\n    feats = pad_sequence(feats, batch_first=True)\n    targets = torch.cat(targets)\n\n    return feats, targets, torch.tensor(feat_lens), torch.tensor(target_lens)\n\n\n\nclass ASRModel(nn.Module):\n    def __init__(self, vocab_size):\n        super().__init__()\n        self.cnn = nn.Sequential(\n            nn.Conv1d(80, 192, 3, padding=1),\n            nn.ReLU(),\n            nn.Conv1d(192, 192, 3, padding=1),\n            nn.ReLU(),\n        )\n        self.lstm = nn.LSTM(\n            192, 192, num_layers=3,\n            bidirectional=True, batch_first=True\n        )\n        self.fc = nn.Linear(384, vocab_size)\n\n    def forward(self, x):\n    # CNN in fp16 is OK\n        x = x.transpose(1, 2).contiguous()\n        x = self.cnn(x)\n        x = x.transpose(1, 2).contiguous()\n    \n        # 🚨 FORCE LSTM TO FP32 (CRITICAL)\n        x = x.float()\n        x, _ = self.lstm(x)\n    \n        return self.fc(x)\n\nall_audio = sorted(glob.glob(f\"{TRAIN_AUDIO}/*.wav\"))\n\ntrain_files, val_files = train_test_split(\n    all_audio, test_size=0.1, random_state=42\n)\n\ntrain_ds = BanglaASRDataset(train_files)\nval_ds   = BanglaASRDataset(val_files)\n\ntrain_loader = DataLoader(train_ds, batch_size=1, shuffle=True, collate_fn=collate_fn)\nval_loader   = DataLoader(val_ds, batch_size=1, shuffle=False, collate_fn=collate_fn)\n\n\n\n# ===============================\n# CORRECT CTC TRAINING LOOP\n# ===============================\n\nmodel = ASRModel(len(vocab)).to(device)\n\ncriterion = nn.CTCLoss(blank=0)   # 🚫 no zero_infinity\noptimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)\n\nscaler = GradScaler(\"cuda\")\naccum_steps = 4\nnum_epochs = 20\n\nfor epoch in range(num_epochs):\n    model.train()\n    optimizer.zero_grad()\n\n    total_loss = 0.0\n    num_batches = 0\n\n    for step, batch in enumerate(train_loader):\n        if batch is None:\n            continue\n\n        feats, targets, feat_lens, target_lens = batch\n        feats = feats.to(device)\n        targets = targets.to(device)\n\n        # ---- forward ----\n        with autocast(\"cuda\"):\n            logits = model(feats)                      # (B, T, V)\n            log_probs = logits.log_softmax(2)            # (B, T, V)\n            log_probs = log_probs.transpose(0, 1)        # (T, B, V)\n\n            loss = criterion(\n                log_probs,\n                targets,\n                feat_lens,\n                target_lens\n            )\n\n            loss_scaled = loss / accum_steps\n\n        # ---- backward ----\n        scaler.scale(loss_scaled).backward()\n\n        if (step + 1) % accum_steps == 0:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n\n        total_loss += loss.item()\n        num_batches += 1\n\n    avg_loss = total_loss / max(1, num_batches)\n    print(f\"Epoch {epoch+1}/{num_epochs} | Avg Train Loss: {avg_loss:.2f}\")\n\n\ndef greedy_decode(logits):\n    ids = logits.argmax(dim=-1)\n    prev = -1\n    text = []\n\n    for i in ids[0]:\n        i = i.item()\n        if i != prev and i != 0:\n            text.append(idx2char[i])\n        prev = i\n\n    return \"\".join(text)\n\n\n\nmodel.eval()\nwers = []\n\nwith torch.no_grad():\n    for batch in val_loader:\n        if batch is None:\n            continue\n\n        feats, targets, _, _ = batch\n        feats = feats.to(device)\n\n        logits = model(feats)\n        pred = greedy_decode(logits)\n\n        # ground truth\n        gt = \"\".join([idx2char[i.item()] for i in targets])\n\n        wers.append(wer(gt, pred))\n\nprint(\"Validation WER:\", sum(wers) / len(wers))\n\n\nsubmission = []\nmodel.eval()\n\nwith torch.no_grad():\n    for wav in sorted(glob.glob(f\"{TEST_AUDIO}/*.wav\")):\n        file_id = os.path.basename(wav).replace(\".wav\", \"\")\n\n        try:\n            waveform, sr = torchaudio.load(wav)\n        except:\n            submission.append([file_id, \"\"])\n            continue\n\n        waveform = torchaudio.functional.resample(waveform, sr, 16000)\n        chunks = chunk_audio(waveform, 16000)\n\n        preds = []\n        for c in chunks:\n            m = train_ds.mel(c)\n            m = torch.log(m + 1e-9)\n            m = m.squeeze(0).transpose(0, 1).unsqueeze(0).to(device)\n            logits = model(m)\n            preds.append(greedy_decode(logits))\n\n        submission.append([file_id, \" \".join(preds).strip()])\n\n\ndf = pd.DataFrame(submission, columns=[\"id\", \"transcription\"])\ndf.to_csv(\"submission.csv\", index=False)\ndf.head(32)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-05T04:32:25.504115Z","iopub.execute_input":"2026-02-05T04:32:25.504476Z","iopub.status.idle":"2026-02-05T04:32:33.464Z","shell.execute_reply.started":"2026-02-05T04:32:25.504442Z","shell.execute_reply":"2026-02-05T04:32:33.462724Z"}},"outputs":[],"execution_count":null}]}