FAD-MoE / app.py
Vansh Chugh
fix: drop z-score norm, ckpt never applied it
c0306c8
Raw
History Blame Contribute Delete
5.23 kB
import sys
sys.stdout.reconfigure(line_buffering=True)
try:
import spaces
except ImportError:
# keep @spaces.GPU usable as a no-op; ZeroGPU requires this exact name.
class spaces:
class GPU:
def __init__(self, func=None, duration=60):
self.func = func
def __call__(self, *args, **kwargs):
if self.func is not None:
return self.func(*args, **kwargs)
func = args[0]
return func
import os
import threading
import gradio as gr
import torch
from pyharp import ModelCard, build_endpoint, load_audio
from models.moe_research.w2v2_moe_fz24_aasist import Model
from utils.tools.tools import pad
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
CKPT_PATH = os.path.join(os.path.dirname(__file__), "checkpoints", "fz24_moe_aasist.ckpt")
TRUNCATE = 64600 # ~4s at 16kHz (default: T=64600 samples, per paper)
model = None
model_ready = False # has the device-resident model been built yet?
model_loading = True
model_error = None
def build_model():
"""Builds the detector fresh: frozen wav2vec2 backbone from the HF Hub plus
the trained mixture-of-experts fusion and AASIST head from the checkpoint.
The checkpoint only ever stores the MoE/AASIST weights, never the backbone,
since the backbone is frozen and untouched during training."""
detector = Model()
state_dict = torch.load(CKPT_PATH, map_location="cpu", weights_only=False)["state_dict"]
state_dict = {k[len("model."):]: v for k, v in state_dict.items() if k.startswith("model.")}
detector.load_state_dict(state_dict, strict=False)
detector.eval()
return detector
def warm_cache():
"""Downloads the wav2vec2 backbone and does a dry-run build on CPU, so a
bad checkpoint or a Hub outage surfaces before the first real request
instead of mid-inference. CPU only -- ZeroGPU only intercepts CUDA calls
made inside an @spaces.GPU call, not from a background thread."""
global model_loading, model_error
try:
build_model()
print("Backbone + checkpoint warmed.")
except Exception as e:
model_error = str(e)
print(f"Load error: {e}")
finally:
model_loading = False
threading.Thread(target=warm_cache, daemon=True).start()
model_card = ModelCard(
name="FAD-MoE",
description=(
"Detects AI-generated (spoofed) speech. A frozen wav2vec 2.0 backbone's "
"24 layers are fused by a sparse mixture-of-experts, then classified by "
"an AASIST graph-attention head. Only the first ~4 seconds of the clip "
"are analyzed, per the paper."
),
author=(
"Zhiyong Wang, Ruibo Fu, Zhengqi Wen, Jianhua Tao, Xiaopeng Wang, "
"Yuankun Xie, Xin Qi, Shuchen Shi, Yi Lu, Yukun Liu, Chenxing Li, Xuefei Liu"
),
tags=["deepfake-detection", "speech", "anti-spoofing"],
)
def preprocess(sig):
"""Pad/crop to ~4s, raw waveform, no z-score norm -- matches
asvspoof_data_DA.py, the data module this checkpoint's hparams.yaml
actually names. Mono/16kHz resample added since HARP users upload
arbitrary files, unlike the original's pre-formatted 16kHz mono
ASVspoof set."""
sig = sig.to_mono().resample(16000)
wav = sig.audio_data[0, 0].cpu().numpy()
wav = pad(wav, TRUNCATE)
return torch.tensor(wav, dtype=torch.float32)
@spaces.GPU
@torch.inference_mode()
def process_fn(input_audio_path: str) -> str:
"""Runs the detector on one clip and writes a bonafide/spoofed verdict with
confidence scores to a text file."""
global model, model_ready
if model_loading:
raise gr.Error("Model is still loading, please wait a moment and try again.")
if model_error is not None:
raise gr.Error(f"Model failed to load: {model_error}")
if not model_ready:
model = build_model().to(DEVICE)
model_ready = True
sig = load_audio(input_audio_path)
wav = preprocess(sig).unsqueeze(0).to(DEVICE)
pred, _, _ = model(wav)
probs = torch.nn.functional.softmax(pred, dim=1)[0]
bonafide_prob = probs[1].item()
spoof_prob = probs[0].item()
verdict = "Likely genuine" if bonafide_prob >= spoof_prob else "Likely AI-generated / spoofed"
input_name = os.path.basename(input_audio_path)
out_path = os.path.splitext(input_audio_path)[0] + "_result.txt"
with open(out_path, "w") as f:
f.write(f"{input_name}\n\n")
f.write("AI-Detection Result\n")
f.write(
f"{verdict} (genuine confidence: {bonafide_prob:.1%}, "
f"spoofed confidence: {spoof_prob:.1%})\n"
)
return out_path
with gr.Blocks() as demo:
input_components = [
gr.Audio(type="filepath", label="Input Audio").harp_required(True),
]
output_components = [
gr.File(type="filepath", label="Detection Result").set_info(
"Bonafide/spoofed verdict with confidence scores."
),
]
build_endpoint(
model_card=model_card,
input_components=input_components,
output_components=output_components,
process_fn=process_fn,
)
if __name__ == "__main__":
demo.queue().launch(pwa=True)