| 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 |
|
|
| |
| from gipformer_pure_pytorch import ( |
| Conv2dSubsampling, |
| Zipformer2, |
| Decoder, |
| Joiner, |
| greedy_search |
| ) |
|
|
| |
| 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) |
| |
| 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) |
| encoder_out, encoder_out_lens = self.encoder(x, x_lens, src_key_padding_mask) |
| encoder_out = encoder_out.permute(1, 0, 2) |
| 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 |
| ) |
|
|
| |
| 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}") |
| |
| |
| 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") |
| |
| |
| 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.") |
| |
| |
| model = ModelContainer(encoder, decoder_joiner) |
| |
| |
| 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() |
|
|