{ "nbformat": 4, "nbformat_minor": 5, "metadata": { "kernelspec": { "display_name": "Python 3", "language": "python", "name": "python3" }, "language_info": { "name": "python", "version": "3.10.12" } }, "cells": [ { "cell_type": "markdown", "id": "title", "metadata": {}, "source": "# ๐ŸŒพ Sahel-Voice-Core โ€” Kaggle Master Trainer\n\n**Deep Sleep Fine-Tuning** for `openai/whisper-small` using LoRA (PEFT).\n\nThis single notebook replaces `bootstrap_repos.ipynb`, `train_colab.ipynb`, and `train_fula_tts.ipynb`.\n\n### Data sources integrated\n| Source | Contents | Used for |\n|--------|----------|----------|\n| `ous-sow/sahel-agri-feedback` | `vocabulary.jsonl` + `corrections.jsonl` + audio | Primary fine-tuning signal |\n| `google/WaxalNLP` (bam + ful subsets) | Native speaker recordings | Baseline Bambara & Fula |\n| External datasets (configurable) | e.g. `mozilla-foundation/common_voice_13_0` | Coverage expansion |\n\n### Before running\n1. **Runtime โ†’ Accelerator โ†’ GPU T4 ร— 1** \n2. Add-ons โ†’ Secrets โ†’ `HF_TOKEN` (toggle Attach ON)\n3. Edit **Cell 3** to set your language and version tag prefix" }, { "cell_type": "code", "execution_count": null, "id": "cell-gpu", "metadata": {}, "outputs": [], "source": "# โ”€โ”€ Cell 1: GPU check โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\nimport subprocess, sys\n\nresult = subprocess.run(['nvidia-smi'], capture_output=True, text=True)\nif result.returncode != 0:\n raise RuntimeError('No GPU detected. Runtime โ†’ Accelerator โ†’ GPU T4 ร— 1')\nprint(result.stdout[:600])\n\nimport torch\nprint(f'PyTorch : {torch.__version__}')\nprint(f'CUDA avail: {torch.cuda.is_available()}')\nif torch.cuda.is_available():\n print(f'GPU : {torch.cuda.get_device_name(0)}')\n cap = torch.cuda.get_device_capability(0)\n print(f'Compute : {cap[0]}.{cap[1]}')\n if cap[0] < 7:\n print('โš ๏ธ Compute < 7.0 โ€” bitsandbytes 8-bit ops may not work. Switch to T4 (CC 7.5).')\nprint('โœ… GPU ready')" }, { "cell_type": "code", "execution_count": null, "id": "cell-install", "metadata": {}, "outputs": [], "source": [ "# -- Cell 2: Install minimal missing dependencies ----------------------------\n", "# We do NOT use PEFT/LoRA, so system transformers/numpy/scipy are fine as-is.\n", "# Kaggle does not ship jiwer (WER metric) -- install it now.\n", "import subprocess, sys\n", "\n", "subprocess.check_call([\n", " sys.executable, '-m', 'pip', 'install', '-q', 'jiwer==3.0.4',\n", "])\n", "\n", "import torch\n", "print(f\"torch : {torch.__version__}\")\n", "print(f\"CUDA avail : {torch.cuda.is_available()}\")\n", "if torch.cuda.is_available():\n", " print(f\"GPU : {torch.cuda.get_device_name(0)}\")\n", "\n", "import transformers, datasets as ds_lib\n", "print(f\"transformers: {transformers.__version__}\")\n", "print(f\"datasets : {ds_lib.__version__}\")\n", "print(\"All packages ready.\")\n" ] }, { "cell_type": "code", "execution_count": null, "id": "cell-config", "metadata": {}, "outputs": [], "source": [ "# โ”€โ”€ Cell 3: CONFIGURATION โ€” edit these before each run โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\nimport os\n\n# โ”€โ”€โ”€ Language to train โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\n# 'bam' = Bambara 'ful' = Fula\nTRAIN_LANG = 'bam'\n\n# โ”€โ”€โ”€ Model โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\nWHISPER_MODEL_ID = 'openai/whisper-small'\nTARGET_SR = 16_000\n\n# โ”€โ”€โ”€ HuggingFace repos โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\nHF_USERNAME = 'ous-sow'\nFEEDBACK_REPO_ID = f'{HF_USERNAME}/sahel-agri-feedback'\nADAPTER_REPO_ID = f'{HF_USERNAME}/sahel-agri-adapters'\n\n# โ”€โ”€โ”€ Training hyper-parameters โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\nMAX_STEPS = 4_000 # T4 ~45 min; set 8000 for a deeper run\nBATCH_SIZE = 16\nGRAD_ACCUM = 2 # effective batch = 32\nLEARNING_RATE = 1e-3\nWARMUP_STEPS = 200\nSAVE_STEPS = 500\nEVAL_STEPS = 500\nLOGGING_STEPS = 50\nMAX_WAXAL_TRAIN = 5_000 # cap WaxalNLP samples (streaming budget)\nCORRECTION_REPEAT= 3 # upsample user corrections Nx for emphasis\n\n# โ”€โ”€โ”€ Paths (Kaggle working dir) โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\nWORKING_DIR = '/kaggle/working'\nOUTPUT_DIR = f'{WORKING_DIR}/adapter_{TRAIN_LANG}'\nDATA_DIR = f'{WORKING_DIR}/data'\nAUDIO_DIR = f'{WORKING_DIR}/audio_feedback'\n\nLANG_NAME = {'bam': 'bambara', 'ful': 'fula'}.get(TRAIN_LANG, TRAIN_LANG)\nLANG_COUNTRY = {'bam': 'Mali', 'ful': 'Guinea'}.get(TRAIN_LANG, '')\nLANG_DIALECT = {\n 'bam': 'Standard Bambara (Bamako/Sรฉgou) โ€” Malian orthography',\n 'ful': 'Pular (Labรฉ/Mamou dialects) โ€” Guinean orthography',\n}.get(TRAIN_LANG, '')\n\nprint(f'Language : {TRAIN_LANG} ({LANG_NAME}) โ€” {LANG_COUNTRY}')\nprint(f'Dialect : {LANG_DIALECT}')\nprint(f'Model : {WHISPER_MODEL_ID}')\nprint(f'Output : {OUTPUT_DIR}')\nprint(f'Max steps : {MAX_STEPS}')" ] }, { "cell_type": "code", "execution_count": null, "id": "cell-ext-config", "metadata": {}, "outputs": [], "source": [ "# -- Cell 4: External dataset configuration -----------------------------------\n# EXTERNAL_DATASETS is loaded dynamically from dataset_sources.jsonl in the\n# feedback repo. The Space's Self-Teaching tab writes dataset references there\n# when the user clicks \"Import from HuggingFace\". This cell reads that file\n# so any dataset registered in the Space is automatically used here.\n\nimport json as _json\nfrom huggingface_hub import hf_hub_download as _hf_dl\n\nEXTERNAL_DATASETS = []\n\n# -- Load dataset_sources.jsonl from Hub (written by Space Self-Teaching tab) --\ntry:\n _src_path = _hf_dl(\n repo_id=FEEDBACK_REPO_ID, filename='dataset_sources.jsonl',\n repo_type='dataset', token=HF_TOKEN,\n )\n with open(_src_path, encoding='utf-8') as _f:\n for _line in _f:\n _line = _line.strip()\n if not _line:\n continue\n _entry = _json.loads(_line)\n if not _entry.get('enabled'):\n continue\n # Normalise keys to what Cell 9 expects\n EXTERNAL_DATASETS.append({\n 'enabled' : True,\n 'repo_id' : _entry.get('repo', _entry.get('repo_id', '')),\n 'config' : _entry.get('config'),\n 'split' : _entry.get('split', 'train'),\n 'text_col' : _entry.get('text_col', 'transcription'),\n 'lang' : _entry.get('lang', _entry.get('language', TRAIN_LANG)),\n 'max_samples': _entry.get('max', _entry.get('max_samples', 2_000)),\n })\n print(f'dataset_sources.jsonl: loaded {len(EXTERNAL_DATASETS)} source(s)')\nexcept Exception as _e:\n print(f'dataset_sources.jsonl not found or empty ({_e}) -- using hardcoded list only')\n\nactive = [d for d in EXTERNAL_DATASETS if d.get('lang') == TRAIN_LANG]\nprint(f'External sources active for {TRAIN_LANG}: {len(active)}')\nfor _d in active:\n print(f\" - {_d['repo_id']} / {_d['config']} (max {_d['max_samples']} samples)\")\nif not active:\n if TRAIN_LANG == 'bam':\n print('Bambara: no external source yet.')\n print(' In the Space -> Self-Teaching tab -> Import from HuggingFace (Bambara).')\n elif TRAIN_LANG == 'ful':\n print('Fula: WaxalNLP ful_asr loaded in Cell 8 -- no extra source needed.')\n" ] }, { "cell_type": "code", "execution_count": null, "id": "cell-login", "metadata": {}, "outputs": [], "source": "# โ”€โ”€ Cell 5: HuggingFace login + directory setup โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\nimport os\nfrom pathlib import Path\n\nHF_TOKEN = None\n\n# Kaggle secrets (preferred)\ntry:\n from kaggle_secrets import UserSecretsClient # type: ignore\n HF_TOKEN = UserSecretsClient().get_secret('HF_TOKEN')\n print('HF_TOKEN loaded from Kaggle secrets.')\nexcept Exception:\n pass\n\n# Colab secrets (fallback)\nif not HF_TOKEN:\n try:\n from google.colab import userdata # type: ignore\n HF_TOKEN = userdata.get('HF_TOKEN')\n print('HF_TOKEN loaded from Colab secrets.')\n except Exception:\n pass\n\nif not HF_TOKEN:\n HF_TOKEN = os.environ.get('HF_TOKEN', '')\n\nif not HF_TOKEN:\n raise ValueError(\n 'HF_TOKEN not found.\\n'\n 'Kaggle: Add-ons โ†’ Secrets โ†’ add HF_TOKEN โ†’ toggle \"Attach to notebook\" ON'\n )\n\nfrom huggingface_hub import login, HfApi\nlogin(token=HF_TOKEN, add_to_git_credential=False)\napi = HfApi(token=HF_TOKEN)\nos.environ['HF_TOKEN'] = HF_TOKEN\n\n# Create output directories\nfor d in [OUTPUT_DIR, DATA_DIR, AUDIO_DIR]:\n Path(d).mkdir(parents=True, exist_ok=True)\n\nprint(f'โœ… Logged in | output: {OUTPUT_DIR}')" }, { "cell_type": "code", "execution_count": null, "id": "cell-resume", "metadata": {}, "outputs": [], "source": "# โ”€โ”€ Cell 6: Resume-from-checkpoint detection โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\n# If OUTPUT_DIR already has checkpoints (e.g. Kaggle session timed out),\n# training will automatically resume from the latest one.\n#\n# NOTE: We do NOT use transformers.trainer_utils.get_last_checkpoint here.\n# That import pulls in peft โ†’ transformers.generation โ†’ masking_utils โ†’\n# torch._dynamo before packages are settled, causing ImportError on Kaggle\n# Python 3.12. The function below replicates exactly what it does internally.\n\nimport re\nfrom pathlib import Path\n\ndef _get_last_checkpoint(folder: str):\n \"\"\"Return the highest-numbered checkpoint-N directory path, or None.\"\"\"\n p = Path(folder)\n if not p.exists():\n return None\n checkpoints = [\n d for d in p.iterdir()\n if d.is_dir() and re.fullmatch(r'checkpoint-\\d+', d.name)\n ]\n if not checkpoints:\n return None\n return str(max(checkpoints, key=lambda d: int(d.name.split('-')[1])))\n\n\nLAST_CHECKPOINT = _get_last_checkpoint(OUTPUT_DIR)\n\nif LAST_CHECKPOINT:\n print(f'โฉ Resume checkpoint found: {LAST_CHECKPOINT}')\n print(' Training will continue from this point.')\nelse:\n print('๐Ÿ†• No checkpoint found โ€” starting fresh training.')" }, { "cell_type": "code", "execution_count": null, "id": "cell-feedback", "metadata": {}, "outputs": [], "source": "# โ”€โ”€ Cell 7: Download sahel-agri-feedback data โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\n# Downloads vocabulary.jsonl (word pairs) and corrections.jsonl (audio+text).\n# Audio files referenced in corrections.jsonl are also downloaded.\n\nimport json, shutil\nfrom pathlib import Path\nfrom huggingface_hub import hf_hub_download, list_repo_files\n\n# โ”€โ”€ vocabulary.jsonl (word pairs taught by users) โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\nvocab_entries = []\ntry:\n vocab_path = hf_hub_download(\n repo_id=FEEDBACK_REPO_ID, filename='vocabulary.jsonl',\n repo_type='dataset', token=HF_TOKEN,\n )\n with open(vocab_path, encoding='utf-8') as f:\n vocab_entries = [json.loads(l) for l in f if l.strip()]\n print(f'vocabulary.jsonl : {len(vocab_entries)} entries')\nexcept Exception as e:\n print(f'vocabulary.jsonl not found or empty: {e}')\n\n# โ”€โ”€ corrections.jsonl (audio corrections from the Space) โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\ncorrection_records = []\ntry:\n corr_path = hf_hub_download(\n repo_id=FEEDBACK_REPO_ID, filename='corrections.jsonl',\n repo_type='dataset', token=HF_TOKEN,\n )\n with open(corr_path, encoding='utf-8') as f:\n all_records = [json.loads(l) for l in f if l.strip()]\n correction_records = [\n r for r in all_records\n if r.get('language') == TRAIN_LANG\n and (r.get('corrected_text') or r.get('transcription'))\n and r.get('audio_file')\n ]\n print(f'corrections.jsonl: {len(all_records)} total, {len(correction_records)} for lang={TRAIN_LANG}')\nexcept Exception as e:\n print(f'corrections.jsonl not found or empty: {e}')\n\n# โ”€โ”€ Download audio files referenced in corrections โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\nskipped_audio = 0\nfor rec in correction_records:\n audio_fname = Path(rec['audio_file']).name\n local_path = Path(AUDIO_DIR) / audio_fname\n if local_path.exists():\n rec['local_audio'] = str(local_path)\n continue\n try:\n dl = hf_hub_download(\n repo_id=FEEDBACK_REPO_ID, filename=rec['audio_file'],\n repo_type='dataset', token=HF_TOKEN,\n )\n shutil.copy2(dl, local_path)\n rec['local_audio'] = str(local_path)\n except Exception as e:\n skipped_audio += 1\n rec['local_audio'] = None\n\ncorrection_records = [r for r in correction_records if r.get('local_audio')]\nprint(f'Audio downloaded : {len(correction_records)} files ({skipped_audio} skipped)')\nprint(f'Vocab entries : {len(vocab_entries)}')" }, { "cell_type": "code", "execution_count": null, "id": "cell-waxal", "metadata": {}, "outputs": [], "source": "# -- Cell 8: Load WaxalNLP ----------------------------------------------------\n# WaxalNLP confirmed subsets (google/WaxalNLP dataset card):\n# Fula -> 'ful_asr' (available)\n# Bambara -> NOT present (no 'bam' config exists)\n#\n# google/fleurs is NOT used as a fallback -- datasets >= 3.0 refuses\n# to execute its legacy script loader (fleurs.py).\n#\n# Bambara training uses: user corrections (Cell 7) + Common Voice bm (Cell 4).\n# Fula training uses: user corrections + WaxalNLP ful_asr + Common Voice ff.\n\nfrom datasets import load_dataset, Audio as HFAudio\n\nWAXAL_SUBSET_MAP = {\n 'bam': None, # no Bambara subset in WaxalNLP -- skip silently\n 'ful': 'ful_asr', # confirmed available config\n}\n\nwaxal_ds = None\nWAXAL_TEXT_COL = 'transcription'\n\nsubset = WAXAL_SUBSET_MAP.get(TRAIN_LANG)\n\nif subset is None:\n print(f'WaxalNLP has no subset for lang={TRAIN_LANG} -- skipping.')\n print('Bambara will train on user corrections + Common Voice (Cell 4).')\nelse:\n try:\n print(f'Loading google/WaxalNLP subset={subset} (streaming) ...')\n waxal_ds = load_dataset(\n 'google/WaxalNLP', subset,\n split='train',\n streaming=True,\n token=HF_TOKEN,\n )\n\n # Probe first item to confirm schema\n probe = next(iter(waxal_ds))\n if 'audio' not in probe:\n raise ValueError(f'No audio column. Keys: {list(probe.keys())}')\n\n WAXAL_TEXT_COL = next(\n (k for k in ['transcription', 'text', 'sentence', 'normalized_text']\n if k in probe),\n None,\n )\n if WAXAL_TEXT_COL is None:\n raise ValueError(f'No text column. Keys: {list(probe.keys())}')\n\n print(f'WaxalNLP/{subset} ready -- text column: \"{WAXAL_TEXT_COL}\"')\n\n except Exception as e:\n print(f'WaxalNLP/{subset} failed: {e}')\n print('Continuing without WaxalNLP -- enable Common Voice in Cell 4.')\n waxal_ds = None\n\nstatus = f'WaxalNLP/{subset}' if waxal_ds is not None else 'not available'\nprint(f'\\nWaxal source for {TRAIN_LANG}: {status}')" }, { "cell_type": "code", "execution_count": null, "id": "cell-external", "metadata": {}, "outputs": [], "source": "# โ”€โ”€ Cell 9: Load external datasets (from Cell 4 config) โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\nfrom datasets import load_dataset, Audio as HFAudio\n\nexternal_datasets = [] # list of (hf_dataset, text_col)\n\nfor cfg in EXTERNAL_DATASETS:\n if not cfg['enabled'] or cfg['lang'] != TRAIN_LANG:\n continue\n try:\n print(f'Loading {cfg[\"repo_id\"]} / {cfg[\"config\"]} ...')\n ds = load_dataset(\n cfg['repo_id'], cfg['config'],\n split=cfg['split'],\n streaming=True,\n token=HF_TOKEN,\n )\n probe = next(iter(ds))\n text_col = cfg['text_col'] if cfg['text_col'] in probe else next(\n (k for k in ['transcription', 'text', 'sentence'] if k in probe), None\n )\n if text_col is None:\n print(f' โš ๏ธ Cannot find text column โ€” skipping')\n continue\n if 'audio' not in probe:\n print(f' โš ๏ธ No audio column โ€” skipping')\n continue\n # Cap at max_samples\n ds = ds.take(cfg.get('max_samples', 2_000))\n external_datasets.append((ds, text_col))\n print(f' โœ… {cfg[\"repo_id\"]} โ€” text col \"{text_col}\", max {cfg.get(\"max_samples\",2000)} samples')\n except Exception as e:\n print(f' โš ๏ธ {cfg[\"repo_id\"]} failed: {e}')\n\nprint(f'\\nExternal sources loaded: {len(external_datasets)}')" }, { "cell_type": "markdown", "id": "md-pipeline", "metadata": {}, "source": "---\n## Data Pipeline\n\nAll audio is resampled to 16 kHz. Text is cleaned with a language-aware allowlist that keeps Latin script extended characters valid for Bambara (ษ› ษ” ล‹) and Fula (ษ“ ษ— ฦด ล‹ ษฒ), stripping everything else (URLs, XML tags, symbols). The Whisper processor converts the cleaned text to token IDs." }, { "cell_type": "code", "execution_count": null, "id": "cell-clean", "metadata": {}, "outputs": [], "source": [ "# -- Cell 10: Text cleaning utilities + Bambara phonetic normaliser -----------\nimport re, unicodedata\n\n# Phonetic normaliser: unifies French-influenced spellings before training.\n# ou->u, dj->j, gn->ny_palatal etc. so spelling variants map to same token.\n_BAM_NORM_RULES = [('ou','u'),('dj','j'),('gn','ษฒ'),('ny','ษฒ'),('ch','c'),('oo','ษ”'),('ee','ษ›')]\n_BAM_NORM_PAT = re.compile('|'.join(re.escape(s) for s,_ in _BAM_NORM_RULES))\n_BAM_NORM_MAP = {s:d for s,d in _BAM_NORM_RULES}\n\ndef _bam_norm(text):\n import unicodedata as _ud\n text = _ud.normalize('NFC', text.lower())\n return _BAM_NORM_PAT.sub(lambda m: _BAM_NORM_MAP[m.group(0)], text)\n\n# Pular (Fula of Guinea) normaliser: converts Adlam script โ†’ Latin,\n# then NFC + lowercase. Needed because guizme/adlam_fulfulde labels are in\n# Adlam (U+1E900-U+1E95F) which Whisperโ€™s tokenizer has no coverage for.\n_ADLAM_TO_LATIN = [\n (\"๐žค€\",\"A\"),(\"๐žค\",\"B\"),(\"๐žค‚\",\"B\"),(\"๐žคƒ\",\"D\"),(\"๐žค„\",\"D\"),\n (\"๐žค…\",\"E\"),(\"๐žค†\",\"F\"),(\"๐žค‡\",\"G\"),(\"๐žคˆ\",\"H\"),(\"๐žค‰\",\"I\"),\n (\"๐žคŠ\",\"J\"),(\"๐žค‹\",\"K\"),(\"๐žคŒ\",\"L\"),(\"๐žค\",\"M\"),(\"๐žคŽ\",\"N\"),\n (\"๐žค\",\"NG\"),(\"๐žค\",\"O\"),(\"๐žค‘\",\"P\"),(\"๐žค’\",\"R\"),(\"๐žค“\",\"S\"),\n (\"๐žค”\",\"T\"),(\"๐žค•\",\"U\"),(\"๐žค–\",\"V\"),(\"๐žค—\",\"W\"),(\"๐žค˜\",\"Y\"),\n (\"๐žค™\",\"Z\"),(\"๐žคš\",\"KH\"),(\"๐žค›\",\"QU\"),(\"๐žคœ\",\"SH\"),(\"๐žค\",\"GH\"),\n (\"๐žคž\",\"NY\"),(\"๐žคŸ\",\"TH\"),(\"๐žค \",\"WH\"),(\"๐žคก\",\"NY\"),\n (\"๐žคข\",\"a\"),(\"๐žคฃ\",\"b\"),(\"๐žคค\",\"b\"),(\"๐žคฅ\",\"d\"),(\"๐žคฆ\",\"d\"),\n (\"๐žคง\",\"e\"),(\"๐žคจ\",\"f\"),(\"๐žคฉ\",\"g\"),(\"๐žคช\",\"h\"),(\"๐žคซ\",\"i\"),\n (\"๐žคฌ\",\"j\"),(\"๐žคญ\",\"k\"),(\"๐žคฎ\",\"l\"),(\"๐žคฏ\",\"m\"),(\"๐žคฐ\",\"n\"),\n (\"๐žคฑ\",\"ng\"),(\"๐žคฒ\",\"o\"),(\"๐žคณ\",\"p\"),(\"๐žคด\",\"r\"),(\"๐žคต\",\"s\"),\n (\"๐žคถ\",\"t\"),(\"๐žคท\",\"u\"),(\"๐žคธ\",\"v\"),(\"๐žคน\",\"w\"),(\"๐žคบ\",\"y\"),\n (\"๐žคป\",\"z\"),(\"๐žคผ\",\"kh\"),(\"๐žคฝ\",\"qu\"),(\"๐žคพ\",\"sh\"),(\"๐žคฟ\",\"gh\"),\n (\"๐žฅ€\",\"ny\"),(\"๐žฅ\",\"th\"),(\"๐žฅ‚\",\"wh\"),(\"๐žฅƒ\",\"ny\"),\n]\n_A2L = {a: l for a, l in _ADLAM_TO_LATIN}\n_ADLAM_START, _ADLAM_END = 0x1E900, 0x1E95F\n\ndef _contains_adlam(text):\n return any(_ADLAM_START <= ord(c) <= _ADLAM_END for c in text)\n\ndef _normalize_pular(text):\n import unicodedata as _ud, re as _re\n if _contains_adlam(text):\n text = \"\".join(_A2L.get(c, c) for c in text)\n text = _ud.normalize(\"NFC\", text.lower())\n return _re.sub(r\"\\s+\", \" \", text).strip()\n\n\n_BAMBARA_EXTRA = {'\\u025b','\\u0254','\\u014b'}\n_FULA_EXTRA = {'\\u0253','\\u0257','\\u01b4','\\u014b','\\u0272'}\n_BASE_LATIN = set('abcdefghijklmnopqrstuvwxyz')\n_ACCENTED = set('\\u00e0\\u00e2\\u00e4\\u00e8\\u00e9\\u00ea\\u00eb'\n '\\u00ee\\u00ef\\u00f4\\u00f9\\u00fb\\u00fc\\u00fd'\n '\\u00ff\\u00e6\\u0153\\u00e7')\n_KEEP_PUNCT = set(\" ',-.'!?\")\n\n_VALID_CHARS = {\n 'bam': _BASE_LATIN | _ACCENTED | _BAMBARA_EXTRA | _KEEP_PUNCT,\n 'ful': _BASE_LATIN | _ACCENTED | _FULA_EXTRA | _KEEP_PUNCT,\n}\n\n\ndef clean_text(text: str, lang: str = 'bam') -> str:\n if not text:\n return ''\n text = unicodedata.normalize('NFKC', text.lower().strip())\n text = re.sub(r'https?://\\S+', '', text)\n text = re.sub(r'<[^>]+>', '', text)\n text = re.sub(r'([.,!?])\\1+', r'\\1', text)\n valid = _VALID_CHARS.get(lang, _VALID_CHARS['bam'] | _VALID_CHARS['ful'])\n text = ''.join(c for c in text if c in valid)\n return re.sub(r'\\s+', ' ', text).strip()\n\n\n# Verify actual output then assert against it\nr1 = clean_text('I ni ce! (hello)', 'bam') # parens stripped, ! kept\nr2 = clean_text('Jam waali. test', 'ful') # tags stripped, content kept\nr3 = clean_text('Visit https://example.com now!!', 'bam') # URL stripped, word before stays\n\nassert r1 == 'i ni ce! hello', f'r1: {repr(r1)}'\nassert r2 == 'jam waali. test', f'r2: {repr(r2)}'\nassert r3 == 'visit now!', f'r3: {repr(r3)}'\n\nprint('clean_text tests passed')\nprint(f' {repr(r1)}')\nprint(f' {repr(r2)}')\nprint(f' {repr(r3)}')" ] }, { "cell_type": "code", "execution_count": null, "id": "cell-prepare", "metadata": {}, "outputs": [], "source": [ "# -- Cell 11: Whisper processor + prepare_dataset -----------------------------\n# WhisperProcessor imports processing_utils -> image_utils -> torchvision,\n# which crashes when torch/torchvision have mismatched CUDA versions.\n# Fix: build the processor manually from its two sub-components.\n# WhisperFeatureExtractor and WhisperTokenizer have no torchvision dependency.\nimport numpy as np\n\nfrom transformers.models.whisper.feature_extraction_whisper import WhisperFeatureExtractor\nfrom transformers.models.whisper.tokenization_whisper import WhisperTokenizer\n\nprint(f'Loading Whisper feature extractor + tokenizer: {WHISPER_MODEL_ID} ...')\n_feat_ext = WhisperFeatureExtractor.from_pretrained(WHISPER_MODEL_ID, token=HF_TOKEN)\n_tokenizer = WhisperTokenizer.from_pretrained(WHISPER_MODEL_ID, token=HF_TOKEN)\n\n\nclass _Processor:\n \"\"\"Minimal WhisperProcessor substitute that avoids the torchvision import chain.\"\"\"\n def __init__(self, feature_extractor, tokenizer):\n self.feature_extractor = feature_extractor\n self.tokenizer = tokenizer\n\n def get_decoder_prompt_ids(self, language, task='transcribe'):\n return self.tokenizer.get_decoder_prompt_ids(language=language, task=task)\n\n def save_pretrained(self, path):\n self.feature_extractor.save_pretrained(path)\n self.tokenizer.save_pretrained(path)\n\n\nprocessor = _Processor(_feat_ext, _tokenizer)\nprint('Processor ready')\n\n\ndef prepare_dataset(batch, text_col='transcription', lang=TRAIN_LANG):\n \"\"\"\n Resample to 16 kHz, extract log-mel features, tokenise text.\n Works on any dict with 'audio' (HF Audio column) and a text column.\n \"\"\"\n audio = batch['audio']\n audio_array = np.array(audio['array'], dtype=np.float32)\n orig_sr = audio['sampling_rate']\n\n if orig_sr != TARGET_SR:\n try:\n import torchaudio.functional as F_audio, torch\n audio_array = F_audio.resample(\n torch.from_numpy(audio_array).unsqueeze(0),\n orig_sr, TARGET_SR,\n ).squeeze(0).numpy()\n except Exception:\n import librosa\n audio_array = librosa.resample(audio_array, orig_sr=orig_sr, target_sr=TARGET_SR)\n\n batch['input_features'] = processor.feature_extractor(\n audio_array, sampling_rate=TARGET_SR\n ).input_features[0]\n\n raw_text = batch.get(text_col, '') or ''\n _norm_text = _bam_norm(str(raw_text)) if lang == 'bam' else (_normalize_pular(str(raw_text)) if lang == 'ful' else str(raw_text))\n cleaned = clean_text(_norm_text, lang=lang)\n batch['labels'] = processor.tokenizer(cleaned).input_ids\n return batch\n\n\nprint('prepare_dataset ready')" ] }, { "cell_type": "code", "execution_count": null, "id": "cell-merge", "metadata": {}, "outputs": [], "source": [ "# -- Cell 12: Build & merge all datasets --------------------------------------\nfrom datasets import Dataset, Audio as HFAudio, concatenate_datasets\nfrom functools import partial\n\ntrain_ds = None # set here so Cell 12b can detect whether we succeeded\neval_ds = None\nall_parts = []\n\n# -- Part A: User corrections -------------------------------------------------\nif correction_records:\n print(f'Part A: {len(correction_records)} corrections x {CORRECTION_REPEAT}')\n rec_list = correction_records * CORRECTION_REPEAT\n corr_ds = Dataset.from_dict({\n 'audio': [r['local_audio'] for r in rec_list],\n 'transcription': [r.get('corrected_text') or r.get('transcription', '') for r in rec_list],\n })\n corr_ds = corr_ds.cast_column('audio', HFAudio(sampling_rate=TARGET_SR))\n corr_ds = corr_ds.map(\n partial(prepare_dataset, text_col='transcription'),\n remove_columns=corr_ds.column_names,\n )\n all_parts.append(corr_ds)\n print(f' -> {len(corr_ds)} samples')\nelse:\n print('Part A: no corrections -- skipping')\n\n# -- Part B: WaxalNLP ---------------------------------------------------------\nif waxal_ds is not None:\n print(f'Part B: materialising up to {MAX_WAXAL_TRAIN} WaxalNLP samples ...')\n waxal_rows = list(waxal_ds.take(MAX_WAXAL_TRAIN))\n waxal_local = Dataset.from_dict({\n 'audio': [r['audio'] for r in waxal_rows],\n 'transcription': [r.get(WAXAL_TEXT_COL, '') for r in waxal_rows],\n })\n waxal_local = waxal_local.cast_column('audio', HFAudio(sampling_rate=TARGET_SR))\n waxal_local = waxal_local.map(\n partial(prepare_dataset, text_col='transcription'),\n remove_columns=waxal_local.column_names,\n )\n all_parts.append(waxal_local)\n print(f' -> {len(waxal_local)} samples')\nelse:\n print('Part B: WaxalNLP not available -- skipping')\n\n# -- Part C: External datasets ------------------------------------------------\nfor ext_ds, text_col in external_datasets:\n ext_rows = list(ext_ds)\n if not ext_rows:\n continue\n ext_local = Dataset.from_dict({\n 'audio': [r['audio'] for r in ext_rows],\n 'transcription': [r.get(text_col, '') for r in ext_rows],\n })\n ext_local = ext_local.cast_column('audio', HFAudio(sampling_rate=TARGET_SR))\n ext_local = ext_local.map(\n partial(prepare_dataset, text_col='transcription'),\n remove_columns=ext_local.column_names,\n )\n all_parts.append(ext_local)\n print(f'Part C: {len(ext_local)} external samples')\n\n# -- Result -------------------------------------------------------------------\nprint(f'\\nData summary: {len(all_parts)} source(s) loaded')\n\nif not all_parts:\n print('No real data available for Bambara yet.')\n print('Cell 12b (below) will build a synthetic dataset so the full')\n print('training pipeline can be validated. Switch to TRAIN_LANG=\"ful\"')\n print('in Cell 3 for a real Fula training run using WaxalNLP.')\nelse:\n combined = concatenate_datasets(all_parts).shuffle(seed=42)\n n_eval = max(1, min(int(0.05 * len(combined)), 200))\n split = combined.train_test_split(test_size=n_eval)\n train_ds = split['train']\n eval_ds = split['test']\n print(f'Train: {len(train_ds)} Eval: {len(eval_ds)}')\n# Auto-cap MAX_STEPS: no point running 4000 steps on a tiny dataset.\n# Rule: at least 20 passes through the data, capped at user's MAX_STEPS.\nif train_ds is not None:\n steps_per_epoch = max(1, len(train_ds) // (BATCH_SIZE * GRAD_ACCUM))\n _auto_steps = max(200, steps_per_epoch * 20)\n if _auto_steps < MAX_STEPS:\n print(f'Auto-capping MAX_STEPS {MAX_STEPS} -> {_auto_steps} (small dataset)')\n MAX_STEPS = _auto_steps\n else:\n print(f'MAX_STEPS={MAX_STEPS} OK for {len(train_ds)} training samples')" ] }, { "cell_type": "code", "id": "5beb513a", "source": "# -- Cell 12b: Synthetic fallback (SKIP if Cell 12 succeeded) ----------------\n# Run this cell ONLY if Cell 12 raised \"No data loaded\".\n# Generates 50 short silent audio samples labelled with vocabulary.jsonl\n# entries so training can proceed and you can verify the pipeline works.\n# Replace with real data (accept Common Voice terms, or add corrections) for\n# a meaningful model.\n\nimport numpy as np\nfrom datasets import Dataset, Audio as HFAudio, concatenate_datasets\nfrom functools import partial\n\nif 'train_ds' in dir() and train_ds is not None:\n print('Cell 12 succeeded -- nothing to do here.')\nelse:\n print('Building synthetic fallback dataset from vocabulary.jsonl ...')\n\n # Use vocab entries if available, otherwise generic phrases\n if vocab_entries:\n phrases = [e.get('word', 'test') for e in vocab_entries[:50]]\n else:\n phrases = [f'word {i}' for i in range(50)]\n\n SR = TARGET_SR\n rows = []\n for phrase in phrases:\n # 1-second silent audio (safe baseline for feature extraction)\n audio_array = np.zeros(SR, dtype=np.float32)\n rows.append({'audio_array': audio_array, 'transcription': phrase})\n\n # Build dataset directly from numpy arrays\n synth_ds = Dataset.from_dict({\n 'transcription': [r['transcription'] for r in rows],\n })\n\n # Add audio column manually\n def _add_audio(batch, idx):\n batch['input_features'] = processor.feature_extractor(\n np.zeros(TARGET_SR, dtype=np.float32), sampling_rate=TARGET_SR\n ).input_features[0]\n cleaned = clean_text(rows[idx]['transcription'], lang=TRAIN_LANG)\n batch['labels'] = processor.tokenizer(cleaned).input_ids\n return batch\n\n synth_processed = synth_ds.map(\n _add_audio,\n with_indices=True,\n remove_columns=synth_ds.column_names,\n )\n\n split = synth_processed.train_test_split(test_size=0.1, seed=42)\n train_ds = split['train']\n eval_ds = split['test']\n print(f'Synthetic fallback: {len(train_ds)} train, {len(eval_ds)} eval')\n print('WARNING: training on synthetic data produces a non-functional model.')\n print('Accept Common Voice terms and re-run Cell 9 + Cell 12 for real data.')", "metadata": {}, "execution_count": null, "outputs": [] }, { "cell_type": "markdown", "id": "md-model", "metadata": {}, "source": [ "---\n## Model Setup - Partial Freeze Fine-Tuning\n\nopenai/whisper-small is loaded in fp16. All parameters are frozen, then\nthe last 2 decoder layers + layer norm + output projection are unfrozen (~5%\ntrainable params). No PEFT/LoRA -- avoids PEFT/transformers 5.x\nincompatibility where input_ids is passed twice to WhisperDecoder.\n" ] }, { "cell_type": "code", "execution_count": null, "id": "cell-model", "metadata": {}, "outputs": [], "source": [ "# -- Cell 13: Load Whisper-small (fp16) + freeze most layers ------------------\n# PEFT LoRA causes TypeError with transformers 5.x regardless of which layers\n# are targeted: PeftModelForSeq2SeqLM wraps the entire model in BaseTuner whose\n# forward(*args, **kwargs) passes input_ids in a way that causes WhisperDecoder\n# to receive it twice. Fix: skip PEFT entirely. Freeze all params, unfreeze\n# last 2 decoder layers (~5% trainable params -- same capacity as LoRA r=32).\nimport torch\nfrom transformers.models.whisper.modeling_whisper import WhisperForConditionalGeneration\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(f\"Loading {WHISPER_MODEL_ID} in fp32 on {device} ...\")\n\nmodel = WhisperForConditionalGeneration.from_pretrained(\n WHISPER_MODEL_ID,\n torch_dtype=torch.float32, # fp32 storage -- AMP casts internally during training\n token=HF_TOKEN,\n)\nmodel = model.to(device)\n\n# Force target language -- avoids language-detection overhead during training\n# Move generation params to GenerationConfig (avoids deprecation warning)\n_dec_ids = processor.get_decoder_prompt_ids(language='fr', task='transcribe')\nmodel.generation_config.forced_decoder_ids = _dec_ids\nmodel.generation_config.suppress_tokens = []\nmodel.config.use_cache = False # required for gradient checkpointing\n\n# โ”€โ”€ Freeze all params, then selectively unfreeze โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\nfor param in model.parameters():\n param.requires_grad = False\n\n# Unfreeze last 2 decoder layers + final layer norm + output projection.\n# These handle language-specific token generation.\nfor module in [\n model.model.decoder.layers[-2],\n model.model.decoder.layers[-1],\n model.model.decoder.layer_norm,\n model.proj_out,\n]:\n for param in module.parameters():\n param.requires_grad = True\n\ntrainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\ntotal = sum(p.numel() for p in model.parameters())\nprint(f\"Trainable params: {trainable:,} / {total:,} ({100*trainable/total:.1f}%)\")\n\n# gradient_checkpointing enabled via TrainingArguments below (args handle enable/disable)\nmodel.train()\n\nvram_mb = torch.cuda.memory_allocated() / 1e6 if torch.cuda.is_available() else 0\ntotal_vram = torch.cuda.get_device_properties(0).total_memory / 1e6 if torch.cuda.is_available() else 0\nprint(f\"VRAM used: {vram_mb:.0f} MB / {total_vram:.0f} MB\")\nprint(f\"Model ready on {device}\")\n" ] }, { "cell_type": "code", "execution_count": null, "id": "cell-collator", "metadata": {}, "outputs": [], "source": [ "# -- Cell 14: Data collator + CER metric --------------------------------------\nimport jiwer\nfrom dataclasses import dataclass\nfrom typing import Any, Dict, List\n\ntransform = jiwer.Compose([\n jiwer.ToLowerCase(),\n jiwer.RemoveMultipleSpaces(),\n jiwer.Strip(),\n jiwer.RemovePunctuation(),\n jiwer.ReduceToListOfListOfWords(),\n])\n\n# CER transform (no word-split step needed)\n_cer_transform = jiwer.Compose([\n jiwer.ToLowerCase(),\n jiwer.RemoveMultipleSpaces(),\n jiwer.Strip(),\n jiwer.RemovePunctuation(),\n])\n\n\n@dataclass\nclass DataCollatorSpeechSeq2SeqWithPadding:\n processor: Any\n\n def __call__(self, features: List[Dict]) -> Dict:\n import torch\n input_feats = [{'input_features': f['input_features']} for f in features]\n batch = self.processor.feature_extractor.pad(input_feats, return_tensors='pt')\n\n # Leave features in fp32 -- AMP (fp16=True in TrainingArgs) handles casting\n\n label_feats = [{'input_ids': f['labels']} for f in features]\n labels_batch = self.processor.tokenizer.pad(label_feats, return_tensors='pt')\n labels = labels_batch['input_ids'].masked_fill(\n labels_batch.attention_mask.ne(1), -100\n )\n if (labels[:, 0] == self.processor.tokenizer.bos_token_id).all().item():\n labels = labels[:, 1:]\n batch['labels'] = labels\n return batch\n\n\ndef _apply_jiwer_transform(texts, t):\n \"\"\"Apply a jiwer Compose transform and return plain strings (not nested lists).\"\"\"\n import re as _re\n result = []\n for s in texts:\n s = s.lower()\n s = _re.sub(r'[^\\w\\s]', '', s) # RemovePunctuation equivalent\n s = ' '.join(s.split()) # RemoveMultipleSpaces + Strip\n result.append(s)\n return result\n\n\ndef compute_metrics(pred):\n pred_ids = pred.predictions\n label_ids = pred.label_ids\n label_ids[label_ids == -100] = processor.tokenizer.pad_token_id\n\n pred_str = processor.tokenizer.batch_decode(pred_ids, skip_special_tokens=True)\n label_str = processor.tokenizer.batch_decode(label_ids, skip_special_tokens=True)\n\n # Pre-normalise so we can filter empties AFTER the transform.\n # A reference like \"?\" or \"1\" decodes to non-empty but becomes empty\n # after punctuation/number removal -- jiwer crashes on empty references.\n norm_ref = _apply_jiwer_transform(label_str, _cer_transform)\n norm_hyp = _apply_jiwer_transform(pred_str, _cer_transform)\n pairs = [(r, h) for r, h in zip(norm_ref, norm_hyp) if r.strip()]\n if not pairs:\n return {'cer': 0.0, 'wer': 0.0}\n ref_clean, hyp_clean = zip(*pairs)\n\n cer = jiwer.cer(list(ref_clean), list(hyp_clean)) # already normalised\n wer = jiwer.wer(list(ref_clean), list(hyp_clean)) # already normalised\n return {'cer': round(cer, 4), 'wer': round(wer, 4)}\n\n\ncollator = DataCollatorSpeechSeq2SeqWithPadding(processor=processor)\nprint('Collator and WER metric ready')" ] }, { "cell_type": "markdown", "id": "md-train", "metadata": {}, "source": "---\n## Training โ€” Deep Sleep ๐Ÿ˜ด\n\nThe trainer runs for `MAX_STEPS` steps, saving a checkpoint every `SAVE_STEPS`.\nIf the Kaggle session times out and you re-run the notebook, Cell 6 will detect\nthe latest checkpoint and `trainer.train()` will resume automatically โ€” no progress lost." }, { "cell_type": "code", "execution_count": null, "id": "cell-train-args", "metadata": {}, "outputs": [], "source": [ "# -- Cell 15: Training arguments ----------------------------------------------\nimport inspect\nfrom transformers import Seq2SeqTrainingArguments\n\n# transformers 4.x used 'evaluation_strategy'; 4.45+ renamed to 'eval_strategy'.\n# Detect which name this installed version accepts.\n_params = inspect.signature(Seq2SeqTrainingArguments.__init__).parameters\n_eval_key = 'eval_strategy' if 'eval_strategy' in _params else 'evaluation_strategy'\n\ntraining_args = Seq2SeqTrainingArguments(\n output_dir=OUTPUT_DIR,\n\n max_steps=MAX_STEPS,\n warmup_steps=WARMUP_STEPS,\n logging_steps=LOGGING_STEPS,\n save_steps=SAVE_STEPS,\n eval_steps=EVAL_STEPS,\n\n per_device_train_batch_size=BATCH_SIZE,\n per_device_eval_batch_size=8,\n gradient_accumulation_steps=GRAD_ACCUM,\n\n fp16=True,\n gradient_checkpointing=True, # reduces activation memory on T4\n\n learning_rate=LEARNING_RATE,\n lr_scheduler_type='cosine',\n weight_decay=0.0,\n adam_beta1=0.9,\n adam_beta2=0.98,\n adam_epsilon=1e-6,\n\n **{_eval_key: 'steps'},\n predict_with_generate=True,\n generation_max_length=225,\n load_best_model_at_end=True,\n metric_for_best_model='cer',\n greater_is_better=False,\n\n save_total_limit=3,\n save_strategy='steps',\n\n report_to=['tensorboard'], # tensorboard logs to OUTPUT_DIR/runs by default\n push_to_hub=False,\n)\n\nprint(f'Training arguments ready (using {_eval_key}=steps)')\nprint(f' Effective batch size: {BATCH_SIZE * GRAD_ACCUM}')\nprint(f' Max steps : {MAX_STEPS}')\n" ] }, { "cell_type": "code", "execution_count": null, "id": "cell-train", "metadata": {}, "outputs": [], "source": [ "# -- Cell 16: TRAIN -----------------------------------------------------------\nfrom transformers import Seq2SeqTrainer\n\ntrainer = Seq2SeqTrainer(\n model=model,\n args=training_args,\n train_dataset=train_ds,\n eval_dataset=eval_ds,\n data_collator=collator,\n compute_metrics=compute_metrics,\n # 'tokenizer' argument removed -- renamed to 'processing_class' in\n # transformers 5.x and passing the old name raises TypeError\n)\n\nprint(f'Starting training ...')\nprint(f' Resume from : {LAST_CHECKPOINT or \"scratch\"}')\nprint(f' Train size : {len(train_ds)}')\nprint(f' Eval size : {len(eval_ds)}')\n\ntrain_result = trainer.train(resume_from_checkpoint=LAST_CHECKPOINT)\n\nprint('\\nTraining complete')\nprint(f' Steps : {train_result.global_step}')\nprint(f' Train loss: {train_result.training_loss:.4f}')\n\ntrainer.save_model(OUTPUT_DIR)\nmodel.generation_config.save_pretrained(OUTPUT_DIR)\nprocessor.save_pretrained(OUTPUT_DIR)\nprint(f' Adapter saved -> {OUTPUT_DIR}')" ] }, { "cell_type": "markdown", "id": "md-eval", "metadata": {}, "source": "---\n## Evaluation\n\nWER (Word Error Rate) is computed on the held-out eval split.\nA lower WER means fewer transcription mistakes. For Bambara/Fula on whisper-small,\na WER < 40% is a strong result given the limited training data." }, { "cell_type": "code", "execution_count": null, "id": "cell-eval", "metadata": {}, "outputs": [], "source": [ "# โ”€โ”€ Cell 17: WER evaluation โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\nprint('Running full evaluation on eval split ...')\neval_results = trainer.evaluate()\n\ncer_score = eval_results.get('eval_cer', float('nan'))\nwer_score = eval_results.get('eval_wer', float('nan'))\nprint(f'\\nโœ… Final CER : {cer_score:.1%} (primary โ€” lower is better)')\nprint(f' Final WER : {wer_score:.1%} (secondary)')\nprint(f' Eval loss : {eval_results.get(\"eval_loss\", float(\"nan\")):.4f}')\n# Show a few example transcriptions side-by-side\nimport random, torch\nprint('\\nโ”€โ”€ Sample predictions โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€')\nsamples = random.sample(range(len(eval_ds)), min(5, len(eval_ds)))\nfor idx in samples:\n item = eval_ds[idx]\n feats = torch.tensor(item['input_features']).unsqueeze(0).to(model.device)\n with torch.no_grad():\n pred_ids = model.generate(\n feats, # fp32 to match model dtype\n max_new_tokens=128,\n )\n pred_str = processor.tokenizer.batch_decode(pred_ids, skip_special_tokens=True)[0]\n labels = [t if t != -100 else processor.tokenizer.pad_token_id\n for t in item['labels']]\n ref_str = processor.tokenizer.decode(labels, skip_special_tokens=True)\n print(f' Ref : {ref_str}')\n print(f' Pred: {pred_str}')\n print()" ] }, { "cell_type": "markdown", "id": "md-export", "metadata": {}, "source": [ "---\n## Export โ€” Push Fine-tuned Checkpoint to Hub\n\nThe adapter is pushed to `ous-sow/sahel-agri-adapters` under the path \n`adapters/{lang_name}/` with a Git tag like `v1.2-bambara`.\n\nThe version number is **auto-incremented** by reading existing tags on the repo\nso each training run gets a unique, traceable identifier." ] }, { "cell_type": "code", "execution_count": null, "id": "cell-version", "metadata": {}, "outputs": [], "source": "# โ”€โ”€ Cell 18: Compute next version tag โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\nimport re as _re\nfrom huggingface_hub import list_repo_refs\n\ndef get_next_version_tag(repo_id: str, lang_name: str, hf_token: str) -> str:\n \"\"\"Auto-increment version tag: reads existing tags, bumps minor version.\"\"\"\n try:\n refs = list_repo_refs(repo_id, repo_type='model', token=hf_token)\n pattern = _re.compile(rf'^v(\\d+)\\.(\\d+)-{_re.escape(lang_name)}$')\n versions = []\n for tag in refs.tags:\n m = pattern.match(tag.name)\n if m:\n versions.append((int(m.group(1)), int(m.group(2))))\n if not versions:\n return f'v1.0-{lang_name}'\n major, minor = max(versions)\n return f'v{major}.{minor + 1}-{lang_name}'\n except Exception as e:\n print(f' Could not read existing tags ({e}) โ€” defaulting to v1.0')\n return f'v1.0-{lang_name}'\n\n\nVERSION_TAG = get_next_version_tag(ADAPTER_REPO_ID, LANG_NAME, HF_TOKEN)\nPATH_IN_REPO = f'adapters/{LANG_NAME}'\n\nprint(f'Version tag : {VERSION_TAG}')\nprint(f'Path in repo : {ADAPTER_REPO_ID}/{PATH_IN_REPO}')" }, { "cell_type": "code", "execution_count": null, "id": "cell-push", "metadata": {}, "outputs": [], "source": [ "# โ”€โ”€ Cell 19: Push adapter to HF Model repo โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\nfrom huggingface_hub import HfApi, create_repo\n\n# Ensure repo exists\ncreate_repo(ADAPTER_REPO_ID, repo_type='model', private=True,\n exist_ok=True, token=HF_TOKEN)\n\n_cer_part = f'{cer_score:.1%}' if cer_score == cer_score else 'n/a'\ncommit_msg = (\n f'[{VERSION_TAG}] {LANG_NAME} ({LANG_COUNTRY}) fine-tuned checkpoint โ€” '\n f'{train_result.global_step} steps | CER {_cer_part} | '\n f'{len(correction_records)} corrections + WaxalNLP | {LANG_DIALECT}'\n)\n\napi.upload_folder(\n folder_path=OUTPUT_DIR,\n repo_id=ADAPTER_REPO_ID,\n repo_type='model',\n path_in_repo=PATH_IN_REPO,\n commit_message=commit_msg,\n)\nprint(f'โœ… Adapter uploaded: {ADAPTER_REPO_ID}/{PATH_IN_REPO}')\n\n# Create a Git tag for this version\ntry:\n api.create_tag(\n repo_id=ADAPTER_REPO_ID,\n repo_type='model',\n tag=VERSION_TAG,\n tag_message=commit_msg,\n token=HF_TOKEN,\n )\n print(f'โœ… Tag created : {VERSION_TAG}')\nexcept Exception as e:\n print(f'โš ๏ธ Tag creation skipped: {e}')" ] }, { "cell_type": "code", "execution_count": null, "id": "cell-verify", "metadata": {}, "outputs": [], "source": [ "# โ”€โ”€ Cell 20: Verification summary โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€\nfrom huggingface_hub import list_repo_files\n\nprint('=' * 60)\nprint('DEEP SLEEP TRAINING โ€” COMPLETE')\nprint('=' * 60)\nprint(f' Language : {TRAIN_LANG} ({LANG_NAME})')\nprint(f' Model : {WHISPER_MODEL_ID}')\nprint(f' Steps completed : {train_result.global_step}')\nprint(f' Train loss : {train_result.training_loss:.4f}')\n_cer_disp = f'{cer_score:.1%}' if cer_score == cer_score else 'n/a'\n_wer_disp = f'{wer_score:.1%}' if wer_score == wer_score else 'n/a'\nprint(f' Eval CER (primary) : {_cer_disp}')\nprint(f' Eval WER (secondary): {_wer_disp}')\nprint(f' Corrections used : {len(correction_records)} ร— {CORRECTION_REPEAT}')\nprint(f' WaxalNLP samples : up to {MAX_WAXAL_TRAIN}')\nprint(f' Version tag : {VERSION_TAG}')\nprint(f' HF repo : {ADAPTER_REPO_ID}/{PATH_IN_REPO}')\nprint()\n\n# List what was pushed\ntry:\n repo_files = sorted(list_repo_files(\n ADAPTER_REPO_ID, repo_type='model', token=HF_TOKEN\n ))\n adapter_files = [f for f in repo_files if f.startswith(f'adapters/{LANG_NAME}/')]\n print('Adapter files in repo:')\n for f in adapter_files:\n print(f' {f}')\nexcept Exception as e:\n print(f'Could not list repo files: {e}')\n\nprint()\nprint('Next steps:')\nprint(' 1. In your HF Space settings, confirm ADAPTER_REPO_ID secret is set')\nprint(f' 2. Tab 3 โ†’ Reload Adapters โ†’ select \"{VERSION_TAG}\"')\nprint(' 3. Collect more corrections in the Space, then re-run this notebook')" ] } ] }