WST-Graph: Speech Deepfake Detection

This repository provides TorchScript inference models for WST-Graph: Topology-Preserving Wavelet Scattering Front-End for Speech Deepfake Detection, together with our reproduced AASIST and AASIST-L baselines.

Each file contains the complete waveform-to-logits computation graph and model parameters. Inference does not require the original Python model implementation, Kymatio.

When the code is prepared, we will plase to Github

Model overview

The released WST configurations use:

Configuration Q81 Q82
Averaging scale J 8 8
Wavelets per octave Q (8, 1) (8, 2)
Oversampling 1 1
Retained scattering orders First and second First and second
ALAP temporal bins 64 64
Normalization Log-modulation Log-modulation

Available models

All model families are available with training seeds 9, 114, and 514.

Model family Directory example CUDA filename example
WST-Graph-Q81 q81-seed9/ wst_graph_q81_seed9_batch32_cuda.ts
WST-Graph-Q82 q82-seed9/ wst_graph_q82_seed9_batch32_cuda.ts
Reproduced AASIST rp-aasist-seed9/ reproduce_aasist_seed9_batch32_cuda.ts
Reproduced AASIST-L rp-aasist-l-seed9/ reproduce_aasist_l_seed9_batch32_cuda.ts

CPU filenames omit _cuda, for example wst_graph_q81_seed9_batch32.ts.

Use CPU files on CPU and CUDA files on logical device cuda:0. To select a different physical GPU, set CUDA_VISIBLE_DEVICES before starting Python.

Input and output

Audio must be mono, 16 kHz.

For each recording:

  1. Keep the first 64,000 samples, corresponding to at most four seconds.
  2. Record the retained sample count as input_length.
  3. Right-pad shorter recordings with zeros to 64,000 samples.
Tensor Shape Type
input_values [B, 64000] torch.float32
input_length [B] torch.int64
Output logits [B, 2] torch.float32

Class labels are:

  • 0: bona fide
  • 1: spoof

Quick start

The following example automatically selects the matching CPU or CUDA:

import soundfile as sf
import torch
from huggingface_hub import hf_hub_download

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

variant = "q81"
seed = 9
suffix = "_cuda" if device.type == "cuda" else ""

filename = (
    f"{variant}-seed{seed}/"
    f"wst_graph_{variant}_seed{seed}_batch32{suffix}.ts"
)

model_path = hf_hub_download(
    repo_id="kwokho1/wst-graph",
    filename=filename,
)

model = torch.jit.load(
    model_path,
    map_location=device,
).eval()

# Replace these paths with up to 32 recordings.
audio_paths = ["audio.wav"]

if not 1 <= len(audio_paths) <= 32:
    raise ValueError("Provide between 1 and 32 recordings.")

waveforms = []
lengths = []

for path in audio_paths:
    audio, sample_rate = sf.read(
        path,
        dtype="float32",
        always_2d=True,
    )

    if sample_rate != 16000:
        raise ValueError("Resample the recording to 16 kHz first.")

    # Convert multiple channels to mono.
    audio = audio.mean(axis=1)
    length = min(len(audio), 64000)

    if length == 0:
        raise ValueError(f"Empty recording: {path}")

    waveform = torch.zeros(64000, dtype=torch.float32)
    waveform[:length] = torch.from_numpy(audio[:length])

    waveforms.append(waveform)
    lengths.append(length)

input_values = torch.stack(waveforms).to(device)
input_length = torch.tensor(
    lengths,
    dtype=torch.int64,
    device=device,
)

with torch.inference_mode():
    logits = model(input_values, input_length)

bonafide_score = logits[:, 0] - logits[:, 1]
spoof_score = logits[:, 1] - logits[:, 0]

print("Logits:", logits.cpu())
print("Bona fide scores:", bonafide_score.cpu())

These files are loaded with torch.jit.load.

Acknowledgments

The implementation builds on the wavelet scattering transform provided by Kymatio and the graph architecture introduced in AASIST.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Paper for kwokho1/wst-graph