Skip to main content

DiarizationLM

Python application PyPI Version Python Versions Downloads codecov arxiv Daily Papers HuggingFace HuggingFace Space

Table of contents

Overview

Here we open source some functions and tools used in the DiarizationLM paper.

We also have open source models on Hugging Face: https://huggingface.co/google/DiarizationLM-8b-Fisher-v2

Play with our demo: https://huggingface.co/spaces/diarizers-community/DiarizationLM-GGUF

demo

Disclaimer

This is NOT an official Google product.

Instructions

img

Install the package

You can install the package with:

pip install diarizationlm

Once installed, you can directly use many of the existing functions from the package. For example:

import diarizationlm

src_text = "hello good morning hi how are you pretty good"
src_spk = "1 1 1 2 2 2 2 1 1"
tgt_text = "hello morning hi hey are you be good"
tgt_spk = "1 2 2 2 1 1 2 1"
transferred_spk = diarizationlm.transcript_preserving_speaker_transfer(
    src_text, src_spk, tgt_text, tgt_spk)
print(transferred_spk)

Data format

We assume all internal data are stored in JSON files. An example is testdata/example_data.json. The field "utterances" stores a list of utterances, and in each utterance we have these string fields:

Field Description
"utterance_id" This stores the utterance ID.
"hyp_text" This stores the sequence of hypothesis words, but joined by spaces.
"hyp_spk" This stores the sequence of hypothesis speakers, but joined by spaces.
"hyp_diarized_text" This is the text representation of the hypothesis words and speakers. It can be used for debugging and to build the prompts to LLM.
"ref_*" Similar to the "hyp_*" fields, but these are ground truth reference, rather than hypothesis.

Conversion between representations

In the paper, we mentioned two representations:

  1. The word sequence and speaker sequence representation.
  2. The pure text representation.

Example:

Word sequence:         ["good", "morning", "how", "are", "you"]
Speaker sequence:      [1, 1, 2, 2, 2]
Text representation:   "<spk:1> good morning <spk:2> how are you"

We provide the functions in diarizationlm/utils.py to convert between these two representations:

  • create_diarized_text() converts the word and speaker sequences to the pure text representation.
  • extract_text_and_spk() converts the pure text representation to the word and speaker sequences.

Transcript-preserving speaker transfer (TPST)

TPST is a critical data processing algorithm used in multiple places in our paper.

A Python implementation is available in diarizationlm/utils.py, defined as:

def transcript_preserving_speaker_transfer(
    src_text: str, src_spk: str, tgt_text: str, tgt_spk: str
) -> str

img

Training data preparation

We provide a Python script train_data_prep.py that can be used for preparing the dataset for finetuning LLMs (i.e. the prompt builder module described in the paper). This tool will do these for you:

  1. Segment the prompts and completions based on the input and output length limit.
  2. Optionally apply prefix and suffix to prompts and completions.
  3. Store prompt-completion pairs in different file formats.

The segmentation length, prefix, and suffix are passed in as flags to train_data_prep.py. In Python code, they are configured as PromptOptions defined in utils.py.

We support 3 different output file formats:

Format Description
tfrecord The TFRecord format can be used by various machine learning libraries.
json This format is more human readable and can be used for debugging. It's also useful for finetuning PaLM models via the Google Cloud API.
csv This format can be used by many existing tools. OpenAI also provides a tool to convert csv files to jsonl files.
jsonl This format can be directly used by the OpenAI API for finetuning GPT models.

Example command:

python3 train_data_prep.py \
--input="testdata/example_data.json" \
--output="/tmp/example_data.jsonl" \
--output_type=jsonl \
--emit_input_length=1000 \
--emit_target_length=1000 \
--prompt_suffix=" --> " \
--completion_suffix=" [eod]" \
--input_feature_key="prompt" \
--output_feature_key="completion"

LLM finetuning and inference (OpenAI)

Warning: This step is very costly! Proceed with caution at your own risk. Also GPT models are very different from PaLM models. Reproducibility is not guaranteed!

In our paper, we used Google's internal tools to finetune PaLM 2 models and to run the model inference. Google's policy does not allow us to disclose any details about the tools and the PaLM 2 models.

However, if you are interested in reproducing some of our experiments, one option is to use other alternative LLMs, such as OpenAI's GPT models.

Using the train_data_prep.py tool mentioned above, you can create csv files, and use OpenAI libraries to convert to the jsonl format. Example command:

openai tools fine_tunes.prepare_data -f train_data.csv

Once you have the training data in jsonl format, you can finetune GPT models with the data, either via the API or using OpenAI's web UI. For example:

openai api fine_tunes.create -t "train_data.jsonl"

After you have finetuned a model, we provide a Python script run_finetuned_gpt.py to run the GPT model inference on testing data. You need to provide your --api_key and --engine to the script.

LLM finetuning and inference (Llama)

We open sourced Llama 2 & 3 based models on Hugging Face:

The scripts to finetune these models are available in the unsloth folder.

Completion parser

During inference, the prompts are send to the LLM, and the LLM will generate the completions. We provide a postprocess_completions.py script that serves as the completion parser module as described in the paper. It will:

  1. Truncate the completion suffix, and any text generated after this suffix.
  2. Concatenate the completions of all segments from the same utterance.
  3. Transfer the speakers to the original hypothesis ASR transcript.

Metrics

We provide an implementation of these metrics in metrics.py:

Also, we would like to highlight that the WER, WDER, and cpWER metrics reported in our papers are all micro metrics, i.e. both numerators and denominators are aggregated on the entire dataset.

If you use our json-based data format, you can call the compute_metrics_on_json_dict() function as below:

import diarizationlm

json_dict = {
  "utterances": [
      {
          "utterance_id": "utt1",
          "hyp_text": "hello good morning how are you",
          "hyp_spk": "1 1 1 2 2 2",
          "ref_text": "Hello. Good morning, how are you?",
          "ref_spk": "2 2 2 2 1 1",
      },
      {
          "utterance_id": "utt2",
          "hyp_text": "a b c d e f g h",
          "hyp_spk": "1 1 1 2 2 2 3 2",
          "ref_text": "a bb c e f gg g h ii",
          "ref_spk": "2 2 2 2 3 3 4 3 2",
      },
  ]
}
result = diarizationlm.compute_metrics_on_json_dict(json_dict)
print("WER =", result["WER"])
print("WDER =", result["WDER"])
print("cpWER =", result["cpWER"])
print("SpkCntMAE =", result["SpkCntMAE"])

Or you can our script to produce metrics as below:

python3 compute_metrics_on_json.py \
--input=testdata/example_data.json \
--output=/tmp/example_metrics.json

If you use our postprocess_completions.py script to process the LLM results, you need to specify --hyp_spk_field="hyp_spk_llm" when running compute_metrics_on_json.py.

Also please note that this implementation is different from Google's internal implementation that we used in the paper, but is a best-effort attempt to replicate the results. The biggest differences are from text normalization, such as de-punctuation.

Citation

Our paper is cited as:

@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},
}

Metadata

Release files for diarizationlm 0.1.5

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for diarizationlm 0.1.5
File Size Uploaded
diarizationlm-0.1.5.tar.gz 26.5 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for diarizationlm 0.1.5
File Interpreter ABI Platform
diarizationlm-0.1.5-py3-none-any.whl Python 3 none any Details

Total release size: 51.4 kB

Release files / diarizationlm-0.1.5.tar.gz

Download URL diarizationlm-0.1.5.tar.gz
Size 26.5 kB
Tags Source
SHA-256 checksum
How to use checksums
9e201e4795fb9c39c77ddf04bd48dcf969f89af8b5bb4fdf6c1598a3724246be
BLAKE2b-256 checksum
How to use checksums
10f0712739d4b91a50d2062c03632c17c3af84fc6b2112e64897378220320123
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/5.1.1 CPython/3.11.1

Release files / diarizationlm-0.1.5-py3-none-any.whl

Download URL diarizationlm-0.1.5-py3-none-any.whl
Size 24.8 kB
Tags Python 3
SHA-256 checksum
How to use checksums
044f02fce92aed86380d58229231d24c9118b63f5b190dd6c13742eb34cabc40
BLAKE2b-256 checksum
How to use checksums
62a0c07f99699ad677d079d28609d6d6f9f4b50e2c2130af0119f2fc12edef6a
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/5.1.1 CPython/3.11.1

Release history Release notifications | RSS feed

This release

0.1.5 This release

2 release files

0.1.4

2 release files

0.1.3

2 release files

0.1.2

2 release files

0.1.1

2 release files

0.1.0

2 release files

0.0.9

2 release files

0.0.8

2 release files

0.0.7

2 release files

0.0.6

2 release files

0.0.5

2 release files

0.0.4

2 release files

0.0.3

2 release files

0.0.2

2 release files

0.0.1

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page