DiarizationLM-Gemma-4-E4B-v1
This is not an officially supported Google product.
Overview
DiarizationLM is a Large Language Model framework designed to post-process, correct, and optimize automatic speech recognition (ASR) and speaker diarization outputs.
google/DiarizationLM-Gemma-4-E4B-v1 is built on Google's Gemma 4 E4B (4B dense parameters) foundation model and fine-tuned with Locality-Preserving Oracle Supervision across all 4 canonical speaker diarization benchmark corpora:
- Fisher English (2-speaker conversational telephone speech)
- Callhome American English (2–5 speaker informal telephone conversations)
- ICSI Meeting Corpus (3–9 speaker academic research meetings)
- AMI Meeting Corpus (4-speaker tabletop meetings)
- Foundation model: google/gemma-4-E4B
- Open-source library & scripts: https://github.com/google/speaker-id/tree/master/DiarizationLM
Unlike earlier models trained exclusively on 2-speaker telephone data (such as google/DiarizationLM-8b-Fisher-v2) or naive multi-domain SFT (which suffers from long-monologue speaker identity drift in 4+ speaker meetings), google/DiarizationLM-Gemma-4-E4B-v1 is trained on locality-preserving oracle targets that teach the model to correct lexical turn boundaries and backchannels (1 ≤ L ≤ 5 words) while preserving acoustic speaker anchors across long monologues (L ≥ 6 words). As a result, it achieves statistically significant (p < 0.0001) WDER and cpWER improvements simultaneously across all four benchmarks under standard out-of-the-box Transcript-Preserving Speaker Transfer (transfer_llm_completion), despite having half the parameter count (4B vs. 8B).
Training Configuration
- Base Architecture: Gemma 4 E4B (42 layers, hidden size 2560, hybrid 5:1 sliding-window and global attention, 262,144 vocabulary size)
- LoRA Adapter: Rank r = 256 applied to all attention (
q_proj,k_proj,v_proj,o_proj), MLP (gate_proj,up_proj,down_proj), and per-layer input (per_layer_input_gate,per_layer_projection,per_layer_model_projection) linear projections, merged into 16-bit (bfloat16) base weights and serialized to 4-bit GGUF (Q4_K_MandQ4_0) - Training Objective: Completion-only cross-entropy loss (
<prompt> --> <completion> [eod]) - Training Data: 51,063 Fisher + 20,762 Multi-Corpus (Callhome, ICSI, AMI) locality-preserving prompt-completion pairs
- Optimization: 10,000 steps, global batch size 8, AdamW (beta1 = 0.9, beta2 = 0.99), peak learning rate 1.5e-4 with 500-step linear warmup and cosine decay on 8 Google Cloud TPU v5p chips
- Prompt Segmentation Length: 4,000 characters (maximal sequence length 2,560 tokens)
Model Files Included
model.safetensors: Merged 16-bit (bfloat16) Hugging Facetransformersweights (~16.0 GB)DiarizationLM-Gemma-4-E4B-v1-q4_k_m.gguf: (Recommended GGUF) Serialized 4-bit K-quant Medium (Q4_K_M) GGUF model (~5.30 GB) forllama.cpp/ Ollama /llama-cpp-python(uses 256-element super-blocks inQ4_Kwith sensitiveattn_v/ffn_downand embedding matrices retained in 6-bitQ6_K)DiarizationLM-Gemma-4-E4B-v1-q4_0.gguf: Serialized legacy 4-bit (Q4_0) GGUF model (~5.15 GB) forllama.cpp/ Ollama /llama-cpp-pythonconfig.json,generation_config.json,tokenizer.json,tokenizer_config.json,processor_config.json,chat_template.jinja,special_tokens_map.json: Tokenizer and model configuration files
Benchmark Performance (with 95% Bootstrap Confidence Intervals)
All metrics below are micro-averaged across the full evaluation sets using the USM + turn-to-diarize baseline and scored via Hungarian-matching dynamic programming (diarizationlm.compute_metrics_on_json_dict). Ranges in brackets indicate 95% non-parametric bootstrap confidence intervals (B = 10,000 conversation-level resamples):
| Benchmark | Test Split | WER (%) | System | WDER (%) [95% CI] | cpWER (%) [95% CI] | SpkCntMAE [95% CI] | Paired ΔWDER vs. Baseline (p-value) |
|---|---|---|---|---|---|---|---|
| Fisher | TEST FULL (172 sessions) | 15.37 | Baseline (USM + Turn-to-Diarize)DiarizationLM-8b-Fisher-v2 (Llama 3 8B)DiarizationLM-Gemma-4-E4B-v1 (4B) | 5.32 [4.93, 5.74] 3.28 2.99 [2.65, 3.37] | 20.88 [19.97, 21.88] 18.37 17.62 [16.76, 18.56] | 0.215 [0.151, 0.291] --- 0.093 [0.047, 0.145] | --- --- -2.33% [-2.51, -2.16] (p < 0.0001) |
| Callhome | TEST FULL (20 calls) | 15.22 | Baseline (USM + Turn-to-Diarize)DiarizationLM-8b-Fisher-v2 (Llama 3 8B)DiarizationLM-Gemma-4-E4B-v1 (4B) | 7.74 [6.07, 9.65] 6.66 4.92 [3.46, 6.75] | 24.31 [21.27, 27.43] 23.57 20.69 [17.98, 23.66] | 0.050 [0.000, 0.150] --- 0.000 [0.000, 0.000] | --- --- -2.82% [-3.48, -2.14] (p < 0.0001) |
| ICSI | TEST FULL (3 meetings) | 29.42 | Baseline (USM + Turn-to-Diarize)DiarizationLM-Gemma-4-E4B-v1 (4B) | 14.70 [11.65, 20.29] 14.10 [10.77, 19.94] | 43.90 [39.10, 51.67] 43.32 [38.38, 51.20] | 0.333 [0.000, 1.000] 0.333 [0.000, 1.000] | --- -0.60% [-0.88, -0.35] (p < 0.0001) |
| AMI | TEST WORD FULL (16 meetings) | 24.33 | Baseline (USM + Turn-to-Diarize)DiarizationLM-Gemma-4-E4B-v1 (4B) | 15.68 [10.64, 21.11] 14.89 [9.80, 20.38] | 40.32 [32.43, 48.00] 39.57 [31.57, 47.29] | 0.500 [0.188, 0.812] 0.438 [0.188, 0.750] | --- -0.79% [-1.00, -0.60] (p < 0.0001) |
Usage
1. Python (transformers + diarizationlm)
First, install the required packages:
pip install transformers diarizationlm
Run inference on a GPU with bfloat16:
from diarizationlm import utils
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
MODEL_ID = "google/DiarizationLM-Gemma-4-E4B-v1"
HYPOTHESIS = (
"<speaker:1> Hello, how are you doing <speaker:2> today? I am doing well."
" What about <speaker:1> you? I'm doing well, too. Thank you."
)
print("Loading model...")
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, device_map="cuda")
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID, torch_dtype=torch.bfloat16, device_map="cuda"
)
print("Tokenizing input...")
inputs = tokenizer([HYPOTHESIS + " --> "], return_tensors="pt").to("cuda")
print("Generating completion...")
outputs = model.generate(
**inputs,
max_new_tokens=int(inputs.input_ids.shape[1] * 1.2),
do_sample=False,
use_cache=True,
)
print("Decoding completion...")
completion = tokenizer.batch_decode(
outputs[:, inputs.input_ids.shape[1]:], skip_special_tokens=True
)[0]
completion = utils.truncate_suffix_and_tailing_text(completion, " [eod]")
print("Transferring completion to hypothesis text...")
transferred_completion = utils.transfer_llm_completion(completion, HYPOTHESIS)
print("========================================")
print("Hypothesis:", HYPOTHESIS)
print("========================================")
print("Completion:", completion)
print("========================================")
print("Transferred completion:", transferred_completion)
print("========================================")
2. GGUF (llama.cpp)
You can also run the quantized DiarizationLM-Gemma-4-E4B-v1-q4_k_m.gguf model directly with llama.cpp:
llama-cli \
-m DiarizationLM-Gemma-4-E4B-v1-q4_k_m.gguf \
-p "<speaker:1> Hello, how are you doing <speaker:2> today? I am doing well. What about <speaker:1> you? I'm doing well, too. Thank you. --> " \
--temp 0.0 \
-n 128
Citation
@inproceedings{wang24h_interspeech,
title = {{DiarizationLM: Speaker Diarization Post-Processing with Large Language Models}},
author = {Quan Wang and Yiling Huang and Guanlong Zhao and Evan Clark and Wei Xia and Hank Liao},
year = {2024},
booktitle = {Interspeech 2024},
pages = {3754--3758},
doi = {10.21437/Interspeech.2024-209},
}