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)