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 safetensors.torch import load_file # Import our pure PyTorch Gipformer implementation from gipformer_pure_pytorch import ( Conv2dSubsampling, Zipformer2, Decoder, Joiner, greedy_search ) # 1. We define decoupled wrappers for Encoder and Decoder/Joiner parts class PurePyTorchEncoder(torch.nn.Module): """ Decoupled Encoder that contains the front-end subsampler (encoder_embed) and the main Zipformer encoder. """ def __init__(self, encoder_dims, in_channels=80): super().__init__() self.encoder_embed = Conv2dSubsampling( in_channels=in_channels, out_channels=encoder_dims[0], dropout=0.0 ) self.encoder = Zipformer2( output_downsampling_factor=2, downsampling_factor=[1, 2, 4, 8, 4, 2], num_encoder_layers=[2, 2, 3, 4, 3, 2], encoder_dim=encoder_dims, encoder_unmasked_dim=[192, 192, 256, 256, 256, 192], query_head_dim=[32], pos_head_dim=[4], value_head_dim=[12], pos_dim=48, num_heads=[4, 4, 4, 8, 4, 4], feedforward_dim=[512, 768, 1024, 1536, 1024, 768], cnn_module_kernel=[31, 31, 15, 15, 15, 31], dropout=0.0, warmup_batches=1.0, causal=False ) def forward(self, x: torch.Tensor, x_lens: torch.Tensor): x, x_lens = self.encoder_embed(x, x_lens) # Create padding mask dynamically batch_size = x_lens.size(0) max_len = x.shape[1] seq_range = torch.arange(0, max_len, device=x.device) seq_range_expand = seq_range.unsqueeze(0).expand(batch_size, max_len) seq_length_expand = x_lens.unsqueeze(-1).expand(batch_size, max_len) src_key_padding_mask = seq_range_expand >= seq_length_expand x = x.permute(1, 0, 2) # (N, T, C) -> (T, N, C) encoder_out, encoder_out_lens = self.encoder(x, x_lens, src_key_padding_mask) encoder_out = encoder_out.permute(1, 0, 2) # (T, N, C) -> (N, T, C) return encoder_out, encoder_out_lens class PurePyTorchDecoder(torch.nn.Module): """ Decoupled Decoder that contains the stateless predictor (decoder) and the joint network (joiner). """ def __init__(self, vocab_size=2000, decoder_dim=512, joiner_dim=512): super().__init__() self.decoder = Decoder( vocab_size=vocab_size, decoder_dim=decoder_dim, blank_id=0, context_size=2 ) self.joiner = Joiner( encoder_dim=decoder_dim, decoder_dim=decoder_dim, joiner_dim=joiner_dim, vocab_size=vocab_size ) # A fake model container to allow reuse of the standard greedy_search function class ModelContainer(torch.nn.Module): def __init__(self, encoder, decoder_joiner): super().__init__() self.encoder = encoder self.decoder = decoder_joiner.decoder self.joiner = decoder_joiner.joiner def main(): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Using device: {device}") # 1. Paths to Split Weights encoder_weights_path = "researchs/asr-benchmark/code/checkpoints/gipformer_encoder.safetensors" decoder_weights_path = "researchs/asr-benchmark/code/checkpoints/gipformer_decoder.safetensors" bpe_model_path = hf_hub_download(repo_id="g-group-ai-lab/gipformer-65M-rnnt", filename="bpe.model") # 2. Instantiate decoupled models and load their respective weights print("Loading decoupled Encoder...") encoder = PurePyTorchEncoder(encoder_dims=[192, 256, 384, 512, 384, 256]) encoder_state = load_file(encoder_weights_path) encoder.load_state_dict(encoder_state, strict=True) encoder.to(device).eval() print("Encoder loaded successfully.") print("Loading decoupled Decoder & Joiner...") decoder_joiner = PurePyTorchDecoder(vocab_size=2000) decoder_state = load_file(decoder_weights_path) decoder_joiner.load_state_dict(decoder_state, strict=True) decoder_joiner.to(device).eval() print("Decoder & Joiner loaded successfully.") # Wrap in container for decoder history model = ModelContainer(encoder, decoder_joiner) # 3. Load Tokenizer sp = spm.SentencePieceProcessor() sp.load(bpe_model_path) # 4. 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(): # Run the split encoder encoder_out, encoder_out_lens = encoder(features, feature_lens) # Run the greedy decoder (accesses split decoder and joiner) hyp_tokens = greedy_search( model=model, encoder_out=encoder_out, max_sym_per_frame=1 ) # Decode BPE tokens to text 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()