Views
No views yet
1import torch
2import torchaudio
3from transformers import Wav2Vec2FeatureExtractor
4# model = load your checkpoint wrapped with the BiGRU+Attention classifier
5
6# 1) Load audio and resample to 16k
7waveform, sr = torchaudio.load("example.wav")
8if sr != 16000:
9 waveform = torchaudio.functional.resample(waveform, sr, 16000)
10
11# 2) Ensure 4 seconds length (pad or truncate)
12target_len = 4 * 16000
13if waveform.shape[1] < target_len:
14 pad = target_len - waveform.shape[1]
15 waveform = torch.nn.functional.pad(waveform, (0, pad))
16else:
17 waveform = waveform[:, :target_len]
18
19# 3) Feature extraction (wav2vec2)
20feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained("facebook/wav2vec2-base")
21input_values = feature_extractor(waveform.squeeze(0).numpy(), sampling_rate=16000, return_tensors="pt").input_values
22
23# 4) Forward pass through model
24model.eval()
25with torch.no_grad():
26 logits = model(input_values) # model should return a scalar logit per sample
27 prob = torch.sigmoid(logits).item() # probability of 'fake' class
28
29prediction = 1 if prob >= 0.5 else 0
30confidence = prob if prediction == 1 else 1 - prob