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:

  1. Fisher English (2-speaker conversational telephone speech)
  2. Callhome American English (2–5 speaker informal telephone conversations)
  3. ICSI Meeting Corpus (3–9 speaker academic research meetings)
  4. AMI Meeting Corpus (4-speaker tabletop meetings)

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_M and Q4_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 Face transformers weights (~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) for llama.cpp / Ollama / llama-cpp-python (uses 256-element super-blocks in Q4_K with sensitive attn_v / ffn_down and embedding matrices retained in 6-bit Q6_K)
  • DiarizationLM-Gemma-4-E4B-v1-q4_0.gguf: Serialized legacy 4-bit (Q4_0) GGUF model (~5.15 GB) for llama.cpp / Ollama / llama-cpp-python
  • config.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...")
prompt = f"<|turn>user\n{HYPOTHESIS} --> <turn|>\n<|turn>model\n"
inputs = tokenizer([prompt], 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 $'<|turn>user\n<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. --> <turn|>\n<|turn>model\n' \
  --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},
}
Downloads last month
1,869
Safetensors
Model size
8B params
Tensor type
BF16
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Model tree for google/DiarizationLM-Gemma-4-E4B-v1

Quantized
(30)
this model
Quantizations
1 model

Spaces using google/DiarizationLM-Gemma-4-E4B-v1 2

Paper for google/DiarizationLM-Gemma-4-E4B-v1