jefffffff9 Claude Sonnet 4.6 commited on
Commit
39604b3
·
1 Parent(s): 427d4a2

Fix Cell 6: replace get_last_checkpoint import with inline scanner

Browse files

transformers.trainer_utils.get_last_checkpoint triggers a deep import
chain (peft → generation → masking_utils → torch._dynamo) that breaks
on Kaggle Python 3.12 before packages are settled. The replacement
does the same thing (scan for checkpoint-N dirs, return the highest)
with zero extra imports.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

Files changed (1) hide show
  1. notebooks/kaggle_master_trainer.ipynb +11 -31
notebooks/kaggle_master_trainer.ipynb CHANGED
@@ -2,18 +2,23 @@
2
  "nbformat": 4,
3
  "nbformat_minor": 5,
4
  "metadata": {
5
- "kernelspec": {"display_name": "Python 3", "language": "python", "name": "python3"},
6
- "language_info": {"name": "python", "version": "3.10.12"}
 
 
 
 
 
 
 
7
  },
8
  "cells": [
9
-
10
  {
11
  "cell_type": "markdown",
12
  "id": "title",
13
  "metadata": {},
14
  "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"
15
  },
16
-
17
  {
18
  "cell_type": "code",
19
  "execution_count": null,
@@ -22,7 +27,6 @@
22
  "outputs": [],
23
  "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')"
24
  },
25
-
26
  {
27
  "cell_type": "code",
28
  "execution_count": null,
@@ -31,7 +35,6 @@
31
  "outputs": [],
32
  "source": "# ── Cell 2: Install dependencies ──────────────────────────────────────────────\n# Pinned to match the HF Space exactly so adapters are compatible.\n!pip install -q \\\n torch==2.11.0 torchaudio==2.11.0 \\\n transformers==5.5.0 \\\n datasets==4.8.4 \\\n accelerate==1.13.0 \\\n evaluate==0.4.2 \\\n huggingface-hub==1.9.0 \\\n peft==0.18.1 \\\n bitsandbytes==0.49.2 \\\n librosa==0.10.2 \\\n soundfile==0.12.1 \\\n jiwer==3.0.4 \\\n numpy==2.2.4 \\\n scipy==1.15.2\nprint('✅ Packages installed')"
33
  },
34
-
35
  {
36
  "cell_type": "code",
37
  "execution_count": null,
@@ -40,7 +43,6 @@
40
  "outputs": [],
41
  "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# ─── LoRA hyper-parameters ───────────────────────────────────────────────────\nLORA_R = 32 # rank — higher = more capacity, more VRAM\nLORA_ALPHA = 64 # scaling factor (typically 2× r)\nLORA_DROPOUT = 0.05\nLORA_TARGETS = ['q_proj', 'v_proj'] # Whisper decoder attention\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}')"
42
  },
43
-
44
  {
45
  "cell_type": "code",
46
  "execution_count": null,
@@ -49,7 +51,6 @@
49
  "outputs": [],
50
  "source": "# ── Cell 4: External dataset configuration ───────────────────────────────────\n# Add or remove entries here to include additional HF datasets.\n# Each entry: (repo_id, config_name, split, text_column, language_filter_or_None)\n#\n# Set ENABLED=True to activate a source, False to skip it.\n\nEXTERNAL_DATASETS = [\n {\n 'enabled' : False, # ← set True to include Common Voice\n 'repo_id' : 'mozilla-foundation/common_voice_13_0',\n 'config' : 'bm', # Bambara config code in Common Voice\n 'split' : 'train',\n 'text_col' : 'sentence',\n 'lang' : 'bam',\n 'max_samples': 2_000,\n },\n {\n 'enabled' : False,\n 'repo_id' : 'mozilla-foundation/common_voice_13_0',\n 'config' : 'ff', # Fula/Fulah config code\n 'split' : 'train',\n 'text_col' : 'sentence',\n 'lang' : 'ful',\n 'max_samples': 2_000,\n },\n # Add more datasets here:\n # {\n # 'enabled': True,\n # 'repo_id': 'your/dataset',\n # 'config': 'subset_name',\n # 'split': 'train',\n # 'text_col':'transcription',\n # 'lang': 'bam',\n # 'max_samples': 1000,\n # },\n]\n\nprint(f'External sources configured: {len(EXTERNAL_DATASETS)}')\nprint(f'Active for lang={TRAIN_LANG}: {sum(1 for d in EXTERNAL_DATASETS if d[\"enabled\"] and d[\"lang\"] == TRAIN_LANG)}')"
51
  },
52
-
53
  {
54
  "cell_type": "code",
55
  "execution_count": null,
@@ -58,16 +59,14 @@
58
  "outputs": [],
59
  "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}')"
60
  },
61
-
62
  {
63
  "cell_type": "code",
64
  "execution_count": null,
65
  "id": "cell-resume",
66
  "metadata": {},
67
  "outputs": [],
68
- "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\nfrom transformers.trainer_utils import get_last_checkpoint\n\nLAST_CHECKPOINT = None\nif Path(OUTPUT_DIR).exists():\n LAST_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.')"
69
  },
70
-
71
  {
72
  "cell_type": "code",
73
  "execution_count": null,
@@ -76,7 +75,6 @@
76
  "outputs": [],
77
  "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)}')"
78
  },
79
-
80
  {
81
  "cell_type": "code",
82
  "execution_count": null,
@@ -85,7 +83,6 @@
85
  "outputs": [],
86
  "source": "# ── Cell 8: Load WaxalNLP (with FLEURS fallback) ──────────────────────────────\n# WaxalNLP subsets for ASR: 'bam' (Bambara), 'ful' (Fula)\n# If the ASR subsets are missing, falls back to google/fleurs which has the\n# same language data under different subset codes.\n\nfrom datasets import load_dataset, Audio as HFAudio\n\nWAXAL_SUBSET_MAP = {'bam': 'bam', 'ful': 'ful'}\nFLEURS_SUBSET_MAP = {'bam': 'bam_ML', 'ful': 'ff_SN'}\n\nwaxal_ds = None\n\n# Try WaxalNLP first\ntry:\n subset = WAXAL_SUBSET_MAP[TRAIN_LANG]\n print(f'Loading google/WaxalNLP subset={subset} (streaming) ...')\n waxal_ds = load_dataset(\n 'google/WaxalNLP', subset,\n split='train', streaming=True,\n token=HF_TOKEN,\n )\n # Probe one item to verify it has audio + text\n probe = next(iter(waxal_ds))\n text_key = next(\n (k for k in ['transcription', 'text', 'sentence', 'normalized_text'] if k in probe),\n None\n )\n if text_key is None or 'audio' not in probe:\n raise ValueError(f'WaxalNLP/{subset} has unexpected schema: {list(probe.keys())}')\n print(f'WaxalNLP/{subset} ready — text column: \"{text_key}\"')\n WAXAL_TEXT_COL = text_key\nexcept Exception as e:\n print(f'WaxalNLP not available ({e}) — falling back to google/fleurs')\n subset = FLEURS_SUBSET_MAP[TRAIN_LANG]\n waxal_ds = load_dataset(\n 'google/fleurs', subset,\n split='train', streaming=True,\n token=HF_TOKEN,\n )\n probe = next(iter(waxal_ds))\n WAXAL_TEXT_COL = next(\n (k for k in ['transcription', 'text', 'sentence', 'raw_transcription'] if k in probe),\n list(probe.keys())[0]\n )\n print(f'FLEURS/{subset} ready — text column: \"{WAXAL_TEXT_COL}\"')\n\nprint(f'\\nWaxal/FLEURS source ready for {TRAIN_LANG}')"
87
  },
88
-
89
  {
90
  "cell_type": "code",
91
  "execution_count": null,
@@ -94,14 +91,12 @@
94
  "outputs": [],
95
  "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)}')"
96
  },
97
-
98
  {
99
  "cell_type": "markdown",
100
  "id": "md-pipeline",
101
  "metadata": {},
102
  "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."
103
  },
104
-
105
  {
106
  "cell_type": "code",
107
  "execution_count": null,
@@ -110,7 +105,6 @@
110
  "outputs": [],
111
  "source": "# ── Cell 10: Text cleaning utilities ─────────────────────────────────────────\nimport re, unicodedata\n\n# Extended-Latin characters that appear in Bambara and Fula writing\n_BAMBARA_EXTRA = set('ɛɔŋÉ')\n_FULA_EXTRA = set('ɓɗƴŋɲʼʻɗ')\n_BASE_LATIN = set('abcdefghijklmnopqrstuvwxyz')\n_ACCENTED = set('àâäèéêëîïôùûüýÿæœç')\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\ndef clean_text(text: str, lang: str = 'bam') -> str:\n \"\"\"\n Normalise and filter a transcription for Whisper training.\n Returns lowercase text containing only characters valid for the given language.\n \"\"\"\n if not text:\n return ''\n # NFKC: resolves ligatures and compatibility equivalents\n text = unicodedata.normalize('NFKC', text.lower().strip())\n # Remove URLs\n text = re.sub(r'https?://\\S+', '', text)\n # Remove XML / HTML tags\n text = re.sub(r'<[^>]+>', '', text)\n # Collapse repeated punctuation (.., ??, !!)\n text = re.sub(r'([.,!?])\\1+', r'\\1', text)\n # Keep only allowlisted characters\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 # Collapse whitespace\n return re.sub(r'\\s+', ' ', text).strip()\n\n\n# Smoke test\nassert clean_text('I ni ce! (hello)', 'bam') == 'i ni ce!'\nassert clean_text('Jam waali. <b>test</b>', 'ful') == 'jam waali.'\nprint('✅ clean_text tests passed')"
112
  },
113
-
114
  {
115
  "cell_type": "code",
116
  "execution_count": null,
@@ -119,7 +113,6 @@
119
  "outputs": [],
120
  "source": "# ── Cell 11: Whisper processor + prepare_dataset ─────────────────────────────\nimport numpy as np\nfrom transformers import WhisperProcessor\n\nprint(f'Loading processor: {WHISPER_MODEL_ID} ...')\nprocessor = WhisperProcessor.from_pretrained(WHISPER_MODEL_ID, token=HF_TOKEN)\nprint('✅ Processor ready')\n\n\ndef prepare_dataset(batch, text_col='transcription', lang=TRAIN_LANG):\n \"\"\"\n Unified feature-extraction for all data sources.\n Resamples audio to 16 kHz, cleans text, returns Whisper input_features + labels.\n \"\"\"\n audio = batch['audio']\n audio_array = np.array(audio['array'], dtype=np.float32)\n orig_sr = audio['sampling_rate']\n\n # Resample if needed\n if orig_sr != TARGET_SR:\n try:\n import torchaudio.functional as F_audio\n import 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 # Whisper log-mel features\n batch['input_features'] = processor.feature_extractor(\n audio_array, sampling_rate=TARGET_SR\n ).input_features[0]\n\n # Text → token IDs\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\n return batch\n\n\nprint('✅ prepare_dataset ready')"
121
  },
122
-
123
  {
124
  "cell_type": "code",
125
  "execution_count": null,
@@ -128,14 +121,12 @@
128
  "outputs": [],
129
  "source": "# ── Cell 12: Build & merge all datasets ──────────────────────────────────────\n# Order of priority: user corrections (upsampled) → WaxalNLP → external\nfrom datasets import Dataset, Audio as HFAudio, concatenate_datasets\nfrom functools import partial\n\nall_parts = []\n\n# ── Part A: User corrections from sahel-agri-feedback ─────────────────────\nif correction_records:\n print(f'Part A: {len(correction_records)} corrections × {CORRECTION_REPEAT} upsample')\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': [\n r.get('corrected_text') or r.get('transcription', '') for r in rec_list\n ],\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 ready')\nelse:\n print('Part A: no corrections available — skipping')\n\n# ── Part B: WaxalNLP / FLEURS ──────────────────────────────────────────────\nif waxal_ds is not None:\n print(f'Part B: materialising up to {MAX_WAXAL_TRAIN} WaxalNLP samples ...')\n waxal_capped = waxal_ds.take(MAX_WAXAL_TRAIN)\n waxal_rows = list(waxal_capped)\n print(f' Fetched {len(waxal_rows)} rows')\n\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 num_proc=2,\n )\n all_parts.append(waxal_local)\n print(f' → {len(waxal_local)} samples ready')\n\n# ── Part C: External datasets ─────────────────────────────────────────────\nfor ext_ds, text_col in external_datasets:\n ext_rows = list(ext_ds) # already capped by .take() in Cell 9\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 num_proc=2,\n )\n all_parts.append(ext_local)\n print(f'Part C: {len(ext_local)} external samples ready')\n\n# ── Merge + train/eval split ──────────────────────────────────────────────\nif not all_parts:\n raise RuntimeError('No data loaded from any source — check config and HF_TOKEN')\n\ncombined = concatenate_datasets(all_parts).shuffle(seed=42)\nsplit = combined.train_test_split(test_size=min(0.05, 200 / len(combined)))\ntrain_ds = split['train']\neval_ds = split['test']\n\nprint(f'\\n✅ Dataset ready')\nprint(f' Train : {len(train_ds)}')\nprint(f' Eval : {len(eval_ds)}')"
130
  },
131
-
132
  {
133
  "cell_type": "markdown",
134
  "id": "md-model",
135
  "metadata": {},
136
  "source": "---\n## Model Setup — 8-bit Deep Sleep\n\n`openai/whisper-small` is loaded in **8-bit** via `BitsAndBytesConfig` to fit in T4's 15 GB VRAM.\n`prepare_model_for_kbit_training` enables gradient checkpointing and casts norms to float32.\nLoRA is injected only into the **decoder** attention projections — the encoder stays frozen,\nso gradient memory stays minimal while the language-specific output improves."
137
  },
138
-
139
  {
140
  "cell_type": "code",
141
  "execution_count": null,
@@ -144,7 +135,6 @@
144
  "outputs": [],
145
  "source": "# ── Cell 13: Load Whisper-small (8-bit) + LoRA ────────────────────────────────\nimport torch\nfrom transformers import WhisperForConditionalGeneration, BitsAndBytesConfig\nfrom peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training, TaskType\n\n# 8-bit quantisation — halves VRAM vs fp16 with negligible quality loss\nbnb_config = BitsAndBytesConfig(load_in_8bit=True)\n\nprint(f'Loading {WHISPER_MODEL_ID} in 8-bit ...')\nmodel = WhisperForConditionalGeneration.from_pretrained(\n WHISPER_MODEL_ID,\n quantization_config=bnb_config,\n device_map='auto',\n token=HF_TOKEN,\n)\n\n# Required before applying LoRA to a k-bit model:\n# - enables gradient checkpointing\n# - casts LayerNorm & output embedding to float32 for stability\nmodel = prepare_model_for_kbit_training(model)\n\n# Force Whisper to generate in the target language (no language detection overhead)\nmodel.config.forced_decoder_ids = processor.get_decoder_prompt_ids(\n language='fr' if TRAIN_LANG in ('bam', 'ful') else TRAIN_LANG,\n task='transcribe',\n)\nmodel.config.suppress_tokens = []\n\n# LoRA: inject adapters into decoder Q and V projections\nlora_cfg = LoraConfig(\n r=LORA_R,\n lora_alpha=LORA_ALPHA,\n target_modules=LORA_TARGETS,\n lora_dropout=LORA_DROPOUT,\n bias='none',\n task_type=TaskType.SEQ_2_SEQ_LM,\n)\nmodel = get_peft_model(model, lora_cfg)\nmodel.print_trainable_parameters()\nprint(f'\\n✅ Model ready on {next(model.parameters()).device}')"
146
  },
147
-
148
  {
149
  "cell_type": "code",
150
  "execution_count": null,
@@ -153,14 +143,12 @@
153
  "outputs": [],
154
  "source": "# ── Cell 14: Data collator + WER metric ──────────────────────────────────────\nimport evaluate\nfrom dataclasses import dataclass\nfrom typing import Any, Dict, List\n\nwer_metric = evaluate.load('wer')\n\n\n@dataclass\nclass DataCollatorSpeechSeq2SeqWithPadding:\n \"\"\"\n Pads input_features (mel spectrograms) and labels independently.\n Labels use -100 for padding so CrossEntropyLoss ignores them.\n Strips the BOS token from label sequences (Whisper convention).\n \"\"\"\n processor: Any\n\n def __call__(self, features: List[Dict]) -> Dict:\n # Pad mel features\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 # Pad labels\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 # Remove leading BOS token (Whisper adds it during generation, not training)\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 \"\"\"WER on the eval set (lower is better).\"\"\"\n pred_ids = pred.predictions\n label_ids = pred.label_ids\n # Replace -100 padding back to pad_token_id for decoding\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 wer = wer_metric.compute(predictions=pred_str, references=label_str)\n return {'wer': round(wer, 4)}\n\n\ncollator = DataCollatorSpeechSeq2SeqWithPadding(processor=processor)\nprint('✅ Collator and WER metric ready')"
155
  },
156
-
157
  {
158
  "cell_type": "markdown",
159
  "id": "md-train",
160
  "metadata": {},
161
  "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."
162
  },
163
-
164
  {
165
  "cell_type": "code",
166
  "execution_count": null,
@@ -169,7 +157,6 @@
169
  "outputs": [],
170
  "source": "# ── Cell 15: Training arguments ───────────────────────────────────────────────\nfrom transformers import Seq2SeqTrainingArguments\n\ntraining_args = Seq2SeqTrainingArguments(\n output_dir=OUTPUT_DIR,\n\n # Steps\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 # Batch\n per_device_train_batch_size=BATCH_SIZE,\n per_device_eval_batch_size=8,\n gradient_accumulation_steps=GRAD_ACCUM,\n\n # Precision & memory\n fp16=True,\n gradient_checkpointing=True,\n\n # Optimiser\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 # Evaluation\n evaluation_strategy='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 # Checkpointing\n save_total_limit=3,\n save_strategy='steps',\n\n # Logging\n report_to=['tensorboard'],\n logging_dir=f'{OUTPUT_DIR}/logs',\n\n # Push — disabled during training, done manually in the export cell\n push_to_hub=False,\n)\nprint('✅ Training arguments configured')\nprint(f' Effective batch size: {BATCH_SIZE * GRAD_ACCUM}')\nprint(f' Max steps : {MAX_STEPS}')\nprint(f' Checkpoint every : {SAVE_STEPS} steps')"
171
  },
172
-
173
  {
174
  "cell_type": "code",
175
  "execution_count": null,
@@ -178,14 +165,12 @@
178
  "outputs": [],
179
  "source": "# ── Cell 16: TRAIN — Deep Sleep ───────────────────────────────────────────────\n# This is the long-running cell. Expected wall time on T4:\n# 2000 steps ≈ 25 min | 4000 steps ≈ 50 min | 8000 steps ≈ 100 min\n#\n# If the session times out, just re-run the notebook from Cell 1.\n# LAST_CHECKPOINT (Cell 6) will be detected and training resumes automatically.\n\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=processor.feature_extractor,\n)\n\nprint(f'Starting Deep Sleep training ...')\nprint(f' Resume from : {LAST_CHECKPOINT or \"scratch\"}')\nprint(f' Train size : {len(train_ds)}')\nprint(f' Eval size : {len(eval_ds)}')\nprint()\n\ntrain_result = trainer.train(resume_from_checkpoint=LAST_CHECKPOINT)\n\nprint('\\n✅ Training complete')\nprint(f' Steps : {train_result.global_step}')\nprint(f' Train loss: {train_result.training_loss:.4f}')\n\n# Save the best model weights locally\ntrainer.save_model(OUTPUT_DIR)\nprocessor.save_pretrained(OUTPUT_DIR)\nprint(f' Adapter saved → {OUTPUT_DIR}')"
180
  },
181
-
182
  {
183
  "cell_type": "markdown",
184
  "id": "md-eval",
185
  "metadata": {},
186
  "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."
187
  },
188
-
189
  {
190
  "cell_type": "code",
191
  "execution_count": null,
@@ -194,14 +179,12 @@
194
  "outputs": [],
195
  "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 forced_decoder_ids=model.config.forced_decoder_ids,\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()"
196
  },
197
-
198
  {
199
  "cell_type": "markdown",
200
  "id": "md-export",
201
  "metadata": {},
202
  "source": "---\n## Export — Push LoRA Adapter 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."
203
  },
204
-
205
  {
206
  "cell_type": "code",
207
  "execution_count": null,
@@ -210,7 +193,6 @@
210
  "outputs": [],
211
  "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}')"
212
  },
213
-
214
  {
215
  "cell_type": "code",
216
  "execution_count": null,
@@ -219,7 +201,6 @@
219
  "outputs": [],
220
  "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} LoRA adapter — '\n f'{train_result.global_step} steps | '\n f'WER {wer_score:.1%} | '\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}')"
221
  },
222
-
223
  {
224
  "cell_type": "code",
225
  "execution_count": null,
@@ -228,6 +209,5 @@
228
  "outputs": [],
229
  "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}')\nprint(f' Eval WER : {wer_score:.1%}')\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')"
230
  }
231
-
232
  ]
233
- }
 
2
  "nbformat": 4,
3
  "nbformat_minor": 5,
4
  "metadata": {
5
+ "kernelspec": {
6
+ "display_name": "Python 3",
7
+ "language": "python",
8
+ "name": "python3"
9
+ },
10
+ "language_info": {
11
+ "name": "python",
12
+ "version": "3.10.12"
13
+ }
14
  },
15
  "cells": [
 
16
  {
17
  "cell_type": "markdown",
18
  "id": "title",
19
  "metadata": {},
20
  "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"
21
  },
 
22
  {
23
  "cell_type": "code",
24
  "execution_count": null,
 
27
  "outputs": [],
28
  "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')"
29
  },
 
30
  {
31
  "cell_type": "code",
32
  "execution_count": null,
 
35
  "outputs": [],
36
  "source": "# ── Cell 2: Install dependencies ──────────────────────────────────────────────\n# Pinned to match the HF Space exactly so adapters are compatible.\n!pip install -q \\\n torch==2.11.0 torchaudio==2.11.0 \\\n transformers==5.5.0 \\\n datasets==4.8.4 \\\n accelerate==1.13.0 \\\n evaluate==0.4.2 \\\n huggingface-hub==1.9.0 \\\n peft==0.18.1 \\\n bitsandbytes==0.49.2 \\\n librosa==0.10.2 \\\n soundfile==0.12.1 \\\n jiwer==3.0.4 \\\n numpy==2.2.4 \\\n scipy==1.15.2\nprint('✅ Packages installed')"
37
  },
 
38
  {
39
  "cell_type": "code",
40
  "execution_count": null,
 
43
  "outputs": [],
44
  "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# ─── LoRA hyper-parameters ───────────────────────────────────────────────────\nLORA_R = 32 # rank — higher = more capacity, more VRAM\nLORA_ALPHA = 64 # scaling factor (typically 2× r)\nLORA_DROPOUT = 0.05\nLORA_TARGETS = ['q_proj', 'v_proj'] # Whisper decoder attention\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}')"
45
  },
 
46
  {
47
  "cell_type": "code",
48
  "execution_count": null,
 
51
  "outputs": [],
52
  "source": "# ── Cell 4: External dataset configuration ───────────────────────────────────\n# Add or remove entries here to include additional HF datasets.\n# Each entry: (repo_id, config_name, split, text_column, language_filter_or_None)\n#\n# Set ENABLED=True to activate a source, False to skip it.\n\nEXTERNAL_DATASETS = [\n {\n 'enabled' : False, # ← set True to include Common Voice\n 'repo_id' : 'mozilla-foundation/common_voice_13_0',\n 'config' : 'bm', # Bambara config code in Common Voice\n 'split' : 'train',\n 'text_col' : 'sentence',\n 'lang' : 'bam',\n 'max_samples': 2_000,\n },\n {\n 'enabled' : False,\n 'repo_id' : 'mozilla-foundation/common_voice_13_0',\n 'config' : 'ff', # Fula/Fulah config code\n 'split' : 'train',\n 'text_col' : 'sentence',\n 'lang' : 'ful',\n 'max_samples': 2_000,\n },\n # Add more datasets here:\n # {\n # 'enabled': True,\n # 'repo_id': 'your/dataset',\n # 'config': 'subset_name',\n # 'split': 'train',\n # 'text_col':'transcription',\n # 'lang': 'bam',\n # 'max_samples': 1000,\n # },\n]\n\nprint(f'External sources configured: {len(EXTERNAL_DATASETS)}')\nprint(f'Active for lang={TRAIN_LANG}: {sum(1 for d in EXTERNAL_DATASETS if d[\"enabled\"] and d[\"lang\"] == TRAIN_LANG)}')"
53
  },
 
54
  {
55
  "cell_type": "code",
56
  "execution_count": null,
 
59
  "outputs": [],
60
  "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}')"
61
  },
 
62
  {
63
  "cell_type": "code",
64
  "execution_count": null,
65
  "id": "cell-resume",
66
  "metadata": {},
67
  "outputs": [],
68
+ "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.')"
69
  },
 
70
  {
71
  "cell_type": "code",
72
  "execution_count": null,
 
75
  "outputs": [],
76
  "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)}')"
77
  },
 
78
  {
79
  "cell_type": "code",
80
  "execution_count": null,
 
83
  "outputs": [],
84
  "source": "# ── Cell 8: Load WaxalNLP (with FLEURS fallback) ──────────────────────────────\n# WaxalNLP subsets for ASR: 'bam' (Bambara), 'ful' (Fula)\n# If the ASR subsets are missing, falls back to google/fleurs which has the\n# same language data under different subset codes.\n\nfrom datasets import load_dataset, Audio as HFAudio\n\nWAXAL_SUBSET_MAP = {'bam': 'bam', 'ful': 'ful'}\nFLEURS_SUBSET_MAP = {'bam': 'bam_ML', 'ful': 'ff_SN'}\n\nwaxal_ds = None\n\n# Try WaxalNLP first\ntry:\n subset = WAXAL_SUBSET_MAP[TRAIN_LANG]\n print(f'Loading google/WaxalNLP subset={subset} (streaming) ...')\n waxal_ds = load_dataset(\n 'google/WaxalNLP', subset,\n split='train', streaming=True,\n token=HF_TOKEN,\n )\n # Probe one item to verify it has audio + text\n probe = next(iter(waxal_ds))\n text_key = next(\n (k for k in ['transcription', 'text', 'sentence', 'normalized_text'] if k in probe),\n None\n )\n if text_key is None or 'audio' not in probe:\n raise ValueError(f'WaxalNLP/{subset} has unexpected schema: {list(probe.keys())}')\n print(f'WaxalNLP/{subset} ready — text column: \"{text_key}\"')\n WAXAL_TEXT_COL = text_key\nexcept Exception as e:\n print(f'WaxalNLP not available ({e}) — falling back to google/fleurs')\n subset = FLEURS_SUBSET_MAP[TRAIN_LANG]\n waxal_ds = load_dataset(\n 'google/fleurs', subset,\n split='train', streaming=True,\n token=HF_TOKEN,\n )\n probe = next(iter(waxal_ds))\n WAXAL_TEXT_COL = next(\n (k for k in ['transcription', 'text', 'sentence', 'raw_transcription'] if k in probe),\n list(probe.keys())[0]\n )\n print(f'FLEURS/{subset} ready — text column: \"{WAXAL_TEXT_COL}\"')\n\nprint(f'\\nWaxal/FLEURS source ready for {TRAIN_LANG}')"
85
  },
 
86
  {
87
  "cell_type": "code",
88
  "execution_count": null,
 
91
  "outputs": [],
92
  "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)}')"
93
  },
 
94
  {
95
  "cell_type": "markdown",
96
  "id": "md-pipeline",
97
  "metadata": {},
98
  "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."
99
  },
 
100
  {
101
  "cell_type": "code",
102
  "execution_count": null,
 
105
  "outputs": [],
106
  "source": "# ── Cell 10: Text cleaning utilities ─────────────────────────────────────────\nimport re, unicodedata\n\n# Extended-Latin characters that appear in Bambara and Fula writing\n_BAMBARA_EXTRA = set('ɛɔŋÉ')\n_FULA_EXTRA = set('ɓɗƴŋɲʼʻɗ')\n_BASE_LATIN = set('abcdefghijklmnopqrstuvwxyz')\n_ACCENTED = set('àâäèéêëîïôùûüýÿæœç')\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\ndef clean_text(text: str, lang: str = 'bam') -> str:\n \"\"\"\n Normalise and filter a transcription for Whisper training.\n Returns lowercase text containing only characters valid for the given language.\n \"\"\"\n if not text:\n return ''\n # NFKC: resolves ligatures and compatibility equivalents\n text = unicodedata.normalize('NFKC', text.lower().strip())\n # Remove URLs\n text = re.sub(r'https?://\\S+', '', text)\n # Remove XML / HTML tags\n text = re.sub(r'<[^>]+>', '', text)\n # Collapse repeated punctuation (.., ??, !!)\n text = re.sub(r'([.,!?])\\1+', r'\\1', text)\n # Keep only allowlisted characters\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 # Collapse whitespace\n return re.sub(r'\\s+', ' ', text).strip()\n\n\n# Smoke test\nassert clean_text('I ni ce! (hello)', 'bam') == 'i ni ce!'\nassert clean_text('Jam waali. <b>test</b>', 'ful') == 'jam waali.'\nprint('✅ clean_text tests passed')"
107
  },
 
108
  {
109
  "cell_type": "code",
110
  "execution_count": null,
 
113
  "outputs": [],
114
  "source": "# ── Cell 11: Whisper processor + prepare_dataset ─────────────────────────────\nimport numpy as np\nfrom transformers import WhisperProcessor\n\nprint(f'Loading processor: {WHISPER_MODEL_ID} ...')\nprocessor = WhisperProcessor.from_pretrained(WHISPER_MODEL_ID, token=HF_TOKEN)\nprint('✅ Processor ready')\n\n\ndef prepare_dataset(batch, text_col='transcription', lang=TRAIN_LANG):\n \"\"\"\n Unified feature-extraction for all data sources.\n Resamples audio to 16 kHz, cleans text, returns Whisper input_features + labels.\n \"\"\"\n audio = batch['audio']\n audio_array = np.array(audio['array'], dtype=np.float32)\n orig_sr = audio['sampling_rate']\n\n # Resample if needed\n if orig_sr != TARGET_SR:\n try:\n import torchaudio.functional as F_audio\n import 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 # Whisper log-mel features\n batch['input_features'] = processor.feature_extractor(\n audio_array, sampling_rate=TARGET_SR\n ).input_features[0]\n\n # Text → token IDs\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\n return batch\n\n\nprint('✅ prepare_dataset ready')"
115
  },
 
116
  {
117
  "cell_type": "code",
118
  "execution_count": null,
 
121
  "outputs": [],
122
  "source": "# ── Cell 12: Build & merge all datasets ──────────────────────────────────────\n# Order of priority: user corrections (upsampled) → WaxalNLP → external\nfrom datasets import Dataset, Audio as HFAudio, concatenate_datasets\nfrom functools import partial\n\nall_parts = []\n\n# ── Part A: User corrections from sahel-agri-feedback ─────────────────────\nif correction_records:\n print(f'Part A: {len(correction_records)} corrections × {CORRECTION_REPEAT} upsample')\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': [\n r.get('corrected_text') or r.get('transcription', '') for r in rec_list\n ],\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 ready')\nelse:\n print('Part A: no corrections available — skipping')\n\n# ── Part B: WaxalNLP / FLEURS ──────────────────────────────────────────────\nif waxal_ds is not None:\n print(f'Part B: materialising up to {MAX_WAXAL_TRAIN} WaxalNLP samples ...')\n waxal_capped = waxal_ds.take(MAX_WAXAL_TRAIN)\n waxal_rows = list(waxal_capped)\n print(f' Fetched {len(waxal_rows)} rows')\n\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 num_proc=2,\n )\n all_parts.append(waxal_local)\n print(f' → {len(waxal_local)} samples ready')\n\n# ── Part C: External datasets ─────────────────────────────────────────────\nfor ext_ds, text_col in external_datasets:\n ext_rows = list(ext_ds) # already capped by .take() in Cell 9\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 num_proc=2,\n )\n all_parts.append(ext_local)\n print(f'Part C: {len(ext_local)} external samples ready')\n\n# ── Merge + train/eval split ──────────────────────────────────────────────\nif not all_parts:\n raise RuntimeError('No data loaded from any source — check config and HF_TOKEN')\n\ncombined = concatenate_datasets(all_parts).shuffle(seed=42)\nsplit = combined.train_test_split(test_size=min(0.05, 200 / len(combined)))\ntrain_ds = split['train']\neval_ds = split['test']\n\nprint(f'\\n✅ Dataset ready')\nprint(f' Train : {len(train_ds)}')\nprint(f' Eval : {len(eval_ds)}')"
123
  },
 
124
  {
125
  "cell_type": "markdown",
126
  "id": "md-model",
127
  "metadata": {},
128
  "source": "---\n## Model Setup — 8-bit Deep Sleep\n\n`openai/whisper-small` is loaded in **8-bit** via `BitsAndBytesConfig` to fit in T4's 15 GB VRAM.\n`prepare_model_for_kbit_training` enables gradient checkpointing and casts norms to float32.\nLoRA is injected only into the **decoder** attention projections — the encoder stays frozen,\nso gradient memory stays minimal while the language-specific output improves."
129
  },
 
130
  {
131
  "cell_type": "code",
132
  "execution_count": null,
 
135
  "outputs": [],
136
  "source": "# ── Cell 13: Load Whisper-small (8-bit) + LoRA ────────────────────────────────\nimport torch\nfrom transformers import WhisperForConditionalGeneration, BitsAndBytesConfig\nfrom peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training, TaskType\n\n# 8-bit quantisation — halves VRAM vs fp16 with negligible quality loss\nbnb_config = BitsAndBytesConfig(load_in_8bit=True)\n\nprint(f'Loading {WHISPER_MODEL_ID} in 8-bit ...')\nmodel = WhisperForConditionalGeneration.from_pretrained(\n WHISPER_MODEL_ID,\n quantization_config=bnb_config,\n device_map='auto',\n token=HF_TOKEN,\n)\n\n# Required before applying LoRA to a k-bit model:\n# - enables gradient checkpointing\n# - casts LayerNorm & output embedding to float32 for stability\nmodel = prepare_model_for_kbit_training(model)\n\n# Force Whisper to generate in the target language (no language detection overhead)\nmodel.config.forced_decoder_ids = processor.get_decoder_prompt_ids(\n language='fr' if TRAIN_LANG in ('bam', 'ful') else TRAIN_LANG,\n task='transcribe',\n)\nmodel.config.suppress_tokens = []\n\n# LoRA: inject adapters into decoder Q and V projections\nlora_cfg = LoraConfig(\n r=LORA_R,\n lora_alpha=LORA_ALPHA,\n target_modules=LORA_TARGETS,\n lora_dropout=LORA_DROPOUT,\n bias='none',\n task_type=TaskType.SEQ_2_SEQ_LM,\n)\nmodel = get_peft_model(model, lora_cfg)\nmodel.print_trainable_parameters()\nprint(f'\\n✅ Model ready on {next(model.parameters()).device}')"
137
  },
 
138
  {
139
  "cell_type": "code",
140
  "execution_count": null,
 
143
  "outputs": [],
144
  "source": "# ── Cell 14: Data collator + WER metric ──────────────────────────────────────\nimport evaluate\nfrom dataclasses import dataclass\nfrom typing import Any, Dict, List\n\nwer_metric = evaluate.load('wer')\n\n\n@dataclass\nclass DataCollatorSpeechSeq2SeqWithPadding:\n \"\"\"\n Pads input_features (mel spectrograms) and labels independently.\n Labels use -100 for padding so CrossEntropyLoss ignores them.\n Strips the BOS token from label sequences (Whisper convention).\n \"\"\"\n processor: Any\n\n def __call__(self, features: List[Dict]) -> Dict:\n # Pad mel features\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 # Pad labels\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 # Remove leading BOS token (Whisper adds it during generation, not training)\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 \"\"\"WER on the eval set (lower is better).\"\"\"\n pred_ids = pred.predictions\n label_ids = pred.label_ids\n # Replace -100 padding back to pad_token_id for decoding\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 wer = wer_metric.compute(predictions=pred_str, references=label_str)\n return {'wer': round(wer, 4)}\n\n\ncollator = DataCollatorSpeechSeq2SeqWithPadding(processor=processor)\nprint('✅ Collator and WER metric ready')"
145
  },
 
146
  {
147
  "cell_type": "markdown",
148
  "id": "md-train",
149
  "metadata": {},
150
  "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."
151
  },
 
152
  {
153
  "cell_type": "code",
154
  "execution_count": null,
 
157
  "outputs": [],
158
  "source": "# ── Cell 15: Training arguments ───────────────────────────────────────────────\nfrom transformers import Seq2SeqTrainingArguments\n\ntraining_args = Seq2SeqTrainingArguments(\n output_dir=OUTPUT_DIR,\n\n # Steps\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 # Batch\n per_device_train_batch_size=BATCH_SIZE,\n per_device_eval_batch_size=8,\n gradient_accumulation_steps=GRAD_ACCUM,\n\n # Precision & memory\n fp16=True,\n gradient_checkpointing=True,\n\n # Optimiser\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 # Evaluation\n evaluation_strategy='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 # Checkpointing\n save_total_limit=3,\n save_strategy='steps',\n\n # Logging\n report_to=['tensorboard'],\n logging_dir=f'{OUTPUT_DIR}/logs',\n\n # Push — disabled during training, done manually in the export cell\n push_to_hub=False,\n)\nprint('✅ Training arguments configured')\nprint(f' Effective batch size: {BATCH_SIZE * GRAD_ACCUM}')\nprint(f' Max steps : {MAX_STEPS}')\nprint(f' Checkpoint every : {SAVE_STEPS} steps')"
159
  },
 
160
  {
161
  "cell_type": "code",
162
  "execution_count": null,
 
165
  "outputs": [],
166
  "source": "# ── Cell 16: TRAIN — Deep Sleep ───────────────────────────────────────────────\n# This is the long-running cell. Expected wall time on T4:\n# 2000 steps ≈ 25 min | 4000 steps ≈ 50 min | 8000 steps ≈ 100 min\n#\n# If the session times out, just re-run the notebook from Cell 1.\n# LAST_CHECKPOINT (Cell 6) will be detected and training resumes automatically.\n\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=processor.feature_extractor,\n)\n\nprint(f'Starting Deep Sleep training ...')\nprint(f' Resume from : {LAST_CHECKPOINT or \"scratch\"}')\nprint(f' Train size : {len(train_ds)}')\nprint(f' Eval size : {len(eval_ds)}')\nprint()\n\ntrain_result = trainer.train(resume_from_checkpoint=LAST_CHECKPOINT)\n\nprint('\\n✅ Training complete')\nprint(f' Steps : {train_result.global_step}')\nprint(f' Train loss: {train_result.training_loss:.4f}')\n\n# Save the best model weights locally\ntrainer.save_model(OUTPUT_DIR)\nprocessor.save_pretrained(OUTPUT_DIR)\nprint(f' Adapter saved → {OUTPUT_DIR}')"
167
  },
 
168
  {
169
  "cell_type": "markdown",
170
  "id": "md-eval",
171
  "metadata": {},
172
  "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."
173
  },
 
174
  {
175
  "cell_type": "code",
176
  "execution_count": null,
 
179
  "outputs": [],
180
  "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 forced_decoder_ids=model.config.forced_decoder_ids,\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()"
181
  },
 
182
  {
183
  "cell_type": "markdown",
184
  "id": "md-export",
185
  "metadata": {},
186
  "source": "---\n## Export — Push LoRA Adapter 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."
187
  },
 
188
  {
189
  "cell_type": "code",
190
  "execution_count": null,
 
193
  "outputs": [],
194
  "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}')"
195
  },
 
196
  {
197
  "cell_type": "code",
198
  "execution_count": null,
 
201
  "outputs": [],
202
  "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} LoRA adapter — '\n f'{train_result.global_step} steps | '\n f'WER {wer_score:.1%} | '\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}')"
203
  },
 
204
  {
205
  "cell_type": "code",
206
  "execution_count": null,
 
209
  "outputs": [],
210
  "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}')\nprint(f' Eval WER : {wer_score:.1%}')\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')"
211
  }
 
212
  ]
213
+ }