Here is the official fine-tuning code used for Vāgdhenu, plus a clearer, self-contained version you can study and adapt.
1. Official script from the repo
training/finetune_indicf5.py (exactly as shipped)
1"""Fine-tune IndicF5 (F5/flow-matching) on pilot_reciter 5h. 2Warm-start from the GRN-corrected checkpoint loaded directly into the CFM. 3DDP via `accelerate launch --multi_gpu`.""" 4 5import argparse 6import torch 7from f5_tts.infer.utils_infer import load_model 8from f5_tts.model import DiT, Trainer 9from f5_tts.model.dataset import load_dataset 10 11ap = argparse.ArgumentParser() 12ap.add_argument("--vocab", required=True) 13ap.add_argument("--warm", required=True) # IndicF5 / previous checkpoint 14ap.add_argument("--data_dir", required=True) # prepared dataset folder 15ap.add_argument("--save_dir", required=True) 16ap.add_argument("--wandb_name", required=True) 17ap.add_argument("--epochs", type=int, default=600) 18ap.add_argument("--lr", type=float, default=1e-5) 19ap.add_argument("--bs", type=int, default=19200) # frame batch size 20ap.add_argument("--bstype", default="frame") 21ap.add_argument("--warmup", type=int, default=500) 22ap.add_argument("--save_per", type=int, default=2000) 23a = ap.parse_args() 24 25# Production architecture (same as Vāgdhenu) 26CFG = dict(dim=1024, depth=22, heads=16, ff_mult=2, text_dim=512, conv_layers=4) 27 28cfm = load_model(DiT, CFG, mel_spec_type="vocos", vocab_file=a.vocab, device="cpu") 29 30ck = torch.load(a.warm, map_location="cpu", weights_only=True) 31sd = ck.get("model_state_dict", ck) 32miss, unexp = cfm.load_state_dict(sd, strict=False) 33print("[warm-start] missing(non-melspec):", 34 len([m for m in miss if "mel_spec" not in m]), 35 "| unexpected:", len(unexp), flush=True) 36 37trainer = Trainer( 38 cfm, 39 epochs=a.epochs, 40 learning_rate=a.lr, 41 num_warmup_updates=a.warmup, 42 save_per_updates=a.save_per, 43 last_per_steps=a.save_per, 44 checkpoint_path=a.save_dir, 45 batch_size=a.bs, 46 batch_size_type=a.bstype, 47 max_samples=64, 48 grad_accumulation_steps=1, 49 max_grad_norm=1.0, 50 logger="wandb", 51 wandb_project="indicf5-sanskrit", 52 wandb_run_name=a.wandb_name, 53 mel_spec_type="vocos", 54 log_samples=False, 55) 56 57MELKW = dict( 58 n_fft=1024, 59 hop_length=256, 60 win_length=1024, 61 n_mel_channels=100, 62 target_sample_rate=24000, 63 mel_spec_type="vocos", 64) 65 66train_dataset = load_dataset( 67 "indicf5", 68 "custom", 69 dataset_type="CustomDatasetPath", 70 mel_spec_kwargs=MELKW, 71 data_dir=a.data_dir, 72) 73 74trainer.train(train_dataset)
How they launched it (production recipe)
1# After scripts/setup.sh and data preparation 2accelerate launch --multi_gpu training/finetune_indicf5.py \ 3 --vocab path/to/vocab.txt \ 4 --warm path/to/indicf5_or_previous.pt \ 5 --data_dir path/to/prepared_kannada_clips \ 6 --save_dir checkpoints/vagdhenu_ft \ 7 --wandb_name kannada_5h_ft \ 8 --epochs 600 \ 9 --lr 1e-5 \ 10 --bs 19200 \ 11 --bstype frame
Key production settings from the tech report:
- LR = 1e-5
- Frame batch ≈ 12800–19200 (~28 GB VRAM)
- bf16
- 600 epochs on ~5 h single-speaker chant data
- Kannada-script text (never raw Devanagari)
- Later a short voice-steering retrain on ~179 paired clips produced the final production voice
2. Educational / self-contained fine-tune loop
(You can run this after installing IndicF5 / F5-TTS)
1""" 2Educational fine-tune script for Vāgdhenu-style IndicF5 DiT. 3Assumes: 4 - IndicF5 / F5-TTS is installed 5 - Dataset is prepared as CustomDatasetPath (audio + text_kannada) 6 - You have a warm-start checkpoint (.pt or .safetensors) 7""" 8 9import os 10import argparse 11from pathlib import Path 12 13import torch 14from accelerate import Accelerator 15from torch.utils.data import DataLoader 16 17# F5-TTS / IndicF5 imports 18from f5_tts.model import DiT, CFM, Trainer 19from f5_tts.model.dataset import load_dataset 20from f5_tts.infer.utils_infer import load_model 21 22 23def parse_args(): 24 p = argparse.ArgumentParser(description="Fine-tune IndicF5 for Sanskrit chant (Vāgdhenu style)") 25 p.add_argument("--vocab", type=str, required=True, help="path to vocab.txt") 26 p.add_argument("--warm", type=str, required=True, help="warm-start checkpoint (.pt)") 27 p.add_argument("--data_dir", type=str, required=True, help="prepared dataset directory") 28 p.add_argument("--save_dir", type=str, default="checkpoints/vagdhenu_ft") 29 p.add_argument("--epochs", type=int, default=600) 30 p.add_argument("--lr", type=float, default=1e-5) 31 p.add_argument("--batch_frames", type=int, default=16000) # reduce if OOM 32 p.add_argument("--warmup_steps", type=int, default=500) 33 p.add_argument("--save_every", type=int, default=2000) 34 p.add_argument("--max_samples", type=int, default=64) 35 p.add_argument("--wandb_project", type=str, default="vagdhenu-finetune") 36 p.add_argument("--wandb_run", type=str, default="sanskrit_chant_ft") 37 p.add_argument("--bf16", action="store_true", default=True) 38 return p.parse_args() 39 40 41def main(): 42 args = parse_args() 43 os.makedirs(args.save_dir, exist_ok=True) 44 45 # ------------------------------------------------------------------ 46 # 1. Model (exact production architecture) 47 # ------------------------------------------------------------------ 48 CFG = dict( 49 dim=1024, 50 depth=22, 51 heads=16, 52 ff_mult=2, 53 text_dim=512, 54 conv_layers=4, 55 ) 56 57 # load_model returns a CFM wrapping the DiT 58 model = load_model( 59 DiT, 60 CFG, 61 mel_spec_type="vocos", 62 vocab_file=args.vocab, 63 device="cpu", 64 ) 65 66 # ------------------------------------------------------------------ 67 # 2. Warm-start 68 # ------------------------------------------------------------------ 69 ckpt = torch.load(args.warm, map_location="cpu", weights_only=True) 70 state = ckpt.get("ema_model_state_dict", ckpt.get("model_state_dict", ckpt)) 71 72 # strip possible prefixes from EMA / compiled models 73 cleaned = {} 74 for k, v in state.items(): 75 new_k = k 76 for prefix in ("ema_model.", "module.", "_orig_mod."): 77 if new_k.startswith(prefix): 78 new_k = new_k[len(prefix):] 79 cleaned[new_k] = v 80 81 missing, unexpected = model.load_state_dict(cleaned, strict=False) 82 print(f"[warm-start] missing: {len(missing)} | unexpected: {len(unexpected)}") 83 84 # ------------------------------------------------------------------ 85 # 3. Dataset (Kannada-routed text is mandatory) 86 # ------------------------------------------------------------------ 87 mel_kwargs = dict( 88 n_fft=1024, 89 hop_length=256, 90 win_length=1024, 91 n_mel_channels=100, 92 target_sample_rate=24000, 93 mel_spec_type="vocos", 94 ) 95 96 train_ds = load_dataset( 97 "indicf5", 98 "custom", 99 dataset_type="CustomDatasetPath", 100 mel_spec_kwargs=mel_kwargs, 101 data_dir=args.data_dir, 102 ) 103 104 # ------------------------------------------------------------------ 105 # 4. Trainer (F5-TTS official Trainer) 106 # ------------------------------------------------------------------ 107 trainer = Trainer( 108 model, 109 epochs=args.epochs, 110 learning_rate=args.lr, 111 num_warmup_updates=args.warmup_steps, 112 save_per_updates=args.save_every, 113 last_per_steps=args.save_every, 114 checkpoint_path=args.save_dir, 115 batch_size=args.batch_frames, 116 batch_size_type="frame", # critical: frame-based batching 117 max_samples=args.max_samples, 118 grad_accumulation_steps=1, 119 max_grad_norm=1.0, 120 logger="wandb", 121 wandb_project=args.wandb_project, 122 wandb_run_name=args.wandb_run, 123 mel_spec_type="vocos", 124 log_samples=False, 125 ) 126 127 print("Starting fine-tune …") 128 print(f" epochs : {args.epochs}") 129 print(f" lr : {args.lr}") 130 print(f" batch (frames) : {args.batch_frames}") 131 print(f" save dir : {args.save_dir}") 132 133 trainer.train(train_ds) 134 print("Fine-tuning finished.") 135 136 137if __name__ == "__main__": 138 main()
3. Data preparation checklist (most important part)
Vāgdhenu expects a CustomDatasetPath layout roughly like:
data_dir/
├── audio/ # 24 kHz mono wavs
│ ├── clip001.wav
│ └── …
└── metadata.csv # or the format load_dataset expects
metadata.csv columns used by the project:
file_nametext_kannada← this is the text the model sees- (optional)
text_devanagari,text_slp1,duration,meter
Critical rules from the tech report
- Always convert Sanskrit → Kannada script (never feed raw Devanagari).
- Keep long vowels, terminal visarga, retroflex aspirates, and conjuncts clean.
- Reference clips for inference must obey the half-reference rule.
- ~5 h of carefully recorded single-speaker chant is enough for a strong clone.
You can load the public corpus with:
1from datasets import load_dataset 2ds = load_dataset("prathoshap/vagdhenu-data", "style_b")
4. Quick start commands
1# 1. Environment (from the repo) 2git clone https://github.com/prathoshap/vagdhenu.git 3cd vagdhenu 4bash scripts/setup.sh # installs torch 2.4.1+cu121, IndicF5, BigVGAN, weights 5 6# 2. Prepare your data (Kannada text + 24 kHz wavs) 7# → put them under e.g. data/my_chant/ 8 9# 3. Fine-tune (single GPU example) 10python training/finetune_indicf5.py \ 11 --vocab models/vocab.txt \ 12 --warm models/indicf5_base.pt \ 13 --data_dir data/my_chant \ 14 --save_dir checkpoints/my_ft \ 15 --wandb_name my_sanskrit_ft \ 16 --epochs 300 \ 17 --lr 1e-5 \ 18 --bs 12000 19 20# Multi-GPU 21accelerate launch --multi_gpu training/finetune_indicf5.py ...
5. After fine-tuning
- The Trainer saves checkpoints as
.pt(often withema_model_state_dict). - For production Vāgdhenu they also did a short voice-steering retrain on paired clips.
- Always keep the fine-tuned BigVGAN-v2 vocoder (
voc_bigvgan_EMA_….pth); the original Vocos vocoder is not good enough for long vowels.
This is the exact fine-tuning path that produced the ~4.6 MOS Sanskrit chant model. Start with the official finetune_indicf5.py, keep the architecture fixed, use Kannada text, and the low learning rate + large frame batch.