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:
- Keep the first 64,000 samples, corresponding to at most four seconds.
- Record the retained sample count as
input_length. - 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.