Koko-TTS: Ultra-Lightweight Flow-Matching Text-to-Speech

Koko-TTS is an ultra-lightweight, high-fidelity Text-to-Speech (TTS) model with ~24.6M parameters. Built upon a Flow-Matching framework, it combines a Matcha-style UNet architecture conditioned on RoPE-based text representations with the highly efficient Vocos 24kHz neural vocoder. The model is designed for ultra-fast, real-time speech synthesis without compromising on audio quality.

Model Details

Metric / Parameter Value
Model Size ~24.6M parameters
Sample Rate 24,000 Hz
Vocoder Vocos Mel 24kHz (charactr/vocos-mel-24khz)
Available Voices 128 Speaker IDs (0 to 127)
Training Dataset saki22/libritts-r-128spk-vocos-mel
Language English (Model is trained on English, but the tokenizer is multilingual)

Audio Samples

Speaker Transcript Audio Sample
Speaker 0 "The morning was quiet, and a gentle breeze moved through the trees. Somewhere in the distance, birds were singing, while the first light of day slowly filled the sky."
Speaker 46 "Well, here we are, take a breath, relax, and listen, sometimes, a quiet moment is all we need."
Speaker 123 "The quick brown fox jumps over the lazy dog. This is a demonstration of koko, an ultra-lightweight, high-quality text-to-speech model designed to combine fast inference with exceptional audio quality."

Quickstart

1. Installation

Ensure you have the required libraries installed:

pip install -q torch torchaudio transformers vocos tokenizers

2. Inference

Generating speech is straightforward using the transformers library.

import torch
import torchaudio
from transformers import AutoModel

device = "cuda" if torch.cuda.is_available() else "cpu"

# Load the model with custom code execution enabled
model = AutoModel.from_pretrained(
    "saki22/koko-tts", 
    trust_remote_code=True
).to(device)

# Generate speech waveform
audio = model.inference(
    text="Hello! This is Koko-TTS running fast and smooth.",
    spk_id=0,               # Choose a speaker between 0 and 127
    temperature=0.667,      # Controls variance/expressiveness
    cfg_strength=1.5,       # Classifier-Free Guidance strength
    n_steps=16,             # Number of ODE solver steps
    solver="euler",         # ODE solver type (euler or midpoint)
    sway_coef=-1.0,         # Sway schedule coefficient
    length_scale=1.0        # Speech speed/pace control (<1.0 faster, >1.0 slower)
)

# Save the generated audio to a .wav file
torchaudio.save("output.wav", audio.unsqueeze(0), sample_rate=24000)

# Optional: Play directly if using Jupyter Notebook / Google Colab
# from IPython.display import Audio
# Audio(audio.numpy(), rate=24000, autoplay=True)

Fine-Tuning Guide

Koko-TTS is highly modular and designed to be easily fine-tuned on custom datasets. Fine-tuning requires two extracted targets from your audio: 100-dimensional Mel-Spectrograms and Token Durations (the number of mel frames corresponding to each character/token).

1. Preprocessing (Mel & Duration Extraction)

A. Mel Spectrogram Extraction (Vocos)

Audio must be resampled to 24,000 Hz. Use the Vocos feature extractor to produce matching 100-channel mel-spectrograms:

import torch
from vocos import Vocos

device = "cuda" if torch.cuda.is_available() else "cpu"
vocos = Vocos.from_pretrained("charactr/vocos-mel-24khz").to(device)

def extract_mel(audio_tensor_24k):
    # audio_tensor_24k shape: [1, T_samples]
    with torch.no_grad():
        mel = vocos.feature_extractor(audio_tensor_24k.to(device)).squeeze(0)
    # Returns mel shape: [100, T_frames]
    return mel

B. Duration Extraction (Forced Alignment)

durations is an integer tensor (torch.long) with the same sequence length as input_ids, specifying how many mel frames each character/token lasts (the sum of durations must equal the total number of mel frames).

You can compute accurate token-level durations using a forced aligner such as torchaudio.pipelines.MMS_FA:

import torch
import torchaudio

device = "cuda" if torch.cuda.is_available() else "cpu"

bundle = torchaudio.pipelines.MMS_FA
aligner = bundle.get_model().to(device)
tokenizer = bundle.get_tokenizer()
aligner_dict = bundle.get_dict()

def extract_durations(audio_tensor_16k, text, total_mel_frames):
    tokens = tokenizer(text)
    token_ids = torch.tensor([[aligner_dict[c] for c in tokens]], device=device)

    with torch.no_grad():
        emission, _ = aligner(audio_tensor_16k.to(device))
        spans = torchaudio.functional.forced_align(emission, token_ids)

    # Frame lengths from forced alignment
    durations = torch.tensor([s.end - s.start for s in spans[0]], dtype=torch.float32)

    # Scale alignment frames to match Vocos 24kHz mel frames
    scale = total_mel_frames / durations.sum().clamp(min=1.0)
    durations = torch.clamp(torch.round(durations * scale), min=1).long()
    return durations

2. Training Loop Example

import torch
from transformers import AutoModel

device = "cuda" if torch.cuda.is_available() else "cpu"

# Load model for training
model = AutoModel.from_pretrained("saki22/koko-tts", trust_remote_code=True).to(device)
model.train()

optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-2)

for batch in dataloader:
    optimizer.zero_grad()

    # Forward pass calculates Flow Matching loss and Duration loss
    outputs = model(
        input_ids=batch["input_ids"].to(device),
        durations=batch["durations"].to(device),          # Required for training duration predictor
        mel_target=batch["mel_target"].to(device),        # 100-dim mel spectrogram
        mel_lengths=batch["mel_lengths"].to(device),      # Length of each mel sequence
        spk_id=batch["speaker_ids"].to(device)            # Speaker ID (0-127)
    )

    loss = outputs["loss"]
    loss.backward()

    # Gradient clipping is recommended for stable training
    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
    optimizer.step()

    print(
        f"Total Loss: {loss.item():.4f} | "
        f"Mel Flow: {outputs['flow_loss'].item():.4f} | "
        f"Duration: {outputs['duration_loss'].item():.4f}"
    )

Architecture Highlights

  • Text Encoder: A RoPE-based Transformer encoder that processes characters natively.
  • Duration Predictor: A robust convolution-based module conditioned on speaker embeddings.
  • Decoder: A Matcha-style 1D UNet utilizing SnakeBeta activations, ResNet blocks, and Multi-head Self-Attention, trained via continuous normalizing flows (Flow-Matching).

License & Citation

This project is open-sourced under the Apache-2.0 License. If you use Koko-TTS in your research or project, please cite it as:

@misc{koko_tts_2026,
  author = {saki22},
  title = {Koko-TTS: Lightweight Flow-Matching Text-to-Speech},
  year = {2026},
  publisher = {Hugging Face},
  howpublished = {\url{https://huggingface.co/saki22/koko-tts}}
}
Downloads last month
288
Safetensors
Model size
24.6M params
Tensor type
F32
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Dataset used to train saki22/koko-tts

Space using saki22/koko-tts 1