gipformer-extract / infer.py
giangndm's picture
Upload infer.py with huggingface_hub
f636e7f verified
Raw
History Blame Contribute Delete
3.29 kB
import io
import time
import torch
import torchaudio
import torchaudio.compliance.kaldi as kaldi
import sentencepiece as spm
from huggingface_hub import hf_hub_download
from datasets import load_dataset, Audio
import soundfile as sf
from encoder import PurePyTorchEncoder
from decoder import PurePyTorchDecoder, ModelContainer, greedy_search
def main():
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")
# 1. Instantiate decoupled models from Hugging Face Hub
repo_id = "giangndm/gipformer-extract"
print(f"Loading decoupled Encoder from HF Hub ({repo_id})...")
encoder = PurePyTorchEncoder.from_pretrained(repo_id=repo_id, device=device)
encoder.eval()
print("Encoder loaded successfully.")
print(f"Loading decoupled Decoder & Joiner from HF Hub ({repo_id})...")
decoder_joiner = PurePyTorchDecoder.from_pretrained(repo_id=repo_id, device=device)
decoder_joiner.eval()
print("Decoder & Joiner loaded successfully.")
# Wrap in container for decoder history
model = ModelContainer(encoder, decoder_joiner)
# 2. Download/Load Tokenizer
bpe_model_path = hf_hub_download(repo_id="g-group-ai-lab/gipformer-65M-rnnt", filename="bpe.model")
sp = spm.SentencePieceProcessor()
sp.load(bpe_model_path)
# 3. Load dataset
print("Loading test samples from HuggingFace dataset...")
dataset = load_dataset("luvox-ai/golden-eval-set", split="train")
dataset = dataset.cast_column("audio", Audio(decode=False))
# Test on first 5 samples
for i in range(5):
sample = dataset[i]
ref_text = sample["transcription"]
# Decode audio manually
audio_dict = sample["audio"]
if audio_dict.get("bytes") is not None:
speech, sr = sf.read(io.BytesIO(audio_dict["bytes"]), dtype="float32")
else:
speech, sr = sf.read(audio_dict["path"], dtype="float32")
if speech.ndim > 1:
speech = speech.mean(axis=1)
speech_tensor = torch.from_numpy(speech).float().to(device)
if sr != 16000:
speech_tensor = torchaudio.functional.resample(speech_tensor, sr, 16000)
# Extract features
features = kaldi.fbank(
speech_tensor.unsqueeze(0),
num_mel_bins=80,
frame_shift=10.0,
frame_length=25.0,
dither=0.0,
sample_frequency=16000,
snip_edges=False,
high_freq=-400
).to(device).unsqueeze(0)
feature_lens = torch.tensor([features.size(1)], dtype=torch.int32, device=device)
# Run inference using the split models
with torch.no_grad():
encoder_out, encoder_out_lens = encoder(features, feature_lens)
hyp_tokens = greedy_search(
model=model,
encoder_out=encoder_out,
max_sym_per_frame=1
)
decoded_text = sp.decode(hyp_tokens)
print(f"\nSample {i+1}:")
print(f" Reference: '{ref_text}'")
print(f" Split-Model Hypothesis: '{decoded_text}'")
if __name__ == "__main__":
main()