{"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":"none","dataSources":[{"sourceType":"competition","sourceId":52324,"databundleVersionId":6229904},{"sourceType":"datasetVersion","sourceId":6334427,"datasetId":3646378,"databundleVersionId":6414868},{"sourceType":"datasetVersion","sourceId":6707460,"datasetId":3865741,"databundleVersionId":6791840},{"sourceType":"datasetVersion","sourceId":14893332,"datasetId":9528796,"databundleVersionId":15757409},{"sourceType":"datasetVersion","sourceId":14852667,"datasetId":2066599,"databundleVersionId":15713052},{"sourceType":"datasetVersion","sourceId":6657244,"datasetId":3694609,"databundleVersionId":6741190},{"sourceType":"datasetVersion","sourceId":4143520,"datasetId":2447262,"databundleVersionId":4200057},{"sourceType":"datasetVersion","sourceId":6686513,"datasetId":3856730,"databundleVersionId":6770803}],"dockerImageVersionId":31259,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from pathlib import Path\nimport torch\nimport os\nimport base64\nimport pandas as pd\nimport io\nimport re","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T22:51:08.897760Z","iopub.execute_input":"2026-02-19T22:51:08.898881Z","iopub.status.idle":"2026-02-19T22:51:08.920310Z","shell.execute_reply.started":"2026-02-19T22:51:08.898772Z","shell.execute_reply":"2026-02-19T22:51:08.918935Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from transformers import pipeline\nimport torch\n\nMODEL = \"/kaggle/input/datasets/tugstugi/bengali-ai-asr-submission/bengali-whisper-medium\"\n\npipe = pipeline(\n    task=\"automatic-speech-recognition\",\n    model=MODEL,\n    tokenizer=MODEL,\n    chunk_length_s=20.1,\n    device=0 if torch.cuda.is_available() else -1,\n    batch_size=1\n)\n\npipe.model.config.forced_decoder_ids = (\n    pipe.tokenizer.get_decoder_prompt_ids(language=\"bn\", task=\"transcribe\")\n)\n\nprint(\"ASR Loaded\")\n\nfrom transformers import AutoModelForTokenClassification, AutoTokenizer\nimport torch\nimport torch.nn.functional as F\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nprint(\"Punctuation Device:\", DEVICE)\n\nPUNCT_MODELS = [\n    '/kaggle/input/datasets/tugstugi/bengali-ai-asr-submission/punct-model-6layers/',\n    '/kaggle/input/datasets/tugstugi/bengali-ai-asr-submission/punct-model-8layers/',\n    '/kaggle/input/datasets/tugstugi/bengali-ai-asr-submission/punct-model-11layers/',\n    '/kaggle/input/datasets/tugstugi/bengali-ai-asr-submission/punct-model-12layers/'\n]\n\npunct_models = [\n    AutoModelForTokenClassification.from_pretrained(f).to(DEVICE).eval()\n    for f in PUNCT_MODELS\n]\n\ntokenizer = AutoTokenizer.from_pretrained(PUNCT_MODELS[0])\n\nPUNCT_WEIGHTS = torch.FloatTensor([[1.0,1.4,1.0,0.8]]).to(DEVICE)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T22:51:08.923383Z","iopub.execute_input":"2026-02-19T22:51:08.924099Z","iopub.status.idle":"2026-02-19T22:51:23.851596Z","shell.execute_reply.started":"2026-02-19T22:51:08.924040Z","shell.execute_reply":"2026-02-19T22:51:23.848540Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\ndef fix_repetition(text, max_count=8):\n    uniq={}\n    words=text.split()\n    for w in words:\n        uniq[w]=uniq.get(w,0)+1\n    for w,c in uniq.items():\n        if c>max_count:\n            words=[x for x in words if x!=w]\n    return \" \".join(words)\n\ndef punctuate(text):\n\n    input_ids = tokenizer(text).input_ids\n\n    with torch.no_grad():\n\n        logits = F.softmax(\n            punct_models[0](\n                input_ids=torch.LongTensor([input_ids]).to(DEVICE)\n            ).logits[0,1:-1],\n            dim=1\n        )\n\n        for model in punct_models[1:]:\n            logits += F.softmax(\n                model(\n                    input_ids=torch.LongTensor([input_ids]).to(DEVICE)\n                ).logits[0,1:-1],\n                dim=1\n            )\n\n        logits = logits / len(punct_models)\n        logits *= PUNCT_WEIGHTS\n\n        label_ids = torch.argmax(logits, dim=-1)\n\n        tokens = tokenizer(text,add_special_tokens=False).input_ids\n\n        punct_text=\"\"\n\n        for i,t in enumerate(tokens):\n\n            tok=tokenizer.decode(t)\n\n            if '##' not in tok:\n                punct_text+=\" \"+tok\n            else:\n                punct_text+=tok[2:]\n\n            punct_text+=['','।',',','?'][label_ids[i].item()]\n\n    punct_text=punct_text.strip()\n\n    if punct_text[-1] not in ['।','? ',',']:\n        punct_text+='।'\n\n    return punct_text\n\ndef write_base64_to_wav(encoded_audio, sample_id):\n    \"\"\"\n    Converts base64 audio to temporary WAV file\n    and returns filepath for Whisper pipeline.\n    \"\"\"\n\n    audio_bytes = base64.b64decode(encoded_audio)\n\n    tmp_path = f\"/kaggle/working/tmp_{sample_id}.wav\"\n\n    with open(tmp_path, \"wb\") as f:\n        f.write(audio_bytes)\n\n    return tmp_path\n\n\n\n############################################\n# -------------- WER -------------------- #\n############################################\n\ndef wer(reference, hypothesis):\n    r = reference.split()\n    h = hypothesis.split()\n\n    d = [[0] * (len(h) + 1) for _ in range(len(r) + 1)]\n\n    for i in range(len(r) + 1):\n        d[i][0] = i\n    for j in range(len(h) + 1):\n        d[0][j] = j\n\n    for i in range(1, len(r) + 1):\n        for j in range(1, len(h) + 1):\n            if r[i - 1] == h[j - 1]:\n                cost = 0\n            else:\n                cost = 1\n            d[i][j] = min(\n                d[i - 1][j] + 1,\n                d[i][j - 1] + 1,\n                d[i - 1][j - 1] + cost,\n            )\n\n    return d[len(r)][len(h)] / max(len(r), 1)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T22:51:23.853803Z","iopub.execute_input":"2026-02-19T22:51:23.854408Z","iopub.status.idle":"2026-02-19T22:51:23.886690Z","shell.execute_reply.started":"2026-02-19T22:51:23.854355Z","shell.execute_reply":"2026-02-19T22:51:23.883641Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ndef score(solution: pd.DataFrame,\n          submission: pd.DataFrame,\n          row_id_column_name: str) -> float:\n    \"\"\"\n    Computes the Mean Word Error Rate (WER) between\n    ground-truth Bengali sentences and ASR transcriptions\n    generated from base64-encoded audio in the submission.\n\n    Expected submission format:\n        id,audio_base64\n\n    Parameters\n    ----------\n    solution : pd.DataFrame\n        Ground truth dataframe containing reference sentences.\n        Must include:\n            - row_id_column_name\n            - 'sentence' column\n\n    submission : pd.DataFrame\n        Submission dataframe containing:\n            - row_id_column_name\n            - 'audio_base64' column (base64-encoded WAV audio)\n\n    row_id_column_name : str\n        Column name used as the unique identifier for each sample.\n\n    Returns\n    -------\n    float\n        Mean Word Error Rate (lower is better).\n        Returns 1.0 if submission is invalid or empty.\n    \"\"\"\n\n    if \"base64_audio\" not in submission.columns:\n        return 1.0\n\n    sub=dict(zip(submission[row_id_column_name],submission[\"base64_audio\"]))\n\n    total=0.0\n    count=0\n\n    for _,row in solution.iterrows():\n\n        if \"Usage\" in solution.columns and row[\"Usage\"]==\"Ignored\":\n            continue\n\n        sid=(row[row_id_column_name])\n\n        if sid not in sub:\n            print(type(sid))\n            continue\n\n        try:\n            wav=write_base64_to_wav(sub[sid],sid)\n\n            text=pipe(\n                wav,\n                generate_kwargs={\"max_length\":260,\"num_beams\":4}\n            )[\"text\"]\n\n            os.remove(wav)\n\n            text=fix_repetition(text)\n            text=punctuate(text)\n\n            print(row['sentence'])\n            print(text)\n\n            total+=wer(row[\"sentence\"],text)\n            count+=1\n\n        except Exception as e:\n            print(\"FAIL:\",e)\n            continue\n\n    if count==0:\n        return 1.0\n\n    return float(total/count)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T22:51:23.889569Z","iopub.execute_input":"2026-02-19T22:51:23.890110Z","iopub.status.idle":"2026-02-19T22:51:23.932194Z","shell.execute_reply.started":"2026-02-19T22:51:23.890071Z","shell.execute_reply":"2026-02-19T22:51:23.930364Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T22:51:55.510579Z","iopub.execute_input":"2026-02-19T22:51:55.511891Z","iopub.status.idle":"2026-02-19T22:53:21.011070Z","shell.execute_reply.started":"2026-02-19T22:51:55.511802Z","shell.execute_reply":"2026-02-19T22:53:21.009837Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}