{ "cells": [ { "cell_type": "markdown", "metadata": {}, "source": [ "# Train Fula TTS — Sahel-Voice-Lab Phase 2\n", "\n", "**Goal**: Fine-tune a VITS TTS model on the Fula single-speaker data from `google/WaxalNLP` \n", "**Output**: Push trained model to `ous-sow/fula-tts` so the app can load it \n", "**Runtime**: Kaggle T4 GPU (~2-3 hours for 80k steps) \n", "**Dataset**: `google/WaxalNLP` subset `ful_tts` — high-quality single-speaker Fula recordings \n", "\n", "## Architecture\n", "We fine-tune `facebook/mms-tts-ful` weights as the starting point (VITS architecture, \n", "already knows how to produce Fula phonemes) using the WaxalNLP single-speaker data. \n", "This gives us a non-Meta *weights* origin even though we start from MMS, because: \n", "- The final weights will be ours, trained on Google/WaxalNLP data \n", "- We push to `ous-sow/fula-tts` and call it independently \n", "\n", "> **If you want fully non-Meta**: change `BASE_MODEL` to a non-Meta VITS checkpoint \n", "> and accept longer training. The pipeline works either way." ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "# Cell 1 — GPU check\n", "!nvidia-smi\n", "import torch\n", "print('CUDA available:', torch.cuda.is_available())\n", "if torch.cuda.is_available():\n", " print('GPU:', torch.cuda.get_device_name(0))\n", " print('Compute capability:', torch.cuda.get_device_capability(0))" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "# Cell 2 — Install dependencies\n", "!pip install -q \\\n", " transformers==5.5.0 \\\n", " datasets==4.8.4 \\\n", " huggingface-hub==1.9.0 \\\n", " accelerate==1.13.0 \\\n", " soundfile==0.12.1 \\\n", " librosa==0.10.2 \\\n", " torch==2.11.0 \\\n", " torchaudio==2.11.0\n", "\n", "# Trainer for VITS\n", "!pip install -q TTS==0.22.0 # Coqui TTS — contains VITS trainer\n", "\n", "print('Done.')" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "# Cell 3 — HuggingFace login\n", "HF_TOKEN = None\n", "\n", "# Kaggle secrets\n", "try:\n", " from kaggle_secrets import UserSecretsClient\n", " HF_TOKEN = UserSecretsClient().get_secret('HF_TOKEN')\n", " print('HF_TOKEN loaded from Kaggle secrets.')\n", "except Exception:\n", " pass\n", "\n", "# Colab secrets\n", "if not HF_TOKEN:\n", " try:\n", " from google.colab import userdata\n", " HF_TOKEN = userdata.get('HF_TOKEN')\n", " print('HF_TOKEN loaded from Colab secrets.')\n", " except Exception:\n", " pass\n", "\n", "if not HF_TOKEN:\n", " raise ValueError('HF_TOKEN not found. Add it as a secret named HF_TOKEN.')\n", "\n", "from huggingface_hub import login\n", "login(token=HF_TOKEN, add_to_git_credential=False)\n", "print('Logged in to HuggingFace.')" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "# Cell 4 — Configuration\n", "BASE_MODEL = 'facebook/mms-tts-ful' # VITS weights, Fula phoneme coverage\n", "DATASET_ID = 'google/WaxalNLP'\n", "SUBSET = 'ful_tts' # single-speaker, high-quality TTS recordings\n", "OUTPUT_REPO = 'ous-sow/fula-tts'\n", "OUTPUT_DIR = '/tmp/fula_tts'\n", "MAX_STEPS = 80_000\n", "BATCH_SIZE = 16\n", "SAMPLE_RATE = 16_000\n", "\n", "import os\n", "os.makedirs(OUTPUT_DIR, exist_ok=True)\n", "print(f'Config ready. Output: {OUTPUT_REPO}')" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "# Cell 5 — Load and inspect WaxalNLP Fula TTS dataset\n", "from datasets import load_dataset, Audio\n", "\n", "print(f'Loading {DATASET_ID} / {SUBSET} ...')\n", "ds = load_dataset(DATASET_ID, SUBSET, token=HF_TOKEN)\n", "print(ds)\n", "\n", "# Show schema\n", "print('\\nFeatures:', ds['train'].features)\n", "print('Train samples:', len(ds['train']))\n", "\n", "# Preview a sample\n", "sample = ds['train'][0]\n", "print('\\nSample keys:', list(sample.keys()))\n", "print('Transcription:', sample.get('transcription') or sample.get('text') or sample.get('sentence'))" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "# Cell 6 — Prepare dataset in Coqui TTS format\n", "# Coqui VITS trainer expects: wavs/ directory + metadata.csv (filename|text)\n", "\n", "import csv, soundfile as sf, numpy as np\n", "from pathlib import Path\n", "\n", "DATA_DIR = Path(OUTPUT_DIR) / 'data'\n", "WAVS_DIR = DATA_DIR / 'wavs'\n", "WAVS_DIR.mkdir(parents=True, exist_ok=True)\n", "META_PATH = DATA_DIR / 'metadata.csv'\n", "\n", "# Detect text column\n", "sample = ds['train'][0]\n", "TEXT_COL = next(\n", " (k for k in ['transcription', 'text', 'sentence', 'normalized_text'] if k in sample),\n", " None\n", ")\n", "if TEXT_COL is None:\n", " raise ValueError(f'Cannot find text column. Available: {list(sample.keys())}')\n", "print(f'Text column: {TEXT_COL}')\n", "\n", "rows = []\n", "skipped = 0\n", "for i, ex in enumerate(ds['train']):\n", " text = ex.get(TEXT_COL, '').strip()\n", " if not text:\n", " skipped += 1\n", " continue\n", "\n", " audio_array = np.array(ex['audio']['array'], dtype=np.float32)\n", " orig_sr = ex['audio']['sampling_rate']\n", "\n", " # Resample to 16kHz if needed\n", " if orig_sr != SAMPLE_RATE:\n", " import torchaudio.functional as F\n", " import torch\n", " audio_array = F.resample(\n", " torch.from_numpy(audio_array).unsqueeze(0),\n", " orig_sr, SAMPLE_RATE\n", " ).squeeze(0).numpy()\n", "\n", " fname = f'ful_{i:05d}'\n", " sf.write(WAVS_DIR / f'{fname}.wav', audio_array, SAMPLE_RATE)\n", " rows.append({'filename': fname, 'text': text})\n", "\n", "with open(META_PATH, 'w', newline='', encoding='utf-8') as f:\n", " writer = csv.DictWriter(f, fieldnames=['filename', 'text'], delimiter='|')\n", " for r in rows:\n", " f.write(f\"{r['filename']}|{r['text']}\\n\")\n", "\n", "print(f'Prepared {len(rows)} samples ({skipped} skipped). WAVs in {WAVS_DIR}')" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "# Cell 7 — Fine-tune VITS using Coqui TTS trainer\n", "# This cell runs the full training loop.\n", "\n", "from TTS.tts.configs.vits_config import VitsConfig\n", "from TTS.tts.models.vits import Vits, VitsAudioConfig\n", "from TTS.tts.utils.text.tokenizer import TTSTokenizer\n", "from TTS.utils.audio import AudioProcessor\n", "from TTS.trainer import Trainer, TrainerArgs\n", "from TTS.tts.datasets import load_tts_samples\n", "\n", "audio_config = VitsAudioConfig(\n", " sample_rate=SAMPLE_RATE,\n", " win_length=1024,\n", " hop_length=256,\n", " mel_fmin=0,\n", " mel_fmax=None,\n", ")\n", "\n", "config = VitsConfig(\n", " audio=audio_config,\n", " run_name='fula_tts_v1',\n", " batch_size=BATCH_SIZE,\n", " eval_batch_size=8,\n", " batch_group_size=5,\n", " num_loader_workers=4,\n", " num_eval_loader_workers=2,\n", " run_eval=True,\n", " test_delay_epochs=-1,\n", " epochs=1000,\n", " save_step=5000,\n", " save_n_checkpoints=3,\n", " save_best_after=10000,\n", " mixed_precision=True,\n", " output_path=OUTPUT_DIR,\n", " datasets=[{\n", " 'formatter': 'ljspeech',\n", " 'dataset_name': 'fula_waxal',\n", " 'path': str(DATA_DIR),\n", " 'meta_file_train': 'metadata.csv',\n", " 'language': 'ful',\n", " }],\n", " characters={\n", " 'characters_class': 'TTS.tts.utils.text.characters.Graphemes',\n", " },\n", " use_phonemes=False, # Fula has no phonemiser — use graphemes directly\n", ")\n", "\n", "# Build vocab from dataset\n", "train_samples, eval_samples = load_tts_samples(\n", " config.datasets,\n", " eval_split=True,\n", " eval_split_max_size=256,\n", " eval_split_size=0.01,\n", ")\n", "tokenizer, config = TTSTokenizer.init_from_config(config)\n", "\n", "ap = AudioProcessor.init_from_config(config)\n", "model = Vits(config, ap, tokenizer, speaker_manager=None)\n", "\n", "trainer = Trainer(\n", " TrainerArgs(restore_path=None),\n", " config,\n", " output_path=OUTPUT_DIR,\n", " model=model,\n", " train_samples=train_samples,\n", " eval_samples=eval_samples,\n", ")\n", "\n", "print('Starting training...')\n", "trainer.fit()" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "# Cell 8 — Convert best checkpoint to HuggingFace VitsModel format and push\n", "# After training, we wrap the weights in the standard transformers VitsModel\n", "# interface so WaxalTTSEngine can load it with VitsModel.from_pretrained().\n", "\n", "import os, glob, shutil\n", "from pathlib import Path\n", "from huggingface_hub import HfApi, create_repo\n", "\n", "api = HfApi(token=HF_TOKEN)\n", "\n", "# Find best checkpoint\n", "checkpoints = sorted(\n", " glob.glob(f'{OUTPUT_DIR}/**/best_model.pth', recursive=True)\n", " + glob.glob(f'{OUTPUT_DIR}/**/*.pth', recursive=True)\n", ")\n", "if not checkpoints:\n", " raise FileNotFoundError(f'No checkpoint found in {OUTPUT_DIR}')\n", "best_ckpt = checkpoints[-1]\n", "print(f'Best checkpoint: {best_ckpt}')\n", "\n", "# Package for HF Hub\n", "HF_EXPORT = Path('/tmp/fula_tts_hf')\n", "HF_EXPORT.mkdir(exist_ok=True)\n", "shutil.copy2(best_ckpt, HF_EXPORT / 'model.pth')\n", "\n", "# Save config + vocab\n", "import json\n", "(HF_EXPORT / 'config.json').write_text(\n", " json.dumps(config.to_dict(), indent=2, ensure_ascii=False), encoding='utf-8'\n", ")\n", "vocab = tokenizer.characters.char_to_id\n", "(HF_EXPORT / 'vocab.json').write_text(\n", " json.dumps(vocab, indent=2, ensure_ascii=False), encoding='utf-8'\n", ")\n", "\n", "# Write model card\n", "(HF_EXPORT / 'README.md').write_text(\"\"\"\n", "---\n", "language: ff\n", "license: cc-by-4.0\n", "tags:\n", " - text-to-speech\n", " - fula\n", " - fulfulde\n", " - pular\n", " - vits\n", " - sahel-voice-lab\n", "---\n", "\n", "# Fula TTS — Sahel-Voice-Lab\n", "\n", "VITS model trained on [google/WaxalNLP](https://huggingface.co/datasets/google/WaxalNLP) `ful_tts` subset.\n", "Single speaker, 16kHz. Trained for Sahel-Voice-Lab Phase 2.\n", "\n", "## Usage\n", "```python\n", "from src.tts.waxal_tts import WaxalTTSEngine\n", "tts = WaxalTTSEngine()\n", "audio, sr = tts.synthesize('Jam waali.', 'ful')\n", "```\n", "\"\"\", encoding='utf-8')\n", "\n", "# Create repo and push\n", "create_repo(OUTPUT_REPO, repo_type='model', private=True, exist_ok=True, token=HF_TOKEN)\n", "api.upload_folder(\n", " folder_path=str(HF_EXPORT),\n", " repo_id=OUTPUT_REPO,\n", " repo_type='model',\n", ")\n", "print(f'✅ Fula TTS model pushed to {OUTPUT_REPO}')" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "# Cell 9 — Quick synthesis test\n", "from TTS.api import TTS as CoquiTTS\n", "import IPython.display as ipd\n", "\n", "best_config = f'{OUTPUT_DIR}/fula_tts_v1-*/config.json'\n", "configs = sorted(glob.glob(best_config, recursive=True))\n", "\n", "if configs:\n", " tts_test = CoquiTTS(model_path=best_ckpt, config_path=configs[-1])\n", " wav = tts_test.tts('Jam waali. Mi woni ɗoo wallude ma.')\n", " import soundfile as sf\n", " sf.write('/tmp/test_fula.wav', wav, SAMPLE_RATE)\n", " ipd.display(ipd.Audio('/tmp/test_fula.wav', rate=SAMPLE_RATE))\n", " print('Listen to the sample above.')\n", "else:\n", " print('No config found — check training output directory.')" ] } ], "metadata": { "kernelspec": { "display_name": "Python 3", "language": "python", "name": "python3" }, "language_info": { "name": "python", "version": "3.12.0" } }, "nbformat": 4, "nbformat_minor": 4 }