{ "cells": [ { "cell_type": "markdown", "metadata": {}, "source": [ "# Fine-tuning `granite-speech-5.0-470m-turboctc`\n", "\n", "Fine-tunes IBM's [470M CTC speech model](https://huggingface.co/ibm-granite/granite-speech-5.0-470m-turboctc)\n", "on [FLEURS German](https://huggingface.co/datasets/google/fleurs) (~10 h, ungated).\n", "\n", "The model is English-only, so German is a real domain shift. Verified end to end on one H100\n", "(~37 min, defaults below unchanged):\n", "\n", "| | baseline | after 8 epochs |\n", "|---|---|---|\n", "| WER | 77.5% | **51.9%** |\n", "| CER | 36.5% | **20.0%** |\n", "\n", "To use your own data, change one cell (\u00a72).\n", "\n", "This is a pure CTC encoder \u2014 no LLM decoder, no chat template, no LoRA. Audio in, text out, one\n", "non-autoregressive forward pass.\n", "\n", "**Runtime \u2192 Change runtime type \u2192 GPU** before running." ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 1. Setup\n", "\n", "Needs `transformers >= 5.16`, newer than Colab ships." ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "%%capture\n", "# transformers >= 5.16 is required and newer than Colab ships.\n", "!pip install -q -U git+https://github.com/huggingface/transformers.git\n", "!pip install -q -U \"datasets>=4.0\" accelerate evaluate jiwer" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "# `datasets` decodes audio only through torchcodec, which is ABI-linked to torch's\n", "# CUDA runtime but declares no torch dependency -- so a plain `pip install -U\n", "# torchcodec` can grab a build for a different CUDA and every decode then dies with\n", "# OSError: libnvrtc.so.NN: cannot open shared object file\n", "# Reinstalling the same version does not help; the version has to match torch.\n", "import subprocess, sys\n", "\n", "import torch\n", "\n", "TORCHCODEC_FOR_TORCH = { # torch minor -> torchcodec release built against it\n", " \"2.9\": \"0.9.1\", \"2.10\": \"0.10.0\", \"2.11\": \"0.11.1\",\n", " \"2.12\": \"0.12.0\", \"2.13\": \"0.15.0\", \"2.14\": \"0.16.0\",\n", "}\n", "torch_mm = \".\".join(torch.__version__.split(\".\")[:2])\n", "pin = TORCHCODEC_FOR_TORCH.get(torch_mm)\n", "print(f\"torch {torch.__version__} -> torchcodec {pin or 'unknown (leaving as-is)'}\")\n", "\n", "if pin:\n", " subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"-q\",\n", " \"--force-reinstall\", \"--no-deps\", f\"torchcodec=={pin}\"], check=True)\n", " print(\"installed; now Runtime > Restart session, then continue from the next cell.\")" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "*Restart the session once after installing, then run from here.* The next cell verifies the\n", "environment \u2014 if audio decoding is broken it says so here rather than mid-training." ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "# Import datasets FIRST: it initializes torchcodec's decoder, which must load its\n", "# libraries before torch claims them.\n", "from datasets import Audio, Dataset, Features, Value, load_dataset # noqa: F401\n", "import torch, transformers\n", "from packaging.version import parse as V\n", "\n", "assert V(transformers.__version__).release >= (5, 16), (\n", " f\"need transformers>=5.16, got {transformers.__version__}; rerun the install cell \"\n", " \"then Runtime > Restart session\")\n", "\n", "# `datasets` decodes audio only through torchcodec, which is ABI-linked to torch.\n", "# `pip install -U torchcodec` can pull a build for a different CUDA than the runtime\n", "# has; the failure then surfaces as OSError: libnvrtc.so.NN on the first decode.\n", "try:\n", " load_dataset(\"hf-internal-testing/librispeech_asr_dummy\", \"clean\",\n", " split=\"validation[:1]\").cast_column(\"audio\", Audio(sampling_rate=16000))[0]\n", "except OSError as e:\n", " raise SystemExit(\n", " f\"Audio decoding is broken ({e}).\\ntorchcodec does not match this torch \"\n", " \"build. Re-run the torchcodec pin cell above (it picks the version for your \"\n", " \"torch), then Runtime > Restart session. Reinstalling the same version will \"\n", " \"not fix it.\") from e\n", "\n", "print(f\"transformers {transformers.__version__} | torch {torch.__version__} \"\n", " f\"| cuda {torch.cuda.is_available()} | audio decode OK\")\n", "if not torch.cuda.is_available():\n", " print(\"WARNING: no GPU -- Runtime > Change runtime type > GPU\")" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "from datasets import Audio, Dataset, Features, Value, load_dataset # import before torch\n", "import numpy as np\n", "import torch\n", "from transformers import AutoProcessor, GraniteSpeech5ForCTC\n", "\n", "MODEL_ID = \"ibm-granite/granite-speech-5.0-470m-turboctc\"\n", "\n", "processor = AutoProcessor.from_pretrained(MODEL_ID)\n", "model = GraniteSpeech5ForCTC.from_pretrained(MODEL_ID, dtype=torch.float32)\n", "\n", "SAMPLING_RATE = processor.feature_extractor.sampling_rate\n", "BLANK_ID = model.config.pad_token_id # CTC blank == pad == 0\n", "\n", "print(f\"{model.num_parameters()/1e6:.0f}M params | {SAMPLING_RATE} Hz | blank id {BLANK_ID}\")" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 2. Data\n", "\n", "Any `Dataset` with these three columns works \u2014 swap `load_data()` for your own source:\n", "\n", "| column | type |\n", "|---|---|\n", "| `audio` | `{\"array\": float32[...], \"sampling_rate\": 16000}` |\n", "| `text` | `str` |\n", "| `input_length` | duration in seconds |" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "N_EVAL = 400\n", "MAX_SECONDS = 16.0\n", "\n", "\n", "def gen_rows():\n", " \"\"\"Yield one clip at a time. FLEURS calls the transcript `transcription`.\"\"\"\n", " for split in (\"train\", \"validation\"):\n", " stream = load_dataset(\n", " \"google/fleurs\", \"de_de\", split=split, streaming=True,\n", " ).cast_column(\"audio\", Audio(sampling_rate=SAMPLING_RATE))\n", " for row in stream:\n", " audio = row[\"audio\"]\n", " dur = len(audio[\"array\"]) / audio[\"sampling_rate\"]\n", " if dur <= MAX_SECONDS and row[\"transcription\"].strip():\n", " yield {\n", " \"audio\": {\"array\": np.asarray(audio[\"array\"], dtype=np.float32),\n", " \"sampling_rate\": audio[\"sampling_rate\"]},\n", " \"text\": row[\"transcription\"],\n", " \"input_length\": dur,\n", " }\n", "\n", "\n", "# from_generator streams straight to an on-disk Arrow file, so only one clip is ever\n", "# in RAM. Building a Python list first and calling Dataset.from_list() needs ~4 GB\n", "# for this corpus (plus a full copy during the map below) and OOMs a free Colab.\n", "full = Dataset.from_generator(gen_rows, features=Features({\n", " \"audio\": Audio(sampling_rate=SAMPLING_RATE),\n", " \"text\": Value(\"string\"),\n", " \"input_length\": Value(\"float32\"),\n", "}))\n", "assert {\"audio\", \"text\", \"input_length\"} <= set(full.column_names)\n", "\n", "split = full.train_test_split(test_size=N_EVAL, seed=0)\n", "train_raw, eval_raw = split[\"train\"], split[\"test\"]\n", "print(f\"train {len(train_raw)} | eval {len(eval_raw)}\")\n", "print(train_raw[0][\"text\"][:90])" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "### Normalize\n", "\n", "Two separate jobs. **Targets** get only lowercasing and character filtering \u2014 `KEEP_CHARS` must\n", "include your language's letters, since anything dropped here the model can never learn to emit.\n", "**Scoring** uses Whisper's `BasicTextNormalizer` (use `EnglishTextNormalizer` only for English)." ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "import re\n", "\n", "from transformers.models.whisper.english_normalizer import BasicTextNormalizer\n", "\n", "KEEP_CHARS = r\"a-z\u00e4\u00f6\u00fc\u00df' \" # German; English would be r\"a-z' \"\n", "_keep = re.compile(f\"[^{KEEP_CHARS}]+\")\n", "_basic = BasicTextNormalizer()\n", "\n", "\n", "def normalize(text):\n", " return re.sub(r\"\\s+\", \" \", _keep.sub(\" \", text.lower().replace(\"-\", \" \"))).strip()\n", "\n", "\n", "def prepare(ds):\n", " drop = [c for c in ds.column_names if c not in (\"audio\", \"text\", \"input_length\")]\n", " # Two things matter here:\n", " # * input_columns= -- without it `datasets` deserializes every clip's audio just\n", " # to read the text (~15 rows/s vs ~11000).\n", " # * writer_batch_size -- map() writes a new Arrow file; a small batch keeps the\n", " # transient buffer tiny instead of holding ~2 GB of audio in RAM.\n", " ds = ds.map(lambda t: {\"text\": normalize(t)}, input_columns=\"text\",\n", " remove_columns=drop, writer_batch_size=64)\n", " return ds.filter(lambda t: len(t) > 0, input_columns=\"text\",\n", " writer_batch_size=64)\n", "\n", "\n", "train_ds, eval_ds = prepare(train_raw), prepare(eval_raw)\n", "print(train_ds[0][\"text\"][:90])" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "### Check the CTC frame budget\n", "\n", "The encoder emits **12.5 frames/s** and CTC needs `len(tokens) <= frames`. Violating clips get a\n", "zeroed loss (`zero_infinity`) \u2014 they train on nothing while looking like a plateau, so drop them.\n", "\n", "If this drops more than a few percent, check your sampling rate or your script: the tokenizer is\n", "English BPE, costing ~0.23 tokens/char for German but **1.9\u20133.0 for Russian, Hindi and Japanese**,\n", "which need 18\u201331 tokens/s and cannot be aligned on this model at all." ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "fe = processor.feature_extractor\n", "SUBSAMPLE = 2 ** len(model.config.encoder_config.subsample_layers)\n", "tok = processor.tokenizer\n", "\n", "\n", "def n_frames(seconds):\n", " mel = int(seconds * SAMPLING_RATE) // fe.hop_length\n", " return (-(-mel // fe.frame_stacking)) // SUBSAMPLE\n", "\n", "\n", "def alignable(text, input_length):\n", " return 0 < len(tok(text, add_special_tokens=False)[\"input_ids\"]) <= n_frames(input_length)\n", "\n", "\n", "before = len(train_ds)\n", "train_ds = train_ds.filter(alignable, input_columns=[\"text\", \"input_length\"])\n", "eval_ds = eval_ds.filter(alignable, input_columns=[\"text\", \"input_length\"])\n", "print(f\"dropped {before - len(train_ds)}/{before} | train {len(train_ds)} eval {len(eval_ds)}\")" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 3. Collator\n", "\n", "One CTC-specific trap: pad labels with the **blank id (0)**, not `-100`. The model computes\n", "`target_lengths = (labels != pad_token_id).sum(-1)` and never looks for `-100`, so `-100` padding is\n", "silently counted as real target tokens." ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "from dataclasses import dataclass\n", "\n", "\n", "@dataclass\n", "class CTCCollator:\n", " processor: object\n", "\n", " def __call__(self, features):\n", " batch = self.processor(\n", " [f[\"audio\"][\"array\"] for f in features],\n", " text=[f[\"text\"] for f in features],\n", " sampling_rate=SAMPLING_RATE, padding=\"longest\", return_tensors=\"pt\",\n", " )\n", " labels = batch[\"labels\"]\n", " batch[\"labels\"] = labels.masked_fill(\n", " labels == self.processor.tokenizer.pad_token_id, BLANK_ID)\n", " return batch\n", "\n", "\n", "collator = CTCCollator(processor)\n", "for k, v in collator([train_ds[i] for i in range(4)]).items():\n", " print(k, tuple(v.shape))" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 4. Metric\n", "\n", "WER and CER from `generate()`. `Trainer` pads gathered batches with `-100`, which the tokenizer\n", "cannot decode, so map those back to blank first." ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "import evaluate\n", "\n", "wer_metric, cer_metric = evaluate.load(\"wer\"), evaluate.load(\"cer\")\n", "\n", "\n", "def compute_metrics(pred):\n", " ids = pred.predictions\n", " if isinstance(ids, tuple):\n", " ids = ids[0]\n", " ids = np.where(np.asarray(ids) < 0, BLANK_ID, ids)\n", " if ids.ndim == 3:\n", " ids = ids.argmax(-1)\n", "\n", " hyps = processor.batch_decode(ids, skip_special_tokens=True)\n", " refs = processor.batch_decode(\n", " np.where(np.asarray(pred.label_ids) < 0, BLANK_ID, pred.label_ids),\n", " skip_special_tokens=True)\n", "\n", " pairs = [(_basic(h), _basic(r)) for h, r in zip(hyps, refs)]\n", " pairs = [(h, r) for h, r in pairs if r]\n", " h, r = zip(*pairs)\n", " return {\"wer\": 100 * wer_metric.compute(predictions=list(h), references=list(r)),\n", " \"cer\": 100 * cer_metric.compute(predictions=list(h), references=list(r))}" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 5. Train\n", "\n", "`Trainer` would evaluate on frame-level logits, which are full of blanks and repeats \u2014 so\n", "`prediction_step` calls `generate()` to get the real greedy CTC decode.\n", "\n", "**Gate on WER, not loss.** CTC loss and WER can move in opposite directions; `load_best_model_at_end`\n", "with `metric_for_best_model=\"wer\"` keeps the checkpoint that actually decoded best." ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "from transformers import Trainer, TrainingArguments\n", "\n", "\n", "class CTCTrainer(Trainer):\n", " def prediction_step(self, model, inputs, prediction_loss_only, ignore_keys=None):\n", " inputs = self._prepare_inputs(inputs)\n", " labels = inputs.get(\"labels\")\n", " with torch.no_grad():\n", " loss = model(**inputs).loss\n", " if prediction_loss_only:\n", " return (loss, None, None)\n", " gen = model.generate(input_features=inputs[\"input_features\"],\n", " attention_mask=inputs.get(\"attention_mask\"))\n", " if gen.shape[-1] < labels.shape[-1]:\n", " gen = torch.nn.functional.pad(\n", " gen, (0, labels.shape[-1] - gen.shape[-1]), value=BLANK_ID)\n", " return (loss, gen, labels)\n", "\n", "\n", "model.gradient_checkpointing_enable()\n", "\n", "args = TrainingArguments(\n", " output_dir=\"granite-turboctc-de\",\n", " per_device_train_batch_size=8,\n", " per_device_eval_batch_size=8,\n", " gradient_accumulation_steps=2,\n", " learning_rate=1e-5,\n", " warmup_steps=100,\n", " num_train_epochs=8,\n", " bf16=torch.cuda.is_bf16_supported(),\n", " fp16=not torch.cuda.is_bf16_supported(),\n", " gradient_checkpointing=True,\n", " train_sampling_strategy=\"group_by_length\",\n", " length_column_name=\"input_length\",\n", " logging_steps=25,\n", " eval_strategy=\"steps\",\n", " eval_steps=50,\n", " save_strategy=\"steps\",\n", " save_steps=50,\n", " save_total_limit=2,\n", " load_best_model_at_end=True,\n", " metric_for_best_model=\"wer\",\n", " greater_is_better=False,\n", " remove_unused_columns=False,\n", " report_to=\"none\",\n", ")\n", "\n", "trainer = CTCTrainer(\n", " model=model, args=args, train_dataset=train_ds, eval_dataset=eval_ds,\n", " data_collator=collator, processing_class=processor, compute_metrics=compute_metrics,\n", ")" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "Baseline first \u2014 it tells you whether fine-tuning can help at all." ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "print(trainer.evaluate(metric_key_prefix=\"base\"))" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "trainer.train()" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "print(trainer.evaluate(metric_key_prefix=\"final\"))" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 6. Inspect and save" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "model.eval()\n", "rows = [eval_ds[i] for i in range(4)]\n", "inputs = processor([r[\"audio\"][\"array\"] for r in rows], sampling_rate=SAMPLING_RATE,\n", " padding=\"longest\", return_tensors=\"pt\").to(model.device, dtype=model.dtype)\n", "with torch.no_grad():\n", " hyps = processor.batch_decode(model.generate(**inputs), skip_special_tokens=True)\n", "for r, h in zip(rows, hyps):\n", " print(f\"REF: {r['text'][:90]}\\nHYP: {_basic(h)[:90]}\\n\")" ] }, { "cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": [ "model.save_pretrained(\"granite-turboctc-de-final\")\n", "processor.save_pretrained(\"granite-turboctc-de-final\")" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## Notes\n", "\n", "* **On a big GPU**, drop `gradient_checkpointing` and use `per_device_train_batch_size=32,\n", " gradient_accumulation_steps=1` \u2014 measured ~16x faster on an H100 at 21 GB peak. The defaults above\n", " fit a free Colab T4.\n", "* **`learning_rate=1e-5`** suits a large domain shift like this one. If your baseline is already\n", " strong, it is too aggressive \u2014 start at `1e-6` with a longer warmup, or freeze the lower encoder:\n", " `for p in model.encoder.layers[:8].parameters(): p.requires_grad = False`.\n", "* **More data helps most.** ~10 h against a 60k-hour pretrain is thin; expect a working pipeline\n", " rather than a production German model." ] } ], "metadata": { "accelerator": "GPU", "colab": { "provenance": [], "gpuType": "T4" }, "kernelspec": { "display_name": "Python 3", "name": "python3" }, "language_info": { "name": "python" } }, "nbformat": 4, "nbformat_minor": 0 }