{"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":37174,"databundleVersionId":3938797,"sourceType":"competition"},{"sourceId":9812,"sourceType":"datasetVersion","datasetId":5793},{"sourceId":14592101,"sourceType":"datasetVersion","datasetId":9320929}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q transformers torch torchaudio scikit-learn jiwer\n\nfrom typing import Dict, List, Tuple, Optional, Union\nimport os\nimport re\nimport json\nimport random\nimport time\nimport gc\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torchaudio\nimport torchaudio.transforms as T\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom transformers import Wav2Vec2Processor, Wav2Vec2ForCTC, TrainingArguments, Trainer, EarlyStoppingCallback\nfrom dataclasses import dataclass, field\nfrom sklearn.model_selection import train_test_split\nfrom jiwer import wer, cer\n\nos.environ[\"WANDB_DISABLED\"] = \"true\"\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\n\nprint(f\"PyTorch: {torch.__version__}\")\nprint(f\"CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")\n\n# =========================\n# CONFIGURATION (FROZEN ENCODER + HEAD-ONLY)\n# =========================\nclass Config:\n    # Unified dataset (pre-created)\n    UNIFIED_CSV = \"/kaggle/input/unified/unified_asr_dataset.csv\"\n    \n    # Output paths\n    OUTPUT_DIR = \"/kaggle/working/asr_output_v2\"\n    MODEL_DIR = os.path.join(OUTPUT_DIR, \"best_model\")\n    CHECKPOINT_DIR = os.path.join(OUTPUT_DIR, \"checkpoints\")\n    \n    # Audio parameters\n    SAMPLE_RATE = 16000\n    MAX_AUDIO_LENGTH = 20.0\n    MIN_AUDIO_LENGTH = 0.5\n    \n    # Dataset parameters: 2 Lakh (1L en + 1L bn)\n    TOTAL_DATASET_SIZE = 200000  # 2 Lakh\n    EN_SAMPLES = 100000  # 1 Lakh English\n    BN_SAMPLES = 100000  # 1 Lakh Bengali\n    \n    # Dataset split (70:15:15)\n    TRAIN_RATIO = 0.70\n    VAL_RATIO = 0.15\n    TEST_RATIO = 0.15\n    \n    # Training (LOW GPU/RAM OPTIMIZED)\n    BATCH_SIZE = 4\n    GRADIENT_ACCUMULATION = 4\n    NUM_EPOCHS = 1\n    LEARNING_RATE = 5e-5  # Slightly higher for head-only training\n    WARMUP_STEPS = 500\n    SAVE_STEPS = 1000\n    EVAL_STEPS = 1000\n    LOGGING_STEPS = 100\n    EARLY_STOPPING_PATIENCE = 2\n    \n    # Model settings - facebook/wav2vec2-xls-r-300m (frozen encoder)\n    PRETRAINED_MODEL = \"facebook/wav2vec2-xls-r-300m\"\n    FREEZE_ENCODER = True  # Freeze feature extractor + wav2vec2\n    \n    # Dropout settings (for unfrozen CTC head)\n    ATTENTION_DROPOUT = 0.05\n    HIDDEN_DROPOUT = 0.05\n    FEAT_PROJ_DROPOUT = 0.05\n    MASK_TIME_PROB = 0.05\n    \n    # Device settings\n    DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    FP16 = torch.cuda.is_available()\n    GRADIENT_CHECKPOINTING = True\n    \n    @staticmethod\n    def setup_directories():\n        for dir_path in [Config.OUTPUT_DIR, Config.CHECKPOINT_DIR]:\n            os.makedirs(dir_path, exist_ok=True)\n        print(f\"✓ Output directories ready\")\n\nConfig.setup_directories()\nprint(f\"\\n🔧 Configuration:\")\nprint(f\"   Model: {Config.PRETRAINED_MODEL}\")\nprint(f\"   Frozen Encoder: {Config.FREEZE_ENCODER}\")\nprint(f\"   Dataset Size: {Config.TOTAL_DATASET_SIZE} (EN: {Config.EN_SAMPLES}, BN: {Config.BN_SAMPLES})\")\nprint(f\"   Batch Size: {Config.BATCH_SIZE}, Gradient Accumulation: {Config.GRADIENT_ACCUMULATION}\")\nprint(f\"   Effective Batch Size: {Config.BATCH_SIZE * Config.GRADIENT_ACCUMULATION}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T13:58:51.508823Z","iopub.execute_input":"2026-01-23T13:58:51.509172Z","iopub.status.idle":"2026-01-23T13:58:54.796518Z","shell.execute_reply.started":"2026-01-23T13:58:51.509145Z","shell.execute_reply":"2026-01-23T13:58:54.795742Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# LOAD AND FILTER DATASET (2 LAKH SAMPLES)\n# =========================\nprint(f\"Loading unified dataset from {Config.UNIFIED_CSV}...\")\ndf = pd.read_csv(Config.UNIFIED_CSV)\n\nprint(f\"Initial samples: {len(df)}\")\nprint(f\"Columns: {df.columns.tolist()}\")\nprint(f\"\\nLanguage distribution (before filtering):\")\nprint(df['language'].value_counts())\n\n# Filter: 1 Lakh English + 1 Lakh Bengali\nprint(f\"\\n📊 Filtering dataset to 2 Lakh (1L en + 1L bn)...\")\n\n# Sample exactly 100K English and 100K Bengali\ndf_en = df[df['language'] == 'en'].sample(n=min(Config.EN_SAMPLES, len(df[df['language'] == 'en'])), random_state=42)\ndf_bn = df[df['language'] == 'bn'].sample(n=min(Config.BN_SAMPLES, len(df[df['language'] == 'bn'])), random_state=42)\n\n# Combine and shuffle\ndf = pd.concat([df_en, df_bn]).reset_index(drop=True)\ndf = df.sample(frac=1, random_state=42).reset_index(drop=True)\n\nprint(f\"\\n✓ After filtering:\")\nprint(f\"   Total samples: {len(df)}\")\nprint(f\"   Language distribution:\")\nprint(f\"   {df['language'].value_counts().to_dict()}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T13:58:54.798348Z","iopub.execute_input":"2026-01-23T13:58:54.7987Z","iopub.status.idle":"2026-01-23T13:58:56.326364Z","shell.execute_reply.started":"2026-01-23T13:58:54.798668Z","shell.execute_reply":"2026-01-23T13:58:56.325642Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# DATA PREPROCESSING AND VALIDATION\n# =========================\nclass TextCleaner:\n    @staticmethod\n    def clean_text(text: str) -> str:\n        if not isinstance(text, str) or pd.isna(text):\n            return \"\"\n        \n        # Remove URLs and emails\n        text = re.sub(r'http\\S+|www\\S+|\\S+@\\S+', '', text)\n        # Remove multiple spaces\n        text = re.sub(r'\\s+', ' ', text).strip()\n        \n        return text\n    \n    @staticmethod\n    def is_valid_text(text: str) -> bool:\n        if not text or len(text.strip()) < 3:\n            return False\n        if len(text) > 300:\n            return False\n        return True\n\n# FAST: Validate audio files using parallel processing\nprint(\"Validating audio files (parallel processing)...\")\ninitial_len = len(df)\n\nfrom multiprocessing import Pool\n\ndef check_file_exists(path):\n    return os.path.exists(path)\n\n# Use multiprocessing for faster file validation\nwith Pool(processes=8) as pool:\n    file_exists = pool.map(check_file_exists, df[\"audio_path\"], chunksize=1000)\n\ndf = df[[f for f in file_exists]].reset_index(drop=True)\nprint(f\"After validation: {len(df)} ({len(df)/initial_len*100:.1f}%)\")\n\n# Clean text\nprint(\"\\nCleaning text...\")\ndf[\"text\"] = df[\"text\"].apply(TextCleaner.clean_text)\ndf = df[df[\"text\"].apply(TextCleaner.is_valid_text)].reset_index(drop=True)\nprint(f\"After text cleaning: {len(df)} samples\")\n\n# Remove duplicates\nprint(\"\\nRemoving duplicates...\")\ndf = df.drop_duplicates(subset=['text']).reset_index(drop=True)\nprint(f\"After removing duplicates: {len(df)} samples\")\n\nprint(f\"\\n✓ Final dataset: {len(df)} samples\")\nprint(f\"Languages: {df['language'].value_counts().to_dict()}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T13:58:56.327467Z","iopub.execute_input":"2026-01-23T13:58:56.327869Z","iopub.status.idle":"2026-01-23T13:59:22.354623Z","shell.execute_reply.started":"2026-01-23T13:58:56.327833Z","shell.execute_reply":"2026-01-23T13:59:22.353668Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# SPLIT DATASET STRATIFIED BY LANGUAGE\n# =========================\nprint(f\"Splitting dataset {Config.TRAIN_RATIO*100}:{Config.VAL_RATIO*100}:{Config.TEST_RATIO*100}...\")\n\n# Stratified by language\ntrain_val, test = train_test_split(\n    df, \n    test_size=Config.TEST_RATIO, \n    random_state=42,\n    stratify=df['language']\n)\n\nval_ratio_adjusted = Config.VAL_RATIO / (Config.TRAIN_RATIO + Config.VAL_RATIO)\ntrain, val = train_test_split(\n    train_val, \n    test_size=val_ratio_adjusted, \n    random_state=42,\n    stratify=train_val['language']\n)\n\ntrain = train.reset_index(drop=True)\nval = val.reset_index(drop=True)\ntest = test.reset_index(drop=True)\n\nprint(f\"\\n✓ Dataset split:\")\nprint(f\"  Train: {len(train)} ({len(train)/len(df)*100:.1f}%)\")\nprint(f\"  Val:   {len(val)} ({len(val)/len(df)*100:.1f}%)\")\nprint(f\"  Test:  {len(test)} ({len(test)/len(df)*100:.1f}%)\")\n\nprint(f\"\\nTrain language distribution: {train['language'].value_counts().to_dict()}\")\nprint(f\"Val language distribution:   {val['language'].value_counts().to_dict()}\")\nprint(f\"Test language distribution:  {test['language'].value_counts().to_dict()}\")\n\n# Save splits\nfor split_name, split_data in [(\"train\", train), (\"val\", val), (\"test\", test)]:\n    path = os.path.join(Config.OUTPUT_DIR, f\"{split_name}_split.csv\")\n    split_data.to_csv(path, index=False)\n    print(f\"  ✓ {split_name}_split.csv saved\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T13:59:22.357326Z","iopub.execute_input":"2026-01-23T13:59:22.358054Z","iopub.status.idle":"2026-01-23T13:59:22.944043Z","shell.execute_reply.started":"2026-01-23T13:59:22.358018Z","shell.execute_reply":"2026-01-23T13:59:22.943461Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# BENGALI + ENGLISH VOCABULARY FOR WAV2VEC2\n# =========================\nvocab_dict = {\n    \"|\": 0,\n    \"_\": 1,\n    # English characters\n    \"a\": 2,\n    \"b\": 3,\n    \"c\": 4,\n    \"d\": 5,\n    \"e\": 6,\n    \"f\": 7,\n    \"g\": 8,\n    \"h\": 9,\n    \"i\": 10,\n    \"j\": 11,\n    \"k\": 12,\n    \"l\": 13,\n    \"m\": 14,\n    \"n\": 15,\n    \"o\": 16,\n    \"p\": 17,\n    \"r\": 18,\n    \"s\": 19,\n    \"t\": 20,\n    \"u\": 21,\n    \"v\": 22,\n    \"w\": 23,\n    \"x\": 24,\n    \"y\": 25,\n    \"z\": 26,\n    # Special characters\n    '\\x93': 27,\n    '\\x94': 28,\n    \"œ\": 29,\n    \"।\": 30,\n    # Bengali vowels (স্বরবর্ণ)\n    \"ঁ\": 31,\n    \"ং\": 32,\n    \"ঃ\": 33,\n    \"অ\": 34,\n    \"আ\": 35,\n    \"ই\": 36,\n    \"ঈ\": 37,\n    \"উ\": 38,\n    \"ঊ\": 39,\n    \"ঋ\": 40,\n    \"এ\": 41,\n    \"ঐ\": 42,\n    \"ও\": 43,\n    \"ঔ\": 44,\n    # Bengali consonants (ব্যঞ্জনবর্ণ)\n    \"ক\": 45,\n    \"খ\": 46,\n    \"গ\": 47,\n    \"ঘ\": 48,\n    \"ঙ\": 49,\n    \"চ\": 50,\n    \"ছ\": 51,\n    \"জ\": 52,\n    \"ঝ\": 53,\n    \"ঞ\": 54,\n    \"ট\": 55,\n    \"ঠ\": 56,\n    \"ড\": 57,\n    \"ঢ\": 58,\n    \"ণ\": 59,\n    \"ত\": 60,\n    \"থ\": 61,\n    \"দ\": 62,\n    \"ধ\": 63,\n    \"ন\": 64,\n    \"প\": 65,\n    \"ফ\": 66,\n    \"ব\": 67,\n    \"ভ\": 68,\n    \"ম\": 69,\n    \"য\": 70,\n    \"র\": 71,\n    \"ল\": 72,\n    \"শ\": 73,\n    \"ষ\": 74,\n    \"স\": 75,\n    \"হ\": 76,\n    # Bengali diacritics (কার/modifiers)\n    \"়\": 77,\n    \"া\": 78,\n    \"ি\": 79,\n    \"ী\": 80,\n    \"ু\": 81,\n    \"ূ\": 82,\n    \"ৃ\": 83,\n    \"ে\": 84,\n    \"ৈ\": 85,\n    \"ো\": 86,\n    \"ৌ\": 87,\n    \"্\": 88,\n    \"ৎ\": 89,\n    \"ৗ\": 90,\n    # Additional Bengali consonants\n    \"ড়\": 91,\n    \"ঢ়\": 92,\n    \"য়\": 93,\n    # Bengali digits (সংখ্যা)\n    \"०\": 94,\n    \"१\": 95,\n    \"२\": 96,\n    \"३\": 97,\n    \"४\": 98,\n    \"५\": 99,\n    \"६\": 100,\n    \"७\": 101,\n    \"८\": 102,\n    \"९\": 103,\n    \"ৰ\": 104,\n    # Unicode control characters\n    '\\u200c': 105,  # Zero-width non-joiner\n    '\\u200d': 106,  # Zero-width joiner\n    '\\u200e': 107,  # Left-to-right mark\n    # Special tokens\n    \"[UNK]\": 108,\n    \"[PAD]\": 109,\n    \"<s>\": 110,\n    \"</s>\": 111,\n}\n\nprint(f\"✓ Bengali + English vocabulary loaded\")\nprint(f\"  Total characters: {len(vocab_dict)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T13:59:22.94508Z","iopub.execute_input":"2026-01-23T13:59:22.945427Z","iopub.status.idle":"2026-01-23T13:59:22.955263Z","shell.execute_reply.started":"2026-01-23T13:59:22.945402Z","shell.execute_reply":"2026-01-23T13:59:22.954735Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#new\n\n# =========================\n# LOAD PRETRAINED MODEL WITH FROZEN ENCODER\n# =========================\nprint(f\"Loading {Config.PRETRAINED_MODEL}...\")\n\ntry:\n    # Try loading processor directly\n    processor = Wav2Vec2Processor.from_pretrained(Config.PRETRAINED_MODEL)\nexcept (OSError, ValueError, Exception):\n    # If processor not available, create from feature extractor and custom tokenizer\n    print(\"⚠️  Processor not found. Creating from feature extractor and custom Bengali tokenizer...\")\n    from transformers import Wav2Vec2FeatureExtractor, Wav2Vec2CTCTokenizer\n    \n    feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(Config.PRETRAINED_MODEL)\n    \n    # Save vocab file\n    vocab_file = os.path.join(Config.OUTPUT_DIR, \"vocab.json\")\n    with open(vocab_file, 'w', encoding='utf-8') as f:\n        json.dump(vocab_dict, f, ensure_ascii=False, indent=2)\n    \n    # Create tokenizer with Bengali vocabulary\n    tokenizer = Wav2Vec2CTCTokenizer(\n        vocab_file,\n        unk_token=\"[UNK]\",\n        pad_token=\"[PAD]\",\n        word_delimiter_token=\"|\"\n    )\n    processor = Wav2Vec2Processor(\n        feature_extractor=feature_extractor,\n        tokenizer=tokenizer\n    )\n    print(\"✓ Processor created with Bengali + English vocabulary\")\n\nmodel = Wav2Vec2ForCTC.from_pretrained(\n    Config.PRETRAINED_MODEL,\n    attention_dropout=Config.ATTENTION_DROPOUT,\n    hidden_dropout=Config.HIDDEN_DROPOUT,\n    feat_proj_dropout=Config.FEAT_PROJ_DROPOUT,\n    mask_time_prob=Config.MASK_TIME_PROB,\n    ctc_loss_reduction=\"mean\",\n    pad_token_id=processor.tokenizer.pad_token_id,\n)\n\n# ========== RESIZE CTC HEAD TO MATCH CUSTOM VOCAB ==========\nprint(f\"\\n📏 Resizing CTC head to match vocabulary...\")\nprint(f\"  Original vocab size: {model.config.vocab_size}\")\nprint(f\"  Tokenizer vocab size: {len(processor.tokenizer)}\")\n\nold_lm_head = model.lm_head\nold_vocab_size = model.config.vocab_size\nnew_vocab_size = len(processor.tokenizer)\n\n# Create new linear layer for CTC head\nhidden_size = old_lm_head.in_features\nnew_lm_head = torch.nn.Linear(hidden_size, new_vocab_size)\n\n# Initialize with small random values\nwith torch.no_grad():\n    new_lm_head.weight.normal_(mean=0.0, std=0.02)\n    if new_lm_head.bias is not None:\n        new_lm_head.bias.zero_()\n\n# Replace the old lm_head\nmodel.lm_head = new_lm_head\nmodel.config.vocab_size = new_vocab_size\n\nprint(f\"✓ CTC head resized: {old_vocab_size} → {new_vocab_size}\")\n\n# ========== CONFIGURE CTC SETTINGS ==========\nprint(f\"\\n✓ Configuring CTC settings...\")\nmodel.config.pad_token_id = processor.tokenizer.pad_token_id\nmodel.config.eos_token_id = processor.tokenizer.eos_token_id\nmodel.config.bos_token_id = processor.tokenizer.bos_token_id\nmodel.config.ctc_zero_infinity = False\nprint(f\"  Pad token ID: {model.config.pad_token_id}\")\nprint(f\"  Vocab size: {model.config.vocab_size}\")\n\n# FREEZE ENCODER: Only train CTC head\nif Config.FREEZE_ENCODER:\n    print(\"\\n🔒 Freezing encoder layers...\")\n    \n    # Freeze feature extractor\n    model.wav2vec2.feature_extractor._freeze_parameters()\n    \n    # Freeze wav2vec2 encoder\n    for param in model.wav2vec2.parameters():\n        param.requires_grad = False\n    \n    # CTC head remains trainable\n    for param in model.lm_head.parameters():\n        param.requires_grad = True\n    \n    print(\"✓ Encoder frozen, CTC head trainable\")\n\nif Config.GRADIENT_CHECKPOINTING:\n    model.gradient_checkpointing_enable()\n\nmodel = model.to(Config.DEVICE)\n\n# Count trainable parameters\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n\nprint(f\"\\n✓ Model loaded and configured\")\nprint(f\"  Model: {Config.PRETRAINED_MODEL}\")\nprint(f\"  Vocab size: {len(processor.tokenizer)}\")\nprint(f\"  Total parameters: {total_params / 1e6:.1f}M\")\nprint(f\"  Trainable parameters: {trainable_params / 1e6:.1f}M\")\nprint(f\"  Frozen parameters: {(total_params - trainable_params) / 1e6:.1f}M\")\nprint(f\"  Trainable %: {trainable_params / total_params * 100:.2f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T14:10:20.069836Z","iopub.execute_input":"2026-01-23T14:10:20.070883Z","iopub.status.idle":"2026-01-23T14:10:22.817858Z","shell.execute_reply.started":"2026-01-23T14:10:20.070821Z","shell.execute_reply":"2026-01-23T14:10:22.817094Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# LOAD PRETRAINED MODEL WITH FROZEN ENCODER\n# =========================\nprint(f\"Loading {Config.PRETRAINED_MODEL}...\")\n\nprocessor = None\ntry:\n    # Try loading processor directly\n    processor = Wav2Vec2Processor.from_pretrained(Config.PRETRAINED_MODEL)\n    print(\"✓ Processor loaded successfully\")\nexcept Exception as e:\n    # If processor not available, create from feature extractor and custom tokenizer\n    print(f\"⚠️  Processor loading failed: {type(e).__name__}: {str(e)[:100]}\")\n    print(\"Creating from feature extractor and custom Bengali tokenizer...\")\n    \n    from transformers import Wav2Vec2FeatureExtractor, Wav2Vec2CTCTokenizer\n    \n    try:\n        feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(Config.PRETRAINED_MODEL)\n        print(\"✓ Feature extractor loaded\")\n    except Exception as fe:\n        print(f\"⚠️  Feature extractor loading failed: {fe}\")\n        # Create a default feature extractor\n        from transformers import Wav2Vec2FeatureExtractor\n        feature_extractor = Wav2Vec2FeatureExtractor(\n            feature_size=1,\n            sampling_rate=Config.SAMPLE_RATE,\n            padding_value=0.0,\n            do_normalize=True,\n            return_attention_mask=True\n        )\n        print(\"✓ Default feature extractor created\")\n    \n    # Ensure output directory exists\n    os.makedirs(Config.OUTPUT_DIR, exist_ok=True)\n    \n    # Save vocab file\n    vocab_file = os.path.join(Config.OUTPUT_DIR, \"vocab.json\")\n    with open(vocab_file, 'w', encoding='utf-8') as f:\n        json.dump(vocab_dict, f, ensure_ascii=False, indent=2)\n    print(f\"✓ Vocab file saved to {vocab_file}\")\n    \n    # Create tokenizer with Bengali vocabulary\n    try:\n        tokenizer = Wav2Vec2CTCTokenizer(\n            vocab_file,\n            unk_token=\"[UNK]\",\n            pad_token=\"[PAD]\",\n            word_delimiter_token=\"|\"\n        )\n        print(\"✓ Custom tokenizer created\")\n    except Exception as te:\n        print(f\"⚠️  Tokenizer creation failed: {te}\")\n        raise\n    \n    processor = Wav2Vec2Processor(\n        feature_extractor=feature_extractor,\n        tokenizer=tokenizer\n    )\n    print(\"✓ Processor created with Bengali + English vocabulary\")\n\nif processor is None:\n    raise RuntimeError(\"Failed to create or load processor\")\n\n# Load model\nprint(f\"\\nLoading model: {Config.PRETRAINED_MODEL}...\")\ntry:\n    model = Wav2Vec2ForCTC.from_pretrained(\n        Config.PRETRAINED_MODEL,\n        attention_dropout=Config.ATTENTION_DROPOUT,\n        hidden_dropout=Config.HIDDEN_DROPOUT,\n        feat_proj_dropout=Config.FEAT_PROJ_DROPOUT,\n        mask_time_prob=Config.MASK_TIME_PROB,\n        ctc_loss_reduction=\"mean\",\n        pad_token_id=processor.tokenizer.pad_token_id,\n    )\n    print(\"✓ Model loaded successfully\")\nexcept Exception as e:\n    print(f\"❌ Model loading failed: {e}\")\n    raise\n\n# FREEZE ENCODER: Only train CTC head\nif Config.FREEZE_ENCODER:\n    print(\"\\n🔒 Freezing encoder layers...\")\n    \n    try:\n        # Freeze feature extractor\n        model.wav2vec2.feature_extractor._freeze_parameters()\n        print(\"  ✓ Feature extractor frozen\")\n    except Exception as e:\n        print(f\"  ⚠️  Could not freeze feature extractor: {e}\")\n    \n    # Freeze wav2vec2 encoder\n    frozen_count = 0\n    for name, param in model.wav2vec2.named_parameters():\n        param.requires_grad = False\n        frozen_count += 1\n    \n    # CTC head remains trainable\n    trainable_count = 0\n    for param in model.lm_head.parameters():\n        param.requires_grad = True\n        trainable_count += 1\n    \n    print(f\"  ✓ Encoder frozen ({frozen_count} parameters)\")\n    print(f\"  ✓ CTC head trainable ({trainable_count} parameters)\")\n\nif Config.GRADIENT_CHECKPOINTING:\n    try:\n        model.gradient_checkpointing_enable()\n        print(\"✓ Gradient checkpointing enabled\")\n    except Exception as e:\n        print(f\"⚠️  Gradient checkpointing not available: {e}\")\n\nmodel = model.to(Config.DEVICE)\n\n# Count trainable parameters\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n\nprint(f\"\\n✓ Model loaded and configured\")\nprint(f\"  Model: {Config.PRETRAINED_MODEL}\")\nprint(f\"  Vocab size: {len(processor.tokenizer)}\")\nprint(f\"  Total parameters: {total_params / 1e6:.1f}M\")\nprint(f\"  Trainable parameters: {trainable_params / 1e6:.1f}M\")\nprint(f\"  Frozen parameters: {(total_params - trainable_params) / 1e6:.1f}M\")\nprint(f\"  Trainable %: {trainable_params / total_params * 100:.2f}%\")\nprint(f\"  Device: {Config.DEVICE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T14:07:40.42164Z","iopub.execute_input":"2026-01-23T14:07:40.421892Z","iopub.status.idle":"2026-01-23T14:07:42.984252Z","shell.execute_reply.started":"2026-01-23T14:07:40.42187Z","shell.execute_reply":"2026-01-23T14:07:42.983214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# MEMORY-EFFICIENT DATASET CLASS (ON-THE-FLY LOADING)\n# =========================\nclass AudioDataset(Dataset):\n    def __init__(self, dataframe: pd.DataFrame, processor: Wav2Vec2Processor, use_augment: bool = False):\n        self.df = dataframe.reset_index(drop=True)\n        self.processor = processor\n        self.feature_extractor = processor.feature_extractor\n        self.tokenizer = processor.tokenizer\n        self.use_augment = use_augment\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx: int) -> Dict:\n        row = self.df.iloc[idx]\n        \n        try:\n            # Load audio on-the-fly\n            waveform, sr = torchaudio.load(row[\"audio_path\"])\n            \n            # Resample if needed\n            if sr != Config.SAMPLE_RATE:\n                resampler = T.Resample(sr, Config.SAMPLE_RATE)\n                waveform = resampler(waveform)\n            \n            # Convert to mono\n            if waveform.shape[0] > 1:\n                waveform = waveform.mean(dim=0, keepdim=True)\n            \n            waveform = waveform.squeeze(0)\n            \n            # Validate and adjust duration\n            duration = len(waveform) / Config.SAMPLE_RATE\n            if duration < Config.MIN_AUDIO_LENGTH or duration > Config.MAX_AUDIO_LENGTH:\n                max_samples = int(Config.MAX_AUDIO_LENGTH * Config.SAMPLE_RATE)\n                if len(waveform) > max_samples:\n                    waveform = waveform[:max_samples]\n                elif len(waveform) < int(Config.MIN_AUDIO_LENGTH * Config.SAMPLE_RATE):\n                    waveform = torch.nn.functional.pad(\n                        waveform,\n                        (0, int(Config.MIN_AUDIO_LENGTH * Config.SAMPLE_RATE) - len(waveform))\n                    )\n            \n            # Minimal speed augmentation\n            if self.use_augment and random.random() < 0.3:\n                speed = random.choice([0.95, 1.0, 1.05])\n                if speed != 1.0:\n                    new_len = int(len(waveform) / speed)\n                    waveform = torch.nn.functional.interpolate(\n                        waveform.unsqueeze(0).unsqueeze(0),\n                        size=new_len,\n                        mode='linear',\n                        align_corners=False\n                    ).squeeze()\n        \n        except Exception as e:\n            print(f\"Error loading {row['audio_path']}: {e}\")\n            waveform = torch.zeros(int(Config.SAMPLE_RATE * 5))\n        \n        # Process audio using feature extractor\n        inputs = self.feature_extractor(\n            waveform.numpy(),\n            sampling_rate=Config.SAMPLE_RATE,\n            return_tensors=\"pt\",\n            return_attention_mask=True\n        )\n        \n        # Encode text labels using tokenizer\n        labels = self.tokenizer(\n            row[\"text\"],\n            return_tensors=\"pt\",\n            padding=\"longest\",\n            truncation=True\n        )\n        \n        return {\n            \"input_values\": inputs.input_values.squeeze(0),\n            \"attention_mask\": inputs.attention_mask.squeeze(0),\n            \"labels\": labels.input_ids.squeeze(0),\n            \"input_ids\": labels.input_ids.squeeze(0)\n        }\n\nprint(\"✓ Dataset class ready (on-the-fly audio loading)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T14:10:34.369522Z","iopub.execute_input":"2026-01-23T14:10:34.37009Z","iopub.status.idle":"2026-01-23T14:10:34.381149Z","shell.execute_reply.started":"2026-01-23T14:10:34.370056Z","shell.execute_reply":"2026-01-23T14:10:34.380508Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# DATA COLLATOR FOR CTC\n# =========================\n@dataclass\nclass DataCollatorCTCWithPadding:\n    processor: Wav2Vec2Processor\n    padding: Union[bool, str] = \"longest\"\n    \n    def __call__(self, features: List[Dict]) -> Dict[str, torch.Tensor]:\n        input_features = [{\"input_values\": f[\"input_values\"]} for f in features]\n        label_features = [{\"input_ids\": f[\"labels\"]} for f in features]\n        \n        # Pad audio inputs using feature extractor\n        batch = self.processor.feature_extractor.pad(\n            input_features,\n            padding=self.padding,\n            return_tensors=\"pt\",\n            return_attention_mask=True\n        )\n        \n        # Pad labels using tokenizer\n        labels_batch = self.processor.tokenizer.pad(\n            label_features,\n            padding=self.padding,\n            return_tensors=\"pt\"\n        )\n        \n        labels = labels_batch[\"input_ids\"].masked_fill(\n            labels_batch.attention_mask.ne(1), -100\n        )\n        batch[\"labels\"] = labels\n        \n        return batch\n\ndata_collator = DataCollatorCTCWithPadding(processor=processor)\nprint(\"✓ Data collator ready\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T14:10:38.103066Z","iopub.execute_input":"2026-01-23T14:10:38.103612Z","iopub.status.idle":"2026-01-23T14:10:38.111016Z","shell.execute_reply.started":"2026-01-23T14:10:38.103575Z","shell.execute_reply":"2026-01-23T14:10:38.11019Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CREATE DATASETS\n# =========================\nprint(\"Creating datasets...\")\ntrain_dataset = AudioDataset(train, processor, use_augment=True)\nval_dataset = AudioDataset(val, processor, use_augment=False)\ntest_dataset = AudioDataset(test, processor, use_augment=False)\n\nprint(f\"✓ Train: {len(train_dataset)} | Val: {len(val_dataset)} | Test: {len(test_dataset)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T14:10:40.746649Z","iopub.execute_input":"2026-01-23T14:10:40.747556Z","iopub.status.idle":"2026-01-23T14:10:40.861403Z","shell.execute_reply.started":"2026-01-23T14:10:40.747514Z","shell.execute_reply":"2026-01-23T14:10:40.860753Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# TRAINING ARGUMENTS (LOW GPU/RAM OPTIMIZED, FROZEN ENCODER)\n# =========================\ntraining_args = TrainingArguments(\n    output_dir=Config.CHECKPOINT_DIR,\n    group_by_length=False,\n    per_device_train_batch_size=Config.BATCH_SIZE,\n    per_device_eval_batch_size=Config.BATCH_SIZE,\n    gradient_accumulation_steps=Config.GRADIENT_ACCUMULATION,\n    eval_strategy=\"steps\",\n    num_train_epochs=Config.NUM_EPOCHS,\n    gradient_checkpointing=Config.GRADIENT_CHECKPOINTING,\n    fp16=Config.FP16,\n    save_steps=Config.SAVE_STEPS,\n    eval_steps=Config.EVAL_STEPS,\n    logging_steps=Config.LOGGING_STEPS,\n    learning_rate=Config.LEARNING_RATE,\n    warmup_steps=Config.WARMUP_STEPS,\n    save_total_limit=3,\n    load_best_model_at_end=True,\n    metric_for_best_model=\"cer\",\n    greater_is_better=False,\n    lr_scheduler_type=\"cosine\",\n    dataloader_num_workers=2,\n    dataloader_pin_memory=True,\n    remove_unused_columns=False,\n    push_to_hub=False,\n    run_name=\"asr-wav2vec2-frozen-multilingual\",\n)\n\nprint(\"✓ Training arguments configured\")\nprint(f\"   Batch Size: {Config.BATCH_SIZE}\")\nprint(f\"   Gradient Accumulation: {Config.GRADIENT_ACCUMULATION}\")\nprint(f\"   Effective batch size: {Config.BATCH_SIZE * Config.GRADIENT_ACCUMULATION}\")\nprint(f\"   FP16: {Config.FP16}\")\nprint(f\"   Learning Rate: {Config.LEARNING_RATE}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T14:10:42.999662Z","iopub.execute_input":"2026-01-23T14:10:42.999968Z","iopub.status.idle":"2026-01-23T14:10:43.038837Z","shell.execute_reply.started":"2026-01-23T14:10:42.999943Z","shell.execute_reply":"2026-01-23T14:10:43.038258Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# METRICS COMPUTATION\n# =========================\nclass ComputeMetrics:\n    def __init__(self, processor: Wav2Vec2Processor):\n        self.processor = processor\n    \n    def __call__(self, eval_pred):\n        predictions, label_ids = eval_pred\n        pred_ids = np.argmax(predictions, axis=-1)\n        \n        pred_str = self.processor.batch_decode(pred_ids)\n        label_str = self.processor.batch_decode(label_ids, group_tokens=False)\n        \n        wer_score = wer(label_str, pred_str)\n        cer_score = cer(label_str, pred_str)\n        \n        return {\"wer\": wer_score, \"cer\": cer_score}\n\ncompute_metrics = ComputeMetrics(processor)\nprint(\"✓ Metrics function ready\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T14:10:46.420257Z","iopub.execute_input":"2026-01-23T14:10:46.420865Z","iopub.status.idle":"2026-01-23T14:10:46.426661Z","shell.execute_reply.started":"2026-01-23T14:10:46.420831Z","shell.execute_reply":"2026-01-23T14:10:46.425934Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# INITIALIZE TRAINER\n# =========================\ntrainer = Trainer(\n    model=model,\n    data_collator=data_collator,\n    args=training_args,\n    compute_metrics=compute_metrics,\n    train_dataset=train_dataset,\n    eval_dataset=val_dataset,\n    callbacks=[\n        EarlyStoppingCallback(\n            early_stopping_patience=Config.EARLY_STOPPING_PATIENCE,\n            early_stopping_threshold=0.001\n        )\n    ]\n)\n\nprint(\"✓ Trainer initialized (frozen encoder + head-only)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T14:10:48.83861Z","iopub.execute_input":"2026-01-23T14:10:48.839442Z","iopub.status.idle":"2026-01-23T14:10:48.861713Z","shell.execute_reply.started":"2026-01-23T14:10:48.839384Z","shell.execute_reply":"2026-01-23T14:10:48.861099Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#new\n\n\n# =========================\n# START TRAINING (HEAD-ONLY)\n# =========================\nprint(\"\\n\" + \"=\"*70)\nprint(\"🚀 STARTING TRAINING (FROZEN ENCODER + HEAD-ONLY)\")\nprint(\"=\"*70 + \"\\n\")\n\n# Verify model is in correct mode\nmodel.train()\nprint(f\"✓ Model set to training mode\")\n\ntry:\n    train_result = trainer.train()\n    \n    print(\"\\n\" + \"=\"*70)\n    print(\"✅ TRAINING COMPLETED\")\n    print(\"=\"*70)\n    print(f\"Final training loss: {train_result.training_loss:.4f}\")\n    \nexcept RuntimeError as e:\n    if \"out of memory\" in str(e).lower():\n        print(\"⚠️  CUDA Out of Memory! Adjusting...\")\n        gc.collect()\n        torch.cuda.empty_cache()\n        \n        print(\"Reducing batch size for retry...\")\n        trainer.args.per_device_train_batch_size = max(1, Config.BATCH_SIZE // 2)\n        trainer.args.gradient_accumulation_steps = Config.GRADIENT_ACCUMULATION * 2\n        \n        print(f\"New effective batch size: {trainer.args.per_device_train_batch_size * trainer.args.gradient_accumulation_steps}\")\n        print(\"Retrying training...\")\n        \n        try:\n            train_result = trainer.train()\n            print(\"\\n✅ TRAINING COMPLETED (with reduced batch size)\")\n            print(f\"Final training loss: {train_result.training_loss:.4f}\")\n        except Exception as retry_error:\n            print(f\"❌ Training failed: {retry_error}\")\n            raise retry_error\n    else:\n        print(f\"❌ RuntimeError: {e}\")\n        raise e\n        \nexcept Exception as e:\n    print(f\"❌ Unexpected error: {type(e).__name__}: {e}\")\n    raise e\n\nprint(f\"\\n✓ Training pipeline completed successfully\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-23T14:10:58.059073Z","iopub.execute_input":"2026-01-23T14:10:58.059381Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# EVALUATE ON TEST SET\n# =========================\nprint(\"\\nEvaluating on test set...\")\ntest_results = trainer.evaluate(eval_dataset=test_dataset)\n\nprint(\"\\n\" + \"=\"*70)\nprint(\"📊 TEST SET BASELINE RESULTS\")\nprint(\"=\"*70)\nfor key, value in test_results.items():\n    print(f\"{key:.<50} {value:.4f}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# GENERATE PREDICTIONS FOR DETAILED METRICS\n# =========================\nprint(\"\\nGenerating detailed predictions...\")\npredictions = []\nlabels_list = []\ninference_times = []\n\nmodel.eval()\nwith torch.no_grad():\n    for i in tqdm(range(len(test_dataset)), desc=\"Predicting\"):\n        sample = test_dataset[i]\n        \n        input_values = sample[\"input_values\"].unsqueeze(0).to(Config.DEVICE)\n        attention_mask = sample[\"attention_mask\"].unsqueeze(0).to(Config.DEVICE)\n        labels = sample[\"labels\"]\n        \n        start = time.time()\n        outputs = model(input_values, attention_mask=attention_mask)\n        end = time.time()\n        \n        pred_ids = np.argmax(outputs.logits.cpu().numpy(), axis=-1)\n        \n        predictions.append(pred_ids[0])\n        labels_list.append(labels.numpy())\n        inference_times.append(end - start)\n\nprint(f\"✓ Generated {len(predictions)} predictions\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# COMPUTE ALL 4 METRICS: WER, CER, WAcc, RTFx\n# =========================\nprint(\"\\nComputing comprehensive metrics...\")\n\n# Decode predictions and labels\npred_strs = processor.batch_decode(np.array(predictions))\nlabel_strs = processor.batch_decode(np.array(labels_list), group_tokens=False)\n\n# 1. WER (Word Error Rate)\nwer_score = wer(label_strs, pred_strs)\n\n# 2. CER (Character Error Rate)\ncer_score = cer(label_strs, pred_strs)\n\n# 3. WAcc (Word Accuracy)\npred_words_list = [p.split() for p in pred_strs]\nlabel_words_list = [l.split() for l in label_strs]\n\ntotal_words = sum(len(l) for l in label_words_list)\ncorrect_words = 0\n\nfor pred_words, label_words in zip(pred_words_list, label_words_list):\n    for pw, lw in zip(pred_words, label_words):\n        if pw == lw:\n            correct_words += 1\n\nword_accuracy = (correct_words / total_words * 100) if total_words > 0 else 0.0\n\n# 4. RTFx (Real-Time Factor)\ntotal_audio_length = sum([\n    len(test_dataset[i][\"input_values\"]) / Config.SAMPLE_RATE \n    for i in range(len(test_dataset))\n])\ntotal_inference_time = sum(inference_times)\nrtfx = total_inference_time / total_audio_length if total_audio_length > 0 else 0.0\n\nprint(\"\\n\" + \"=\"*70)\nprint(\"📊 COMPREHENSIVE TEST METRICS (FROZEN ENCODER)\")\nprint(\"=\"*70)\n\nprint(f\"\\n1️⃣ Word Error Rate (WER):\")\nprint(f\"   Value: {wer_score*100:.2f}%\")\nprint(f\"   Lower is better. Ideal: 0%\")\n\nprint(f\"\\n2️⃣ Character Error Rate (CER):\")\nprint(f\"   Value: {cer_score*100:.2f}%\")\nprint(f\"   Lower is better. Ideal: 0%\")\n\nprint(f\"\\n3️⃣ Word Accuracy (WAcc):\")\nprint(f\"   Value: {word_accuracy:.2f}%\")\nprint(f\"   Correct words: {correct_words}/{total_words}\")\nprint(f\"   Higher is better. Ideal: 100%\")\n\nprint(f\"\\n4️⃣ Real-Time Factor (RTFx):\")\nprint(f\"   Value: {rtfx:.4f}\")\nprint(f\"   Total inference time: {total_inference_time:.2f}s\")\nprint(f\"   Total audio length: {total_audio_length:.2f}s\")\nprint(f\"   RTFx < 1.0 = Real-time capable\")\n\nprint(\"\\n\" + \"=\"*70)\n\n# Save metrics\nmetrics_dict = {\n    \"model_config\": {\n        \"model\": Config.PRETRAINED_MODEL,\n        \"frozen_encoder\": Config.FREEZE_ENCODER,\n        \"trainable_only\": \"CTC head\"\n    },\n    \"wer_percent\": float(wer_score * 100),\n    \"cer_percent\": float(cer_score * 100),\n    \"word_accuracy_percent\": float(word_accuracy),\n    \"correct_words\": int(correct_words),\n    \"total_words\": int(total_words),\n    \"rtfx\": float(rtfx),\n    \"total_inference_time_seconds\": float(total_inference_time),\n    \"total_audio_length_seconds\": float(total_audio_length),\n    \"test_samples\": len(test_dataset),\n    \"dataset_info\": {\n        \"total_samples\": len(df),\n        \"train_samples\": len(train),\n        \"val_samples\": len(val),\n        \"test_samples\": len(test),\n        \"english_samples\": len(df[df['language'] == 'en']),\n        \"bengali_samples\": len(df[df['language'] == 'bn'])\n    },\n    \"training_config\": {\n        \"batch_size\": Config.BATCH_SIZE,\n        \"gradient_accumulation\": Config.GRADIENT_ACCUMULATION,\n        \"learning_rate\": Config.LEARNING_RATE,\n        \"num_epochs\": Config.NUM_EPOCHS,\n        \"fp16\": Config.FP16\n    }\n}\n\nmetrics_file = os.path.join(Config.OUTPUT_DIR, \"test_metrics_frozen.json\")\nwith open(metrics_file, 'w') as f:\n    json.dump(metrics_dict, f, indent=2)\n\nprint(f\"\\n✓ Metrics saved to {metrics_file}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# BATCH INFERENCE (OPTIMIZED)\n# =========================\nprint(\"\\n\" + \"=\"*70)\nprint(\"⚡ BATCH MODE INFERENCE (OPTIMIZED)\")\nprint(\"=\"*70)\n\ndef batch_inference(dataset, batch_size=8):\n    \"\"\"Efficient batch inference\"\"\"\n    all_preds = []\n    all_refs = []\n    batch_times = []\n    \n    model.eval()\n    with torch.no_grad():\n        for batch_idx in tqdm(range(0, len(dataset), batch_size), desc=\"Batch Inference\"):\n            batch_samples = []\n            batch_labels = []\n            \n            for i in range(batch_idx, min(batch_idx + batch_size, len(dataset))):\n                sample = dataset[i]\n                batch_samples.append(sample)\n                batch_labels.append(test_dataset.df.iloc[i]['text'])\n            \n            # Prepare batch\n            input_values_list = [s[\"input_values\"] for s in batch_samples]\n            attention_masks = [s[\"attention_mask\"] for s in batch_samples]\n            \n            # Pad to same length\n            max_len = max([iv.shape[0] for iv in input_values_list])\n            padded_inputs = []\n            padded_masks = []\n            \n            for iv, am in zip(input_values_list, attention_masks):\n                pad_len = max_len - iv.shape[0]\n                padded_inputs.append(torch.nn.functional.pad(iv, (0, pad_len)))\n                padded_masks.append(torch.nn.functional.pad(am, (0, pad_len), value=0))\n            \n            batch_inputs = torch.stack(padded_inputs).to(Config.DEVICE)\n            batch_masks = torch.stack(padded_masks).to(Config.DEVICE)\n            \n            # Inference\n            start = time.time()\n            outputs = model(batch_inputs, attention_mask=batch_masks)\n            batch_time = time.time() - start\n            \n            # Decode\n            pred_ids = np.argmax(outputs.logits.cpu().numpy(), axis=-1)\n            pred_strs_batch = processor.batch_decode(pred_ids)\n            \n            all_preds.extend(pred_strs_batch)\n            all_refs.extend(batch_labels)\n            batch_times.append(batch_time)\n    \n    return all_preds, all_refs, sum(batch_times)\n\n# Run batch inference on test set\nbatch_preds, batch_refs, batch_time = batch_inference(test_dataset, batch_size=8)\n\nprint(f\"\\n✓ Batch inference completed in {batch_time:.2f}s\")\nprint(f\"  Samples processed: {len(batch_preds)}\")\nprint(f\"  Time per sample: {batch_time / len(batch_preds) * 1000:.2f}ms\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# SHOW SAMPLE PREDICTIONS\n# =========================\nprint(\"\\nSample Predictions (first 10):\")\nprint(\"=\"*100 + \"\\n\")\n\nfor i in range(min(10, len(pred_strs))):\n    print(f\"Sample {i+1}:\")\n    print(f\"  Reference: {label_strs[i][:80]}...\")\n    print(f\"  Predicted: {pred_strs[i][:80]}...\")\n    \n    ref_words = label_strs[i].split()\n    pred_words = pred_strs[i].split()\n    matches = sum(1 for r, p in zip(ref_words, pred_words) if r == p)\n    acc = (matches / len(ref_words) * 100) if len(ref_words) > 0 else 0\n    print(f\"  Sample Accuracy: {acc:.1f}%\\n\")\n\nprint(\"=\"*100)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# SAVE BEST MODEL CHECKPOINT\n# =========================\nprint(\"Saving best model checkpoint...\")\n\nmodel.save_pretrained(Config.MODEL_DIR)\nprocessor.save_pretrained(Config.MODEL_DIR)\n\nprint(f\"✓ Model saved to {Config.MODEL_DIR}\")\n\n# Save training configuration and results\nfinal_config = {\n    \"model_info\": {\n        \"pretrained_model\": Config.PRETRAINED_MODEL,\n        \"frozen_encoder\": Config.FREEZE_ENCODER,\n        \"training_mode\": \"Head-only (CTC layer)\"\n    },\n    \"dataset_info\": {\n        \"source\": Config.UNIFIED_CSV,\n        \"total_samples\": len(df),\n        \"train_samples\": len(train),\n        \"val_samples\": len(val),\n        \"test_samples\": len(test),\n        \"language_distribution\": {\n            \"english\": int(len(df[df['language'] == 'en'])),\n            \"bengali\": int(len(df[df['language'] == 'bn']))\n        }\n    },\n    \"training_config\": {\n        \"batch_size\": Config.BATCH_SIZE,\n        \"gradient_accumulation\": Config.GRADIENT_ACCUMULATION,\n        \"learning_rate\": Config.LEARNING_RATE,\n        \"num_epochs\": Config.NUM_EPOCHS,\n        \"warmup_steps\": Config.WARMUP_STEPS,\n        \"device\": Config.DEVICE,\n        \"fp16_enabled\": Config.FP16,\n        \"gradient_checkpointing\": Config.GRADIENT_CHECKPOINTING\n    },\n    \"test_metrics\": metrics_dict\n}\n\nconfig_file = os.path.join(Config.MODEL_DIR, \"training_config.json\")\nwith open(config_file, 'w') as f:\n    json.dump(final_config, f, indent=2)\n\nprint(f\"✓ Configuration saved\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# GOOGLE DRIVE UPLOAD\n# =========================\nprint(\"\\nAttempting Google Drive backup...\")\n\ntry:\n    from google.colab import auth\n    from googleapiclient.discovery import build\n    from googleapiclient.http import MediaFileUpload\n    \n    print(\"Authenticating with Google Drive...\")\n    auth.authenticate_user()\n    drive_service = build('drive', 'v3')\n    \n    # Create folder\n    folder_metadata = {\n        'name': 'ASR-Wav2Vec2-Frozen-Multilingual-2L',\n        'mimeType': 'application/vnd.google-apps.folder'\n    }\n    folder = drive_service.files().create(body=folder_metadata, fields='id').execute()\n    folder_id = folder.get('id')\n    print(f\"✓ Created folder: {folder_id}\")\n    \n    # Upload critical files\n    files_to_upload = [\n        (os.path.join(Config.MODEL_DIR, \"pytorch_model.bin\"), \"pytorch_model.bin\"),\n        (os.path.join(Config.MODEL_DIR, \"config.json\"), \"config.json\"),\n        (os.path.join(Config.MODEL_DIR, \"training_config.json\"), \"training_config.json\"),\n        (os.path.join(Config.MODEL_DIR, \"preprocessor_config.json\"), \"preprocessor_config.json\"),\n        (metrics_file, \"test_metrics_frozen.json\")\n    ]\n    \n    for filepath, filename in files_to_upload:\n        if os.path.exists(filepath):\n            file_metadata = {'name': filename, 'parents': [folder_id]}\n            media = MediaFileUpload(filepath)\n            drive_service.files().create(body=file_metadata, media_body=media).execute()\n            print(f\"  ✓ {filename}\")\n    \n    print(f\"\\n✓ All files uploaded to Google Drive!\")\n    print(f\"\\nAccess your model:\")\n    print(f\"Folder ID: {folder_id}\")\n    print(f\"Link: https://drive.google.com/drive/folders/{folder_id}\")\n    \nexcept ImportError:\n    print(\"⚠️  Not running on Google Colab. Skipping Drive upload.\")\n    print(f\"Model saved locally at: {Config.MODEL_DIR}\")\nexcept Exception as e:\n    print(f\"⚠️  Drive upload error: {e}\")\n    print(f\"Model still saved at: {Config.MODEL_DIR}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# FINAL SUMMARY\n# =========================\nprint(\"\\n\" + \"=\"*80)\nprint(\"✅ ASR TRAINING PIPELINE COMPLETED (FROZEN ENCODER)\")\nprint(\"=\"*80)\n\nprint(f\"\\n🔒 Model Configuration:\")\nprint(f\"   Model: {Config.PRETRAINED_MODEL}\")\nprint(f\"   Training Mode: Frozen Encoder + Head-Only (CTC)\")\nprint(f\"   Trainable Parameters: {trainable_params / 1e6:.1f}M\")\nprint(f\"   Frozen Parameters: {(total_params - trainable_params) / 1e6:.1f}M\")\n\nprint(f\"\\n📊 Dataset Summary:\")\nprint(f\"   Total samples: {len(df)} (2 Lakh)\")\nprint(f\"   Languages: English ({len(df[df['language'] == 'en'])}) + Bengali ({len(df[df['language'] == 'bn'])})\")\n\nprint(f\"\\n📈 Train/Val/Test Split (70:15:15):\")\nprint(f\"   Training: {len(train)} samples\")\nprint(f\"   Validation: {len(val)} samples\")\nprint(f\"   Testing: {len(test)} samples\")\n\nprint(f\"\\n🎯 Test Performance:\")\nprint(f\"   Word Error Rate (WER):     {wer_score*100:.2f}%\")\nprint(f\"   Character Error Rate (CER): {cer_score*100:.2f}%\")\nprint(f\"   Word Accuracy (WAcc):       {word_accuracy:.2f}%\")\nprint(f\"   Real-Time Factor (RTFx):   {rtfx:.4f}\")\n\nprint(f\"\\n💾 Saved Artifacts:\")\nprint(f\"   Model: {Config.MODEL_DIR}\")\nprint(f\"   Metrics: {metrics_file}\")\nprint(f\"   Checkpoints: {Config.CHECKPOINT_DIR}\")\nprint(f\"   Dataset Splits: {Config.OUTPUT_DIR}\")\n\nprint(f\"\\n⚙️  Optimization Settings:\")\nprint(f\"   Batch Size: {Config.BATCH_SIZE}\")\nprint(f\"   Gradient Accumulation: {Config.GRADIENT_ACCUMULATION}\")\nprint(f\"   Effective Batch: {Config.BATCH_SIZE * Config.GRADIENT_ACCUMULATION}\")\nprint(f\"   Learning Rate: {Config.LEARNING_RATE}\")\nprint(f\"   Device: {Config.DEVICE}\")\nprint(f\"   FP16: {Config.FP16}\")\nprint(f\"   Gradient Checkpointing: {Config.GRADIENT_CHECKPOINTING}\")\n\nprint(f\"\\n⚡ Inference Speed:\")\nprint(f\"   Time per sample: {total_inference_time / len(test_dataset) * 1000:.2f}ms\")\nprint(f\"   RTFx: {rtfx:.4f} (< 1.0 is real-time)\")\n\nprint(\"\\n\" + \"=\"*80)\n","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}]}