train_chatterbox / Memory.md
ibibek's picture
Initial upload: Nepali Chatterbox fine-tuning kit
94336a6 verified
|
Raw History Blame Contribute Delete
15.3 kB
# Nepali Chatterbox Fine-Tune β€” Project Memory
Status as of this writing: **LoRA adapter trained, inference pipeline working, quality
being evaluated.** This doc is a handoff for continuing in a fresh session.
## Goal
Fine-tune Chatterbox TTS to speak **Nepali**, which is not one of Chatterbox's ~23
officially supported languages. Approach: fine-tune *from* the Hindi checkpoint
(`ResembleAI/Chatterbox-Multilingual-hi`) rather than the generic multilingual base,
since Hindi and Nepali share the Devanagari script and are linguistically related β€”
and reuse Hindi's `"[hi]"` language tag for Nepali text ("Option 1") rather than
extending the tokenizer vocabulary with a new `"[ne]"` token ("Option 2", not done).
## Data source (external, read-only)
**`/home/bu/Desktop/ttsdata/make_tts_data/`** β€” a separate, independently-running
project (not part of this fine-tune folder) that generates synthetic Nepali speech
via the Gemini API TTS (`gemini-3.1-flash-tts-preview`, voice **"Leda"**, `ne-NP`,
24kHz) from a source spreadsheet (`speech_text_Nepali.xlsx`). Tracked in
`outputs/manifest.csv` (columns include `id, text, audio_path, status, duration_sec`)
with audio in `outputs/wavs/<id>.wav`.
- As of the last training run: **5,650 successful clips, ~22.84 hours**, single voice.
This generator was *still running* (retrying rate-limited rows, adding more) β€” re-run
`prepare_dataset.py` (see below) to pick up more data later, it's cheap.
- **Never delete, move, or modify anything under `ttsdata/`.** This fine-tune project
only ever reads from it, via symlinks (see `prepare_dataset.py`).
- It's synthetic (Gemini TTS) speech, single voice, not human recordings β€” sets a
quality ceiling: this fine-tune can't sound better than "Leda" does, and inherits
any of her artifacts.
## Where things live
```
/home/bu/Desktop/nepali-finetune/
β”œβ”€β”€ .venv/ # dedicated venv, see "Environment" below
β”œβ”€β”€ pretrained_models/ # base checkpoint (see "Base checkpoint" below)
β”œβ”€β”€ MyTTSDataset/
β”‚ β”œβ”€β”€ metadata.csv # LJSpeech format, built by prepare_dataset.py
β”‚ β”œβ”€β”€ wavs/ # SYMLINKS into ../ttsdata/.../outputs/wavs/
β”‚ └── preprocess/ # .pt tensors cached by train.py (large, regeneratable)
β”œβ”€β”€ speaker_reference/
β”‚ └── reference.wav # copy of ttsdata's sample_Leda.wav (12.6s, 24kHz)
β”œβ”€β”€ chatterbox_output/
β”‚ β”œβ”€β”€ new_lang_adapter/ # FINAL LoRA adapter (real PEFT format)
β”‚ β”œβ”€β”€ checkpoint-500/, -1000/, -1500/, -1770/ # HF Trainer periodic checkpoints
β”‚ β”‚ # (different, non-PEFT state-dict format --
β”‚ β”‚ # see test_tts.py's load_engine())
β”‚ └── runs/ # tensorboard logs
β”œβ”€β”€ prepare_dataset.py # ttsdata manifest -> MyTTSDataset (WRITTEN BY US)
β”œβ”€β”€ test_tts.py # modular inference helpers (WRITTEN BY US)
β”œβ”€β”€ Inference.ipynb # user's notebook, wired to test_tts.py
β”œβ”€β”€ train.py, inference.py, # from the gokhaneraslan/chatterbox-finetuning
β”‚ merge_lora.py, setup.py # toolkit, used mostly as-is (a few patches, below)
β”œβ”€β”€ src/ # toolkit internals, incl. vendored src/chatterbox_/
β”‚ β”œβ”€β”€ config.py # TrainConfig -- all hyperparams
β”‚ β”œβ”€β”€ chatterbox_/tts.py # PATCHED (see "Bugs fixed" below)
β”‚ └── chatterbox_/models/t3/inference/alignment_stream_analyzer.py # PATCHED (2 bugs)
β”œβ”€β”€ train.log # full log of the completed training run
β”œβ”€β”€ requirements.txt # PATCHED (see "Environment" below)
└── Memory.md # this file
```
Sibling project `/home/bu/Desktop/chatterbox/` is the separate, earlier
inference-only Chatterbox setup (English/Hindi, not Nepali) β€” not covered here.
## Environment
Dedicated venv at `.venv`, **not** shared with `/home/bu/Desktop/chatterbox/.venv`.
Key deviations from the toolkit's stock `requirements.txt` (all documented inline
in `requirements.txt` itself):
- **torch==2.9.1+cu128 / torchaudio==2.9.1+cu128**, installed separately from
`https://download.pytorch.org/whl/cu128` β€” the toolkit's pinned `torch==2.6.0` has
no kernels for the RTX 5090 (Blackwell, compute capability sm_120): errors with
*"no kernel image is available for execution on the device."*
- **torchcodec==0.9.1** (version-matched to torch 2.9.1 per
[pytorch/torchcodec's compatibility table](https://github.com/pytorch/torchcodec#installing-torchcodec))
β€” needed for `torchaudio.save()`/`.load()` in torchaudio>=2.9. Also needs a
**system-wide FFmpeg** (major version 4-8); the user installed this via
`sudo apt install ffmpeg` (Ubuntu 22.04, FFmpeg 4.4.2) partway through this project.
- **peft==0.20.0**, not the toolkit's pinned `0.17.1` β€” 0.17.1 imports
`transformers.HybridCache`, which doesn't exist in `transformers==5.2.0`
(`ImportError` on startup).
- `chatterbox-tts==0.1.2` pin dropped entirely β€” the toolkit vendors its own copy of
the model code (`src/chatterbox_/`) and never actually imports the pip package;
the pin was only there to pull in transitive deps, which are now listed explicitly
in `requirements.txt` (transformers==5.2.0, diffusers==0.29.0, resemble-perth,
conformer==0.3.2, s3tokenizer, tokenizers, einops, scipy).
- `resemble-perth` unpinned (PyPI only has up to 1.0.1 for this Python version).
To recreate: `python3 -m venv .venv && source .venv/bin/activate && pip install
torch==2.9.1+cu128 torchaudio==2.9.1+cu128 --index-url https://download.pytorch.org/whl/cu128
&& pip install -r requirements.txt torchcodec==0.9.1`
**Jupyter kernel**: registered as `nepali-finetune` (display name "Nepali Finetune")
via `python -m ipykernel install --user --name nepali-finetune --display-name
"Nepali Finetune"` from inside this venv. `Inference.ipynb` must use this kernel, not
the sibling chatterbox project's kernel (which lacks `peft` and `src.chatterbox_`).
## Base checkpoint ("Option 1": reuse Hindi, no vocab extension)
The toolkit's `setup.py` normally downloads the **English** base checkpoint. Instead,
`pretrained_models/` was populated manually, reusing files already cached (via
huggingface_hub) from the sibling chatterbox project's earlier work:
| File in `pretrained_models/` | Actual source |
|---|---|
| `t3_cfg.safetensors` | `t3_hi.safetensors` from `ResembleAI/Chatterbox-Multilingual-hi` (Hindi T3, vocab=2454) |
| `s3gen.safetensors` | `s3gen_v3.safetensors` from the same Hindi repo |
| `tokenizer.json` | `grapheme_mtl_merged_expanded_v1.json` (multilingual grapheme vocab, 2454 tokens, `"[hi]"` = token id 722) |
| `ve.safetensors` | from base `ResembleAI/chatterbox` repo (voice encoder, language-independent) |
| `conds.pt` | from base `ResembleAI/chatterbox` repo (English default-voice conditioning; not functionally used since we always pass an explicit `audio_prompt_path`, but required to exist by the toolkit's `check_pretrained_models()` file check) |
`src/config.py`'s `new_vocab_size = 2454` already matches this exactly, so weight
loading is a straight copy with **no resizing/mean-init needed** (confirmed in logs:
*"Embedding layer: 2454 tokens preserved."* / *"Output head: 2454 tokens preserved."*).
## "Option 1" tag handling β€” the trickiest correctness detail
The toolkit's vendored `ChatterboxTTS` uses the plain `EnTokenizer`, which has **no**
concept of upstream's `"[lang]"` tag prepending (unlike `MTLTokenizer.encode(text,
language_id=...)`). To reuse the `"[hi]"` tag for Nepali anyway, text must be manually
pre-tagged and normalized *before* it reaches `EnTokenizer`.
**Verified empirically** (in the live conversation, not just assumed) that:
```python
MTLTokenizer.encode(text, language_id="hi")
```
produces **identical token ids** to:
```python
EnTokenizer.encode(f"[hi]{unicodedata.normalize('NFKD', text.lower())}")
```
This is implemented in two places, both doing the same transform:
- `prepare_dataset.py` β€” bakes it into the `normalized_text` column of `metadata.csv`
(which `preprocess_ljspeech.py` prefers over the raw `raw_text` column)
- `test_tts.py`'s `tag_and_normalize()` β€” applied per-sentence at inference time
If you ever see garbled/wrong-language output, check whether this tagging is being
applied correctly β€” it's the single easiest thing to silently break.
## Code patches made to the vendored toolkit (not upstream fixes, ours)
1. **`src/chatterbox_/tts.py`**, `ChatterboxTTS.from_local()`: hardcoded `T3()`
(defaults to English config, vocab=704). Since our checkpoint is multilingual
(vocab=2454), patched to `T3(T3Config.multilingual())`. Without this, loading
throws a `size mismatch` error on `text_emb.weight`/`text_head.weight`.
2. **`src/chatterbox_/models/t3/inference/alignment_stream_analyzer.py`** β€” this file
is only active for multilingual models (`hp.is_multilingual` gate in `t3.py`).
Two bugs found and fixed here:
- **sdpa/eager crash**: didn't switch `tfmr.config._attn_implementation` from
`'sdpa'` to `'eager'` before setting `output_attentions=True`. Newer
`transformers` (5.2.0) hard-errors on this combination
(*"output_attentions attribute is not supported when using sdpa"*). Fixed by
adding the sdpa→eager fallback (mirrors a fix upstream Chatterbox made in a
newer version of this same file, which was otherwise deleted entirely upstream).
- **Over-aggressive repetition guard**: comment said *"3x same token in a row"*
but the code checked `generated_tokens[-2:]` (only 2 tokens), force-stopping
generation on any 2 consecutive identical speech-codec tokens β€” a common,
usually benign pattern (sustained sounds naturally repeat a token across
consecutive ~40ms frames), not necessarily a hallucination. Fixed to check
`generated_tokens[-3:]` (3 tokens), matching the stated intent. Effect: test
utterance duration went 7.6s β†’ 10.3s on identical text/seed, and remaining
triggers now land on the same token ID consistently near natural sentence-end
(looks like it's now correctly catching trailing silence, not live speech β€”
**not 100% confirmed, pending the user's listening judgment**).
## Training run (completed)
- `prepare_dataset.py` (no args) β†’ 5,650 rows, 22.84h, symlinked wavs.
- Config: `is_turbo=False`, `is_lora=True`, `lora_r=128`, `lora_alpha=256`,
`batch_size=32`, `num_epochs=10`, `learning_rate=1e-4`,
`lora_target_modules=["q_proj","k_proj","v_proj","o_proj","gate_proj","up_proj","down_proj","spkr_enc"]`,
`lora_modules_to_save=["text_emb","text_head"]`.
- Trainable params: **95,629,312 / 631,618,560 (15.14%)**.
- Preprocessing (offline feature extraction: speaker embeddings, S3 speech tokens,
text tokens): ~8 minutes for 5,650 clips (~12 it/s).
- Training: ~60 minutes (3,617s), 1,770 steps (5,650Γ·32Γ—10), 100% GPU util,
~14.4/32.6GB VRAM (compute-bound, not memory-bound β€” didn't try a bigger batch
since it wouldn't have helped throughput).
- Loss: 4.892 (epoch 1) β†’ 2.922 (epoch 10), steady decrease, no instability
(grad_norm stayed in a sane 1.3–2.4 range throughout).
- Launched via `nohup python3 train.py > train.log 2>&1 & disown` β€” fully detached
from the shell so it survives independent of any tool/session timeout. Do this
again for any re-run of comparable length.
- **Before committing to the full run**, a smoke test was done first: `num_epochs=1,
batch_size=4` temporarily in `config.py`, `prepare_dataset.py --limit 8`, ran
`train.py`, confirmed weight loading/LoRA wiring/one training step/checkpoint save
all worked, *then* reverted config and regenerated the full dataset. This caught
the peft version bug and the multilingual-vocab bug cheaply, before an hour-long
run. **Recommend the same pattern for any future config/code changes.**
## Inference / testing tools
- **`test_tts.py`** (import this from notebooks; also runnable standalone):
- `load_engine(checkpoint="new_lang_adapter", device=None, cfg=None)` β€” loads
either the final adapter (PEFT format) or an intermediate Trainer checkpoint by
name (e.g. `"checkpoint-1000"`), auto-detecting which format it's given (checks
for `adapter_config.json`). Intermediate checkpoints have a different state-dict
key structure (`t3.base_model.model....default...`, full weights incl. frozen
base) vs. the final adapter's clean PEFT format
(`base_model.model....lora_A.weight`, adapter-only + `text_emb`/`text_head`).
- `synthesize(engine, text, audio_prompt_path=DEFAULT_REFERENCE, language_tag="hi",
trim_silence=True, seed=None, **gen_kwargs)` β€” splits text into sentences
(Devanagari-aware: handles `ΰ₯€`/`ΰ₯₯`, not just `.?!`), tags+normalizes each, generates,
VAD-trims, concatenates with short pauses. Returns `(sample_rate, np.ndarray)`.
- `tag_and_normalize`, `split_sentences`, `save_wav`, `set_seed` β€” smaller pieces,
all independently reusable.
- `REPO_ROOT` resolved via `__file__` with a `Path.cwd()` fallback (for when the
module's source is pasted into a notebook cell rather than imported, where
`__file__` doesn't exist).
- **`Inference.ipynb`**: 4 cells β€” import, load+synthesize+save demo, inline
`IPython.display.Audio` playback, commented-out template for comparing a
checkpoint. Verified end-to-end via `jupyter nbconvert --execute` with the
`nepali-finetune` kernel.
- **`inference.py`** (toolkit's own script): still present, `TEXT_TO_SAY`/
`AUDIO_PROMPT` module-level vars edited to a Nepali test sentence and
`speaker_reference/reference.wav`, but **superseded by `test_tts.py`** for anything
beyond a one-off CLI run β€” it hardcodes text/path rather than being callable.
## Current status / open question
Training completed successfully and the pipeline produces real, on-topic Nepali
speech in the target voice. The main open question is **subjective audio quality** β€”
specifically whether generation now runs to natural completion (after the repetition-
guard fix) or is still cutting off. Waiting on the user's listening judgment on
`test_output_fixed.wav` / the notebook's `test_output.wav` before deciding next steps.
**If quality needs more work, options in rough order of effort:**
1. Try other seeds / generation params (`temperature`, `repetition_penalty`, `cfg_weight`)
β€” cheap, no retraining needed, use `test_tts.py`.
2. Compare against earlier checkpoints (`load_engine("checkpoint-500")` etc.) β€” maybe
10 epochs over-trained on a single-speaker synthetic set; cheap, no retraining.
3. Retry the ~2,558 rate-limited ("failed", mostly HTTP 429) rows in ttsdata's
manifest to grow the dataset further, then re-run `prepare_dataset.py` + `train.py`.
4. Adjust LoRA hyperparameters (rank/alpha) or epoch count and retrain.
**If quality is good:** run `merge_lora.py` to bake the adapter into a standalone
`.safetensors` file (no PEFT dependency needed at inference time), then do a broader
multi-sentence/multi-seed evaluation.