{"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":[{"sourceType":"competition","sourceId":129276,"databundleVersionId":15506988,"isSourceIdPinned":false}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ============================================================================\n# BUET DL Sprint 4.0 — Bangla Long-Form ASR | FINAL MERGED | v11\n# Target: WER ≤ 0.2030 | Kaggle T4 x1 | ~5-6 hours\n# ============================================================================\n#\n# MERGED FROM:\n#   - User demo code  : alignment filter (words/sec), conservative LR, clean targets\n#   - v10             : suppress_tokens crash fix\n#   - Fixed notebook  : GenerationConfig trainer fix (root cause of WER=0.93)\n#\n# KEY DESIGN DECISIONS (senior DL reasoning):\n#   1. BASE MODEL  : bengaliAI/tugstugi (already Bengali) → head-start over openai/whisper\n#   2. ALIGNMENT   : Skip files where words/sec > 5 → prevents label-audio mismatch\n#   3. LR = 8e-6   : tugstugi is pre-trained; high LR destroys prior knowledge\n#   4. LORA r=32   : r=48/64 caused loss=0.003 = overfit; r=32 w/ dropout=0.05 is right\n#   5. GEN CONFIG  : Explicitly set on trainer.model AFTER trainer construction\n#   6. EVAL BEAMS=1: Greedy eval avoids suppress_tokens[-2] crash entirely\n#   7. WHOLE-FILE  : Train on full files, not chunks → correct label alignment\n#   8. STEPS=1500  : With 8e-6 LR, 1500 steps converges without memorising\n#\n# EXPECTED TRAINING LOG (healthy):\n#   Step 200: WER ~0.50-0.65  (language forcing working → Bengali output)\n#   Step 600: WER ~0.32-0.42\n#   Step 1000: WER ~0.24-0.32\n#   Step 1500: WER ~0.20-0.25\n#   If step 200 WER > 0.90 → language forcing is broken, check [LANG CHECK] output\n#   If step 200 WER < 0.10 → overfitting, reduce LR or add dropout\n# ============================================================================\n\n# ─────────────────────────────────────────────────────────────────────────────\n# CELL 1 · INSTALL\n# ─────────────────────────────────────────────────────────────────────────────\nimport subprocess, sys\n\ndef pip(*args):\n    subprocess.check_call([sys.executable, \"-m\", \"pip\", \"install\", \"-q\", *args])\n\npip(\"transformers>=4.44.0,<4.46.0\", \"accelerate>=0.30.0\")\npip(\"jiwer>=3.0.3\", \"torchaudio>=2.1.0\", \"soundfile\", \"librosa\")\nprint(\"✓ Packages ready\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:38.581575Z","iopub.execute_input":"2026-02-19T19:37:38.582337Z","iopub.status.idle":"2026-02-19T19:37:45.069952Z","shell.execute_reply.started":"2026-02-19T19:37:38.582302Z","shell.execute_reply":"2026-02-19T19:37:45.069099Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 2 · ENV\n# ─────────────────────────────────────────────────────────────────────────────\nimport os\nos.environ[\"CUDA_VISIBLE_DEVICES\"]    = \"0\"\nos.environ[\"TOKENIZERS_PARALLELISM\"]  = \"false\"\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"expandable_segments:True\"\nos.environ[\"HF_HOME\"]                 = \"/kaggle/working/hf_cache\"\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.071239Z","iopub.execute_input":"2026-02-19T19:37:45.071467Z","iopub.status.idle":"2026-02-19T19:37:45.075583Z","shell.execute_reply.started":"2026-02-19T19:37:45.071447Z","shell.execute_reply":"2026-02-19T19:37:45.075055Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 3 · IMPORTS\n# ─────────────────────────────────────────────────────────────────────────────\nimport gc, re, time, unicodedata, warnings, random\nfrom collections import Counter\nfrom dataclasses import dataclass, field\nfrom pathlib import Path\nfrom typing import Dict, List, Optional, Tuple\n\nimport numpy as np\nimport pandas as pd\nimport soundfile as sf\nimport torch\nimport torch.nn as nn\nimport torchaudio\nfrom torch.utils.data import Dataset\nfrom tqdm.auto import tqdm\n\nimport jiwer\nfrom transformers import (\n    EarlyStoppingCallback, GenerationConfig,\n    Seq2SeqTrainer, Seq2SeqTrainingArguments,\n    WhisperForConditionalGeneration, WhisperProcessor,\n)\n\nwarnings.filterwarnings(\"ignore\")\nprint(\"✓ Imports OK\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.076342Z","iopub.execute_input":"2026-02-19T19:37:45.076556Z","iopub.status.idle":"2026-02-19T19:37:45.092244Z","shell.execute_reply.started":"2026-02-19T19:37:45.076537Z","shell.execute_reply":"2026-02-19T19:37:45.091612Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 4 · CONFIG\n# ─────────────────────────────────────────────────────────────────────────────\n@dataclass\nclass CFG:\n    model_name : str = \"bengaliAI/tugstugi_bengaliai-asr_whisper-medium\"\n    language   : str = \"bengali\"\n    lang_code  : str = \"bn\"\n    task       : str = \"transcribe\"\n\n    train_audio : str = (\n        \"/kaggle/input/competitions/dl-sprint-4-0-bengali-long-form-speech-recognition\"\n        \"/transcription/transcription/train/audio\"\n    )\n    train_ann : str = (\n        \"/kaggle/input/competitions/dl-sprint-4-0-bengali-long-form-speech-recognition\"\n        \"/transcription/transcription/train/annotation\"\n    )\n    test_audio : str = (\n        \"/kaggle/input/competitions/dl-sprint-4-0-bengali-long-form-speech-recognition\"\n        \"/transcription/transcription/test/audio\"\n    )\n\n    lora_r       : int   = 32\n    lora_alpha   : int   = 64\n    lora_dropout : float = 0.05\n    lora_targets : List[str] = field(default_factory=lambda: [\n        \"q_proj\", \"v_proj\", \"k_proj\", \"out_proj\"\n    ])\n\n    lr         : float = 8e-6\n    steps      : int   = 1500\n    warmup     : int   = 150\n    batch_size : int   = 4\n    grad_accum : int   = 8\n    weight_decay : float = 0.01\n    ft_eval    : int   = 200\n\n    max_audio    : float = 30.0\n    min_dur      : float = 0.5\n    max_wps      : float = 5.0\n\n    val_ratio : float = 0.10\n    max_val   : int   = 100\n\n    chunk_dur    : float = 24.0\n    chunk_hop    : float = 20.0\n    beam_size    : int   = 5\n    no_repeat_ngram : int = 3\n    rep_penalty  : float = 1.2\n    gen_maxlen   : int   = 180\n\n    grad_ckpt   : bool = True\n    num_workers : int  = 2\n    seed        : int  = 42\n\n    out_dir    : str = \"/kaggle/working/finetuned\"\n    out_merged : str = \"/kaggle/working/merged_model\"\n    submission : str = \"/kaggle/working/submission.csv\"\n    cache      : str = \"/kaggle/working/infer_cache.csv\"\n\n\nC = CFG()\n\ndef set_seed(s: int):\n    random.seed(s); np.random.seed(s); torch.manual_seed(s)\n    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)\n\nset_seed(C.seed)\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nif DEVICE == \"cuda\":\n    vram = torch.cuda.get_device_properties(0).total_memory / 1e9\n    cap  = torch.cuda.get_device_capability()\n    if cap[0] >= 8:\n        TRAIN_DTYPE  = torch.bfloat16\n        INFER_DTYPE  = torch.bfloat16\n        TRAINER_FP16 = False\n        TRAINER_BF16 = True\n    else:\n        TRAIN_DTYPE  = torch.float32\n        INFER_DTYPE  = torch.float16\n        TRAINER_FP16 = True\n        TRAINER_BF16 = False\n\n    print(f\"GPU  : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM : {vram:.1f}GB  sm_{cap[0]}{cap[1]}\")\n\n    if vram < 13:\n        C.batch_size = 2; C.grad_accum = 16\n        print(\"⚠  Low VRAM → batch=2, accum=16\")\n    elif vram < 16:\n        C.batch_size = 3; C.grad_accum = 11\n        print(\"✓  Mid VRAM → batch=3, accum=11\")\n\n    torch.backends.cuda.matmul.allow_tf32 = True\n    torch.backends.cudnn.allow_tf32       = True\nelse:\n    TRAIN_DTYPE = INFER_DTYPE = torch.float32\n    TRAINER_FP16 = TRAINER_BF16 = False\n    print(\"⚠  CPU mode\")\n\nprint(f\"✓ Effective batch: {C.batch_size}×{C.grad_accum}={C.batch_size*C.grad_accum}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.09368Z","iopub.execute_input":"2026-02-19T19:37:45.094101Z","iopub.status.idle":"2026-02-19T19:37:45.116694Z","shell.execute_reply.started":"2026-02-19T19:37:45.094074Z","shell.execute_reply":"2026-02-19T19:37:45.116106Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 5 · SUPPRESS_TOKENS PATCH  (v10 definitive fix)\n# Prevents: IndexError: index -2 out of bounds at every eval step\n# Root cause: tugstugi ships with generation_config.suppress_tokens = []\n# ─────────────────────────────────────────────────────────────────────────────\ndef patch_suppress_tokens(model: WhisperForConditionalGeneration,\n                           processor: WhisperProcessor) -> WhisperForConditionalGeneration:\n    gen_cfg = model.generation_config\n\n    st = getattr(gen_cfg, \"suppress_tokens\", None)\n    if st is not None and len(st) < 2:\n        gen_cfg.suppress_tokens = None          # ← key fix\n        print(f\"  [PATCH] suppress_tokens {st!r} → None\")\n\n    if getattr(gen_cfg, \"prev_sot_token_id\", None) is None:\n        try:\n            sot = processor.tokenizer.convert_tokens_to_ids(\"<|startoftranscript|>\")\n            gen_cfg.prev_sot_token_id = sot if (sot and sot > 0) else 50258\n        except:\n            gen_cfg.prev_sot_token_id = 50258\n        print(f\"  [PATCH] prev_sot_token_id = {gen_cfg.prev_sot_token_id}\")\n\n    # Verify no crash will happen\n    st_a = getattr(gen_cfg, \"suppress_tokens\", None)\n    assert st_a is None or len(st_a) >= 2, f\"Patch failed: suppress_tokens={st_a!r}\"\n    print(f\"  [PATCH] ✓ Safe\")\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.117429Z","iopub.execute_input":"2026-02-19T19:37:45.117611Z","iopub.status.idle":"2026-02-19T19:37:45.130448Z","shell.execute_reply.started":"2026-02-19T19:37:45.117594Z","shell.execute_reply":"2026-02-19T19:37:45.129935Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 6 · NORMALIZER  (merged: user's clean version + NFC + zero-width strip)\n# ─────────────────────────────────────────────────────────────────────────────\nclass BengaliNorm:\n    _ZW    = re.compile(r'[\\u200b-\\u200d\\ufeff\\u200e\\u200f\\u00ad]')\n    _PUNCT = re.compile(r'[।,;:!?\\-–—.\\'\\\"\\(\\)\\[\\]{}\\u0964\\u0965\\u2018\\u2019\\u201c\\u201d]')\n    _DIGIT = re.compile(r'[0-9০-৯]+')\n    _ASCII = re.compile(r'[A-Za-z]+')\n    _SP    = re.compile(r'\\s+')\n\n    def __call__(self, text: str) -> str:\n        if not isinstance(text, str) or not text: return \"\"\n        text = unicodedata.normalize(\"NFC\", text)\n        text = self._ZW.sub(\"\", text)\n        text = self._PUNCT.sub(\" \", text)\n        text = self._DIGIT.sub(\" \", text)\n        text = self._ASCII.sub(lambda m: m.group(0).lower(), text)\n        return self._SP.sub(\" \", text).strip()\n\nnorm = BengaliNorm()   # ← fix: added ()\nprint(\"✓ Normalizer OK\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.13114Z","iopub.execute_input":"2026-02-19T19:37:45.131381Z","iopub.status.idle":"2026-02-19T19:37:45.148085Z","shell.execute_reply.started":"2026-02-19T19:37:45.13135Z","shell.execute_reply":"2026-02-19T19:37:45.147425Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 7 · AUDIO UTILS\n# ─────────────────────────────────────────────────────────────────────────────\ndef get_duration(path: str) -> float:\n    try:   return sf.info(str(path)).duration\n    except: return 0.0\n\ndef load_wav(path: str, max_dur: float = 30.0) -> np.ndarray:\n    \"\"\"Load + resample to 16kHz, truncate to max_dur.\"\"\"\n    try:\n        info = sf.info(str(path))\n        max_f = int(max_dur * info.samplerate)\n        with sf.SoundFile(str(path)) as f:\n            wav = f.read(min(max_f, info.frames), dtype=\"float32\")\n        if wav.ndim == 2: wav = wav.mean(axis=1)\n        if info.samplerate != 16_000:\n            wav = torchaudio.functional.resample(\n                torch.from_numpy(wav).unsqueeze(0),\n                info.samplerate, 16_000).squeeze(0).numpy()\n        return wav.astype(np.float32)\n    except:\n        return np.zeros(int(16_000 * 3), dtype=np.float32)\n\ndef load_wav_slice(path: str, start: float, end: float) -> np.ndarray:\n    \"\"\"Efficient seek-based slice — no full-file load.\"\"\"\n    try:\n        info = sf.info(str(path))\n        sr   = info.samplerate\n        fs   = int(start * sr)\n        fe   = min(int(end * sr), info.frames)\n        with sf.SoundFile(str(path)) as f:\n            f.seek(fs)\n            wav = f.read(max(1, fe - fs), dtype=\"float32\")\n        if wav.ndim == 2: wav = wav.mean(axis=1)\n        if sr != 16_000:\n            wav = torchaudio.functional.resample(\n                torch.from_numpy(wav).unsqueeze(0), sr, 16_000).squeeze(0).numpy()\n        return wav.astype(np.float32)\n    except:\n        return np.zeros(int((end - start) * 16_000), dtype=np.float32)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.149151Z","iopub.execute_input":"2026-02-19T19:37:45.149508Z","iopub.status.idle":"2026-02-19T19:37:45.167856Z","shell.execute_reply.started":"2026-02-19T19:37:45.149487Z","shell.execute_reply":"2026-02-19T19:37:45.167291Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 8 · DATA ALIGNMENT ENGINE\n# ─────────────────────────────────────────────────────────────────────────────\ndef load_data() -> Tuple[List[Dict], List[Dict]]:\n    audio_dir = Path(C.train_audio)\n    ann_dir   = Path(C.train_ann)\n    if not audio_dir.exists():\n        for root, dirs, _ in os.walk(\"/kaggle/input\"):\n            if \"audio\" in dirs and \"annotation\" in dirs:\n                audio_dir = Path(root) / \"audio\"\n                ann_dir   = Path(root) / \"annotation\"\n                print(f\"  [AUTO] Found at {audio_dir.parent}\")\n                break\n    wav_paths = sorted(audio_dir.glob(\"*.wav\"))\n    print(f\"  Found {len(wav_paths)} wav files\")\n    sample_durs = [get_duration(str(w)) for w in wav_paths[:min(50, len(wav_paths))]]\n    if sample_durs:\n        print(f\"  Duration stats (first {len(sample_durs)} files): \"\n              f\"min={min(sample_durs):.1f}s  max={max(sample_durs):.1f}s  \"\n              f\"mean={np.mean(sample_durs):.1f}s\")\n    good, skipped_ann, skipped_align, skipped_dur = [], 0, 0, 0\n    for wp in tqdm(wav_paths, desc=\"  Validating alignment\"):\n        ann = ann_dir / f\"{wp.stem}.txt\"\n        if not ann.exists():\n            skipped_ann += 1; continue\n        text = norm(ann.read_text(encoding=\"utf-8\").strip())\n        if not text or len(text.split()) < 2:\n            skipped_ann += 1; continue\n        dur = get_duration(str(wp))\n        if dur < C.min_dur:  # ← removed \"or dur > 300\"\n            skipped_dur += 1; continue\n        n_words = len(text.split())\n        if dur > 0 and (n_words / dur) > C.max_wps:\n            skipped_align += 1  # keep the sample\n        good.append({\n            \"path\": str(wp),\n            \"text\": text,\n            \"dur\":  dur,  # ← keep full duration\n        })\n    print(f\"  ✓ Usable: {len(good)} | \"\n          f\"No-ann: {skipped_ann} | \"\n          f\"High-wps (kept): {skipped_align} | \"\n          f\"Bad-dur: {skipped_dur}\")\n    if len(good) == 0:\n        raise RuntimeError(\"No usable training samples! Check paths and alignment.\")\n    rng = np.random.default_rng(99)\n    idx = rng.permutation(len(good)).tolist()\n    nv  = min(C.max_val, max(4, int(len(good) * C.val_ratio)))\n    train = [good[i] for i in idx[nv:]]\n    val   = [good[i] for i in idx[:nv]]\n    print(f\"  Split → train: {len(train)}  val: {len(val)}\")\n    return train, val","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.168761Z","iopub.execute_input":"2026-02-19T19:37:45.16906Z","iopub.status.idle":"2026-02-19T19:37:45.186205Z","shell.execute_reply.started":"2026-02-19T19:37:45.169032Z","shell.execute_reply":"2026-02-19T19:37:45.185602Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 9 · SPECAUGMENT + DATASET\n# ─────────────────────────────────────────────────────────────────────────────\nclass SpecAugment:\n    def __call__(self, spec: torch.Tensor) -> torch.Tensor:\n        nm, nf = spec.shape\n        for _ in range(2):\n            f = random.randint(0, 15)\n            f0 = random.randint(0, max(0, nm - f))\n            spec[f0:f0+f, :] = 0\n        for _ in range(2):\n            t = random.randint(0, 35)\n            t0 = random.randint(0, max(0, nf - t))\n            spec[:, t0:t0+t] = 0\n        return spec\n\n\nclass ASRDataset(Dataset):\n    def __init__(self, samples: List[Dict], fe, tok, augment: bool = False):\n        self.samples = samples\n        self.fe      = fe\n        self.tok     = tok\n        self.sa      = SpecAugment() if augment else None\n\n    def __len__(self): return len(self.samples)\n\n    def __getitem__(self, idx: int) -> Dict:\n        s   = self.samples[idx]\n        wav = load_wav(s[\"path\"], max_dur=C.max_audio)\n        if len(wav) < 100: wav = np.zeros(16_000 * 3, dtype=np.float32)\n        feat = self.fe(wav, sampling_rate=16_000, return_tensors=\"pt\").input_features[0]\n        if self.sa is not None: feat = self.sa(feat)\n        labs = self.tok(\n            s[\"text\"], return_tensors=\"pt\",\n            truncation=True, max_length=C.gen_maxlen,\n        ).input_ids[0]\n        return {\"input_features\": feat, \"labels\": labs}\n\n\nclass WhisperCollator:\n    def __init__(self, fe, tok):\n        self.fe = fe; self.tok = tok\n\n    def __call__(self, batch: List[Dict]) -> Dict:\n        feats = self.fe.pad(\n            [{\"input_features\": b[\"input_features\"]} for b in batch],\n            return_tensors=\"pt\",\n        )[\"input_features\"]\n        max_len = max(b[\"labels\"].size(0) for b in batch)\n        labels  = torch.full((len(batch), max_len), -100, dtype=torch.long)\n        for i, b in enumerate(batch):\n            llen = b[\"labels\"].size(0)\n            labels[i, :llen] = b[\"labels\"]\n        return {\"input_features\": feats, \"labels\": labels}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.187066Z","iopub.execute_input":"2026-02-19T19:37:45.187328Z","iopub.status.idle":"2026-02-19T19:37:45.203772Z","shell.execute_reply.started":"2026-02-19T19:37:45.187301Z","shell.execute_reply":"2026-02-19T19:37:45.20311Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 10 · WER METRIC  (with [LANG CHECK] diagnostic)\n# ─────────────────────────────────────────────────────────────────────────────\ndef make_compute_metrics(tokenizer):\n    _n = BengaliNorm()\n\n    def compute_metrics(pred):\n        pred_ids  = pred.predictions\n        label_ids = pred.label_ids.copy()\n        label_ids[label_ids == -100] = tokenizer.pad_token_id\n\n        preds = [_n(s) for s in tokenizer.batch_decode(pred_ids,  skip_special_tokens=True)]\n        refs  = [_n(s) for s in tokenizer.batch_decode(label_ids, skip_special_tokens=True)]\n\n        # ── DIAGNOSTIC: print first 2 pairs to detect language issues ──────\n        print(\"\\n  [LANG CHECK]\")\n        for r, p in list(zip(refs, preds))[:2]:\n            print(f\"    REF: {r[:80]}\")\n            print(f\"    HYP: {p[:80]}\")\n        # ────────────────────────────────────────────────────────────────────\n\n        pairs = [(p, r) for p, r in zip(preds, refs) if r.strip()]\n        if not pairs: return {\"wer\": 1.0}\n        ps, rs = zip(*pairs)\n        wer_val = jiwer.wer(list(rs), list(ps))\n        print(f\"  [WER] {wer_val:.4f}\\n\")\n        return {\"wer\": round(wer_val, 4)}\n\n    return compute_metrics\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.205986Z","iopub.execute_input":"2026-02-19T19:37:45.206244Z","iopub.status.idle":"2026-02-19T19:37:45.219354Z","shell.execute_reply.started":"2026-02-19T19:37:45.206224Z","shell.execute_reply":"2026-02-19T19:37:45.218687Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 11 · CUSTOM LORA\n# ─────────────────────────────────────────────────────────────────────────────\nclass LoRALinear(nn.Module):\n    def __init__(self, base: nn.Linear, r: int, alpha: int, dropout: float):\n        super().__init__()\n        self.base   = base\n        self.lora_A = nn.Linear(base.in_features, r,                bias=False)\n        self.lora_B = nn.Linear(r,                base.out_features, bias=False)\n        self.scale  = alpha / r\n        self.drop   = nn.Dropout(dropout)\n        nn.init.normal_(self.lora_A.weight, std=0.02)\n        nn.init.zeros_(self.lora_B.weight)\n        for p in self.base.parameters(): p.requires_grad = False\n\n    def forward(self, x):\n        return self.base(x) + self.lora_B(self.lora_A(self.drop(x))) * self.scale\n\n    def merge(self) -> nn.Linear:\n        with torch.no_grad():\n            W  = self.base.weight.float()\n            dW = (self.lora_B.weight.float() @ self.lora_A.weight.float()) * self.scale\n        m = nn.Linear(self.base.in_features, self.base.out_features,\n                      bias=(self.base.bias is not None))\n        m.weight.data = (W + dW).to(self.base.weight.dtype)\n        if self.base.bias is not None:\n            m.bias.data = self.base.bias.data.clone()\n        return m\n\n\ndef inject_lora(model: nn.Module) -> nn.Module:\n    targets = set(C.lora_targets)\n    for p in model.parameters(): p.requires_grad = False\n    n = 0\n    for parent in model.modules():\n        for name, child in list(parent.named_children()):\n            if isinstance(child, nn.Linear) and name in targets:\n                setattr(parent, name, LoRALinear(child, C.lora_r, C.lora_alpha, C.lora_dropout))\n                n += 1\n    trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    total     = sum(p.numel() for p in model.parameters())\n    print(f\"  LoRA: {n} layers | {trainable:,}/{total:,} ({100*trainable/total:.2f}%)\")\n    return model\n\n\ndef merge_lora(model: nn.Module) -> nn.Module:\n    for parent in model.modules():\n        for name, child in list(parent.named_children()):\n            if isinstance(child, LoRALinear):\n                setattr(parent, name, child.merge())\n    return model\n\n\nprint(\"✓ LoRA OK\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.220207Z","iopub.execute_input":"2026-02-19T19:37:45.220421Z","iopub.status.idle":"2026-02-19T19:37:45.232919Z","shell.execute_reply.started":"2026-02-19T19:37:45.220391Z","shell.execute_reply":"2026-02-19T19:37:45.232161Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 12 · TRAINING\n# ─────────────────────────────────────────────────────────────────────────────\ndef train() -> Tuple:\n    t0 = time.time()\n    set_seed(C.seed)\n    os.makedirs(C.out_dir, exist_ok=True)\n\n    print(\"=\" * 65)\n    print(\"  v11 TRAINING — target WER ≤ 0.2030\")\n    print(\"=\" * 65)\n\n    print(\"\\n[1/5] Processor …\")\n    processor = WhisperProcessor.from_pretrained(\n        C.model_name, language=C.language, task=C.task)\n    processor.tokenizer.set_prefix_tokens(language=C.lang_code, task=C.task)\n    fe, tok = processor.feature_extractor, processor.tokenizer\n\n    forced_ids = processor.get_decoder_prompt_ids(language=C.lang_code, task=C.task)\n    print(f\"  Forced decoder ids: {forced_ids[:3]} …\")\n\n    print(\"\\n[2/5] Loading data …\")\n    train_data, val_data = load_data()\n    ds_train = ASRDataset(train_data, fe, tok, augment=True)\n    ds_val   = ASRDataset(val_data,   fe, tok, augment=False)\n    collator = WhisperCollator(fe, tok)\n    metrics  = make_compute_metrics(tok)\n\n    print(f\"\\n[3/5] Loading {C.model_name} …\")\n    model = WhisperForConditionalGeneration.from_pretrained(\n        C.model_name, torch_dtype=TRAIN_DTYPE, low_cpu_mem_usage=True)\n\n    model.config.use_cache       = False\n    model.config.suppress_tokens = []\n\n    model.config.forced_decoder_ids            = forced_ids\n    model.generation_config.forced_decoder_ids = forced_ids\n\n    model = patch_suppress_tokens(model, processor)\n\n    if C.grad_ckpt:\n        model.enable_input_require_grads()\n        model.gradient_checkpointing_enable()\n\n    model = inject_lora(model)\n\n    for p in model.parameters():\n        if p.requires_grad and p.dtype not in (TRAIN_DTYPE, torch.int64, torch.int32, torch.bool):\n            p.data = p.data.to(TRAIN_DTYPE)\n\n    model.to(DEVICE)\n\n    _s = ds_train[0]\n    with torch.no_grad():\n        _o = model(\n            input_features=_s[\"input_features\"].unsqueeze(0).to(DEVICE),\n            labels=_s[\"labels\"].unsqueeze(0).where(\n                _s[\"labels\"].unsqueeze(0) != tok.pad_token_id,\n                torch.full_like(_s[\"labels\"].unsqueeze(0), -100)\n            ).to(DEVICE),\n        )\n    loss_val = _o.loss.item()\n    assert loss_val == loss_val, \"NaN loss in sanity check!\"\n    print(f\"  ✓ Sanity loss: {loss_val:.4f}\")\n    if loss_val < 0.3:\n        print(\"  ⚠  Very low sanity loss — possible data leakage!\")\n    if loss_val > 8.0:\n        print(\"  ⚠  Very high sanity loss — model may not have loaded correctly!\")\n    del _s, _o; gc.collect(); torch.cuda.empty_cache()\n\n    print(f\"\\n[4/5] Training …\")\n    print(f\"  LR={C.lr}  steps={C.steps}  LoRA r={C.lora_r}  \"\n          f\"batch={C.batch_size}×{C.grad_accum}={C.batch_size*C.grad_accum} eff\")\n\n    train_args = Seq2SeqTrainingArguments(\n        output_dir                    = C.out_dir,\n        per_device_train_batch_size   = C.batch_size,\n        per_device_eval_batch_size    = max(1, C.batch_size),\n        gradient_accumulation_steps   = C.grad_accum,\n        learning_rate                 = C.lr,\n        warmup_steps                  = C.warmup,\n        max_steps                     = C.steps,\n        lr_scheduler_type             = \"cosine\",\n        weight_decay                  = C.weight_decay,\n        max_grad_norm                 = 1.0,\n        gradient_checkpointing        = C.grad_ckpt,\n        fp16                          = TRAINER_FP16,\n        bf16                          = TRAINER_BF16,\n        evaluation_strategy           = \"steps\",\n        eval_steps                    = C.ft_eval,\n        save_strategy                 = \"steps\",\n        save_steps                    = C.ft_eval,\n        load_best_model_at_end        = True,\n        metric_for_best_model         = \"wer\",\n        greater_is_better             = False,\n        predict_with_generate         = True,\n        generation_max_length         = C.gen_maxlen,\n        generation_num_beams          = 1,\n        logging_steps                 = 20,\n        report_to                     = [\"none\"],\n        save_total_limit              = 2,\n        dataloader_num_workers        = C.num_workers,\n        dataloader_pin_memory         = False,\n        remove_unused_columns         = False,\n        label_names                   = [\"labels\"],\n        push_to_hub                   = False,\n    )\n\n    trainer = Seq2SeqTrainer(\n        model           = model,\n        args            = train_args,\n        train_dataset   = ds_train,\n        eval_dataset    = ds_val,\n        tokenizer       = fe,\n        data_collator   = collator,\n        compute_metrics = metrics,\n        callbacks       = [EarlyStoppingCallback(early_stopping_patience=4)],\n    )\n\n    # ── FIXED: added decoder_start_token_id and bos_token_id ──\n    sot_id = processor.tokenizer.convert_tokens_to_ids(\"<|startoftranscript|>\")\n    trainer_gen_cfg = GenerationConfig(\n        decoder_start_token_id = sot_id,\n        bos_token_id           = sot_id,\n        forced_decoder_ids     = forced_ids,\n        max_new_tokens         = C.gen_maxlen,\n        num_beams              = 1,\n        no_repeat_ngram_size   = C.no_repeat_ngram,\n        suppress_tokens        = None,\n    )\n    trainer.model.generation_config                  = trainer_gen_cfg\n    trainer.model.config.forced_decoder_ids          = forced_ids\n    trainer.model.config.decoder_start_token_id      = sot_id\n    # ──────────────────────────────────────────────────────────\n\n    t_train = time.time()\n    trainer.train()\n\n    best_wer   = trainer.state.best_metric or 1.0\n    train_mins = (time.time() - t_train) / 60\n    print(f\"\\n  ✓ Training done | Best val WER: {best_wer:.4f} | {train_mins:.1f} min\")\n\n    if best_wer > 0.5:\n        print(\"\\n  ══ DIAGNOSIS ══\")\n        print(\"  If [LANG CHECK] showed English → language forcing still broken\")\n        print(\"  If [LANG CHECK] showed Bengali but bad WER → need more steps\")\n        print(\"  If [LANG CHECK] showed empty strings → generation failing\")\n        print(\"  Action: scroll up and read the [LANG CHECK] outputs at eval steps.\")\n\n    model = trainer.model\n    del trainer, ds_train, ds_val\n    gc.collect(); torch.cuda.empty_cache()\n\n    print(\"\\n[5/5] Merging LoRA weights …\")\n    model.cpu()\n    merged = merge_lora(model)\n    os.makedirs(C.out_merged, exist_ok=True)\n    merged.save_pretrained(C.out_merged, safe_serialization=True)\n    processor.save_pretrained(C.out_merged)\n\n    total_mins = (time.time() - t0) / 60\n    print(f\"  ✓ Merged model → {C.out_merged}\")\n    print(f\"  Total time: {total_mins:.1f} min ({total_mins/60:.1f} hr)\")\n    print(f\"  Best WER:   {best_wer:.4f}\")\n    return merged, processor","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.23376Z","iopub.execute_input":"2026-02-19T19:37:45.234092Z","iopub.status.idle":"2026-02-19T19:37:45.253879Z","shell.execute_reply.started":"2026-02-19T19:37:45.234066Z","shell.execute_reply":"2026-02-19T19:37:45.253004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 13 · HALLUCINATION DETECTOR\n# ─────────────────────────────────────────────────────────────────────────────\ndef is_hallucination(text: str) -> bool:\n    if not text: return True\n    words = text.split()\n    if len(words) < 4: return False\n    cnt = Counter(words)\n    if cnt.most_common(1)[0][1] / len(words) > 0.60: return True\n    bg  = [f\"{words[i]} {words[i+1]}\" for i in range(len(words) - 1)]\n    if bg and Counter(bg).most_common(1)[0][1] / len(bg) > 0.55: return True\n    return False\n\ndef postprocess(text: str) -> str:\n    if not text: return \"\"\n    text = unicodedata.normalize(\"NFC\", text)\n    text = re.sub(r'\\s+', ' ', text).strip()\n    text = re.sub(r'\\b(\\S+)\\s+\\1\\b', r'\\1', text)          # dedup single words\n    text = re.sub(r'(\\S+\\s+\\S+)\\s+\\1', r'\\1', text)        # dedup bigrams\n    text = re.sub(r'\\s*\\[.?\\]\\s', ' ', text).strip()      # remove [NOISE] etc\n    return text\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.255109Z","iopub.execute_input":"2026-02-19T19:37:45.25541Z","iopub.status.idle":"2026-02-19T19:37:45.27035Z","shell.execute_reply.started":"2026-02-19T19:37:45.255381Z","shell.execute_reply":"2026-02-19T19:37:45.269569Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 14 · INFERENCE  (sliding window + hallucination fallback)\n# ─────────────────────────────────────────────────────────────────────────────\ndef transcribe_file(wav_path: str, model, processor) -> str:\n    try:    total = sf.info(wav_path).duration\n    except: return \"\"\n    if total < 0.5: return \"\"\n\n    fe     = processor.feature_extractor\n    forced = processor.get_decoder_prompt_ids(language=C.lang_code, task=C.task)\n\n    model.eval()\n    parts, start = [], 0.0\n\n    with torch.no_grad():\n        while start < total:\n            end = min(start + C.chunk_dur, total)\n            wav = load_wav_slice(wav_path, start, end)\n            if len(wav) < 1600: break\n\n            feat = fe(wav, sampling_rate=16_000, return_tensors=\"pt\")\\\n                     .input_features.to(DEVICE, dtype=INFER_DTYPE)\n\n            # Primary: beam search\n            ids = model.generate(\n                feat,\n                forced_decoder_ids   = forced,\n                num_beams            = C.beam_size,\n                no_repeat_ngram_size = C.no_repeat_ngram,\n                repetition_penalty   = C.rep_penalty,\n                max_new_tokens       = C.gen_maxlen,\n            )\n            txt = processor.batch_decode(ids, skip_special_tokens=True)[0].strip()\n\n            # Fallback: temperature sampling if hallucination detected\n            if not txt or is_hallucination(txt):\n                ids2 = model.generate(\n                    feat, forced_decoder_ids=forced,\n                    do_sample=True, temperature=0.2, num_beams=1,\n                    no_repeat_ngram_size=C.no_repeat_ngram,\n                    repetition_penalty=C.rep_penalty,\n                    max_new_tokens=C.gen_maxlen,\n                )\n                txt2 = processor.batch_decode(ids2, skip_special_tokens=True)[0].strip()\n                if txt2 and not is_hallucination(txt2): txt = txt2\n\n            if txt and not is_hallucination(txt):\n                parts.append(txt)\n\n            if end >= total: break\n            start += C.chunk_hop\n\n    if not parts: return \"\"\n    if len(parts) == 1: return postprocess(parts[0])\n\n    # Merge with bigram boundary deduplication\n    merged = [parts[0]]\n    for curr in parts[1:]:\n        pw = merged[-1].split(); cw = curr.split()\n        ovlp = 0\n        for n in range(min(10, len(pw), len(cw)), 0, -1):\n            if pw[-n:] == cw[:n]: ovlp = n; break\n        rest = cw[ovlp:]\n        if rest: merged.append(\" \".join(rest))\n\n    return postprocess(\" \".join(merged))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.271399Z","iopub.execute_input":"2026-02-19T19:37:45.272045Z","iopub.status.idle":"2026-02-19T19:37:45.286692Z","shell.execute_reply.started":"2026-02-19T19:37:45.272024Z","shell.execute_reply":"2026-02-19T19:37:45.28608Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 15 · GENERATE SUBMISSION\n# ─────────────────────────────────────────────────────────────────────────────\ndef find_test_dir() -> Path:\n    p = Path(C.test_audio)\n    if p.exists() and list(p.glob(\"*.wav\")): return p\n    for root, dirs, _ in os.walk(\"/kaggle/input\"):\n        for d in dirs:\n            if d == \"audio\":\n                cand = Path(root) / d\n                if list(cand.glob(\"*.wav\")): return cand\n    raise FileNotFoundError(\"Cannot find test audio directory!\")\n\n\ndef generate_submission(model, processor) -> pd.DataFrame:\n    test_dir  = find_test_dir()\n    wav_files = sorted(test_dir.glob(\"*.wav\"))\n    print(f\"\\n  Transcribing {len(wav_files)} test files …\")\n\n    model = model.to(DEVICE).to(INFER_DTYPE)\n    model.eval()\n\n    # Resume from cache\n    done: Dict[str, str] = {}\n    if Path(C.cache).exists():\n        try:\n            c    = pd.read_csv(C.cache)\n            done = dict(zip(c[\"filename\"].astype(str),\n                            c[\"transcription\"].fillna(\"\").astype(str)))\n            print(f\"  Resumed: {len(done)} cached entries\")\n        except: pass\n\n    rows = []\n    for wf in tqdm(wav_files, desc=\"  Transcribing\"):\n        fid = wf.stem\n        if fid in done:\n            rows.append({\"filename\": fid, \"transcription\": done[fid]}); continue\n        try:\n            txt = transcribe_file(str(wf), model, processor)\n        except Exception as e:\n            print(f\"  ⚠ {wf.name}: {e}\"); txt = \"\"\n\n        rows.append({\"filename\": fid, \"transcription\": txt})\n        print(f\"  {fid}: {txt[:80]}\")\n\n        if len(rows) % 5 == 0:\n            pd.DataFrame(rows).to_csv(C.cache, index=False, encoding=\"utf-8-sig\")\n\n    df = pd.DataFrame(rows)\n    df.to_csv(C.cache,      index=False, encoding=\"utf-8-sig\")\n    df.to_csv(C.submission, index=False, encoding=\"utf-8-sig\")\n\n    # Verify submission format\n    assert list(df.columns) == [\"filename\", \"transcription\"], \\\n        f\"Wrong columns: {list(df.columns)}\"\n    for _, row in df.iterrows():\n        txt = str(row[\"transcription\"])\n        assert unicodedata.normalize(\"NFC\", txt) == txt, f\"Not NFC: {row['filename']}\"\n\n    print(f\"\\n  ✓ {len(df)} rows saved → {C.submission}\")\n    print(f\"  Columns: {list(df.columns)}\")\n    return df\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.287524Z","iopub.execute_input":"2026-02-19T19:37:45.287812Z","iopub.status.idle":"2026-02-19T19:37:45.300459Z","shell.execute_reply.started":"2026-02-19T19:37:45.28779Z","shell.execute_reply":"2026-02-19T19:37:45.29988Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nfrom pathlib import Path\n\ntrain_base = \"/kaggle/input/competitions/dl-sprint-4-0-bengali-long-form-speech-recognition/transcription/transcription/train\"\n\nprint(f\"Train base exists: {Path(train_base).exists()}\")\nprint(\"\\nContents:\")\nfor root, dirs, files in os.walk(train_base):\n    print(f\"\\n📁 {root}\")\n    for f in files[:10]:\n        print(f\"   {f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.301322Z","iopub.execute_input":"2026-02-19T19:37:45.301553Z","iopub.status.idle":"2026-02-19T19:37:45.529184Z","shell.execute_reply.started":"2026-02-19T19:37:45.301529Z","shell.execute_reply":"2026-02-19T19:37:45.528623Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pathlib import Path\nimport os\n\naudio_dir = Path(\"/kaggle/input/competitions/dl-sprint-4-0-bengali-long-form-speech-recognition/transcription/transcription/train/audio\")\nann_dir   = Path(\"/kaggle/input/competitions/dl-sprint-4-0-bengali-long-form-speech-recognition/transcription/transcription/train/annotation\")\n\nwav_paths = sorted(audio_dir.glob(\"*.wav\"))\nprint(f\"WAV files found: {len(wav_paths)}\")\n\nfor wp in wav_paths[:5]:\n    ann = ann_dir / f\"{wp.stem}.txt\"\n    dur = get_duration(str(wp))\n    exists = ann.exists()\n    if exists:\n        text = norm(ann.read_text(encoding=\"utf-8\").strip())\n        words = len(text.split())\n        wps = words/dur if dur > 0 else 0\n        print(f\"{wp.stem} | dur={dur:.1f}s | words={words} | wps={wps:.1f} | ann=✓\")\n    else:\n        print(f\"{wp.stem} | ann=✗ MISSING\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.529908Z","iopub.execute_input":"2026-02-19T19:37:45.530166Z","iopub.status.idle":"2026-02-19T19:37:45.585666Z","shell.execute_reply.started":"2026-02-19T19:37:45.530145Z","shell.execute_reply":"2026-02-19T19:37:45.585181Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 16 · RUN  ← fixed: no __name__ guard needed in Kaggle\n# ─────────────────────────────────────────────────────────────────────────────\nprint(\"=\" * 65)\nprint(\"  PRE-FLIGHT\")\nprint(\"=\" * 65)\n\nchecks = [\n    (\"CUDA available\",     torch.cuda.is_available()),\n    (\"Train audio exists\", Path(C.train_audio).exists() or\n                            any(True for _ in Path(\"/kaggle/input\").rglob(\"*.wav\"))),\n]\nfor name, ok in checks:\n    print(f\"  [{'✓' if ok else '✗'}] {name}\")\n\nprint(\"\\n  Verifying crash patch …\")\n_m = WhisperForConditionalGeneration.from_pretrained(\n    C.model_name, torch_dtype=TRAIN_DTYPE, low_cpu_mem_usage=True)\n_p = WhisperProcessor.from_pretrained(C.model_name)\nprint(f\"  suppress_tokens BEFORE patch: {getattr(_m.generation_config, 'suppress_tokens', None)!r}\")\n_m = patch_suppress_tokens(_m, _p)\nprint(f\"  suppress_tokens AFTER  patch: {getattr(_m.generation_config, 'suppress_tokens', None)!r}\")\ndel _m, _p; gc.collect()\nif torch.cuda.is_available(): torch.cuda.empty_cache()\n\nprint(f\"\\n  Config summary:\")\nprint(f\"    LR={C.lr}  steps={C.steps}  LoRA r={C.lora_r}\")\nprint(f\"    max_wps_filter={C.max_wps}  max_audio={C.max_audio}s\")\nprint(f\"    batch={C.batch_size}×{C.grad_accum}={C.batch_size*C.grad_accum} eff\")\nprint(f\"    Eval every {C.ft_eval} steps\")\nprint()\n\n# ── TRAIN ──\nmodel, processor = train()\n\n# ── INFER ──\nsub = generate_submission(model, processor)\n\nprint(\"\\n\" + \"=\" * 65)\nprint(\"  COMPLETE\")\nprint(f\"  Submit: {C.submission}\")\nprint(\"=\" * 65)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T19:37:45.586434Z","iopub.execute_input":"2026-02-19T19:37:45.586676Z","iopub.status.idle":"2026-02-19T20:29:52.200563Z","shell.execute_reply.started":"2026-02-19T19:37:45.586647Z","shell.execute_reply":"2026-02-19T20:29:52.19941Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# CELL 17 · INFERENCE-ONLY MODE (if model already trained)\n# ─────────────────────────────────────────────────────────────────────────────\n\"\"\"\n# Run this block alone if you already have a merged model saved:\n\nMODEL_DIR = \"/kaggle/working/merged_model\"\nprocessor = WhisperProcessor.from_pretrained(MODEL_DIR)\nmodel     = WhisperForConditionalGeneration.from_pretrained(MODEL_DIR)\nsub       = generate_submission(model, processor)\nprint(sub.head())\n\"\"\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T20:29:52.201605Z","iopub.status.idle":"2026-02-19T20:29:52.20186Z","shell.execute_reply.started":"2026-02-19T20:29:52.201747Z","shell.execute_reply":"2026-02-19T20:29:52.201762Z"}},"outputs":[],"execution_count":null}]}