{ "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)\n\nprint(f'Language : {TRAIN_LANG} ({LANG_NAME})')\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# Data source reality check (as of 2026):\n#\n# Bambara ASR:\n# - google/WaxalNLP -> no 'bam' subset exists\n# - Common Voice 'bm' -> moved to Mozilla Data Collective (not on HF)\n# - PRIMARY SOURCE -> user corrections in sahel-agri-feedback\n#\n# Fula ASR:\n# - google/WaxalNLP 'ful_asr' -> AVAILABLE (handled in Cell 8)\n# - Common Voice 'ff' -> moved to Mozilla Data Collective (not on HF)\n#\n# To add a new dataset later, add an entry with enabled=True.\n# Any HF dataset with an 'audio' column and a text column works.\n\nEXTERNAL_DATASETS = [\n # Example -- uncomment and set enabled=True when a Bambara HF dataset appears:\n # {\n # 'enabled' : False,\n # 'repo_id' : 'MALIBA-AI/bambara-asr', # check HF for availability\n # 'config' : None,\n # 'split' : 'train',\n # 'text_col' : 'transcription',\n # 'lang' : 'bam',\n # 'max_samples': 5_000,\n # },\n]\n\nactive = [d for d in EXTERNAL_DATASETS if d.get('enabled') and d.get('lang') == TRAIN_LANG]\nprint(f'External sources active for {TRAIN_LANG}: {len(active)}')\nif not active:\n if TRAIN_LANG == 'bam':\n print('Bambara: no external HF dataset available.')\n print(' Training will use user corrections from sahel-agri-feedback.')\n print(' Collect corrections via the Space to grow this dataset over time.')\n elif TRAIN_LANG == 'ful':\n print('Fula: WaxalNLP ful_asr loaded in Cell 8 -- no extra external source needed.')" }, { "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 -----------------------------------------\nimport re, unicodedata\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 cleaned = clean_text(str(raw_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 fp16 on {device} ...\")\n\nmodel = WhisperForConditionalGeneration.from_pretrained(\n WHISPER_MODEL_ID,\n torch_dtype=torch.float16,\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# Enable gradient checkpointing -- reduces activation memory at slight speed cost\nmodel.gradient_checkpointing_enable()\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 + WER 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\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 # Cast to fp16 to match the model -- avoids dtype mismatch in conv1\n batch['input_features'] = batch['input_features'].to(torch.float16)\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 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 wer = jiwer.wer(label_str, pred_str,\n hypothesis_transform=transform,\n reference_transform=transform)\n return {'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=False,\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='wer',\n greater_is_better=False,\n\n save_total_limit=3,\n save_strategy='steps',\n\n report_to=['tensorboard'],\n logging_dir=f'{OUTPUT_DIR}/logs',\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\nwer_score = eval_results.get('eval_wer', float('nan'))\nprint(f'\\nšŸ“Š Final WER : {wer_score:.1%}')\nprint(f' Eval loss : {eval_results.get(\"eval_loss\", float(\"nan\")):.4f}')\n\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.half(),\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\ncommit_msg = (\n f'[{VERSION_TAG}] {LANG_NAME} fine-tuned checkpoint — '\n f'{train_result.global_step} steps | '\n f'WER {wer_score:.1%} | ' if wer_score == wer_score else f'WER n/a | '\n f'{len(correction_records)} corrections + WaxalNLP'\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_wer_disp = f'{wer_score:.1%}' if wer_score == wer_score else 'n/a'\nprint(f' Eval WER : {_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')" ] } ] }