| 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}") |
| |
| |
| 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.") |
| |
| |
| model = ModelContainer(encoder, decoder_joiner) |
| |
| |
| 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) |
| |
| |
| 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)) |
| |
| |
| for i in range(5): |
| sample = dataset[i] |
| ref_text = sample["transcription"] |
| |
| |
| 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) |
| |
| |
| 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) |
| |
| |
| 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() |
|
|