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()