Views
No views yet
k2-fsa/icefall Zipformer2 architecture). All external training and speech recognition library dependencies (such as k2, icefall, lhotse, and kaldifeat) have been stripped away.safetensors format:gipformer_encoder.safetensors): Contains the 2D front-end subsampling blocks and the multi-resolution zipformer layers.gipformer_decoder.safetensors): Contains the stateless predictor network and the RNN-T joint network.gipformer_pure_pytorch.py: Self-contained, dependency-free architecture code.gipformer_encoder.safetensors: Pre-trained weights for the subsampler and Zipformer encoder (FP32).gipformer_decoder.safetensors: Pre-trained weights for the stateless predictor and the Joint network (FP32).pytorch_split_asr.py: Minimal end-to-end ASR inference script demonstrating how to run transcription using the split components.pip install torch torchaudio soundfile sentencepiece huggingface_hub safetensorspytorch_split_asr.py script directly. It will automatically download the pre-trained weights and BPE tokenization model, and transcribe sample audio clips:1import torch
2import torchaudio
3import torchaudio.compliance.kaldi as kaldi
4import soundfile as sf
5import sentencepiece as spm
6from huggingface_hub import hf_hub_download
7from safetensors.torch import load_file
8
9# Import architecture classes
10from gipformer_pure_pytorch import (
11 Conv2dSubsampling,
12 Zipformer2,
13 Decoder,
14 Joiner,
15 greedy_search
16)
17
18# 1. We define decoupled wrappers for Encoder and Decoder/Joiner parts
19class PurePyTorchEncoder(torch.nn.Module):
20 def __init__(self, encoder_dims, in_channels=80):
21 super().__init__()
22 self.encoder_embed = Conv2dSubsampling(
23 in_channels=in_channels,
24 out_channels=encoder_dims[0],
25 dropout=0.0
26 )
27 self.encoder = Zipformer2(
28 output_downsampling_factor=2,
29 downsampling_factor=[1, 2, 4, 8, 4, 2],
30 num_encoder_layers=[2, 2, 3, 4, 3, 2],
31 encoder_dim=encoder_dims,
32 encoder_unmasked_dim=[192, 192, 256, 256, 256, 192],
33 query_head_dim=[32],
34 pos_head_dim=[4],
35 value_head_dim=[12],
36 pos_dim=48,
37 num_heads=[4, 4, 4, 8, 4, 4],
38 feedforward_dim=[512, 768, 1024, 1536, 1024, 768],
39 cnn_module_kernel=[31, 31, 15, 15, 15, 31],
40 dropout=0.0,
41 warmup_batches=1.0,
42 causal=False
43 )
44
45 def forward(self, x: torch.Tensor, x_lens: torch.Tensor):
46 x, x_lens = self.encoder_embed(x, x_lens)
47 batch_size = x_lens.size(0)
48 max_len = x.shape[1]
49 seq_range = torch.arange(0, max_len, device=x.device)
50 seq_range_expand = seq_range.unsqueeze(0).expand(batch_size, max_len)
51 seq_length_expand = x_lens.unsqueeze(-1).expand(batch_size, max_len)
52 src_key_padding_mask = seq_range_expand >= seq_length_expand
53
54 x = x.permute(1, 0, 2)
55 encoder_out, encoder_out_lens = self.encoder(x, x_lens, src_key_padding_mask)
56 encoder_out = encoder_out.permute(1, 0, 2)
57 return encoder_out, encoder_out_lens
58
59class PurePyTorchDecoder(torch.nn.Module):
60 def __init__(self, vocab_size=2000, decoder_dim=512, joiner_dim=512):
61 super().__init__()
62 self.decoder = Decoder(
63 vocab_size=vocab_size,
64 decoder_dim=decoder_dim,
65 blank_id=0,
66 context_size=2
67 )
68 self.joiner = Joiner(
69 encoder_dim=decoder_dim,
70 decoder_dim=decoder_dim,
71 joiner_dim=joiner_dim,
72 vocab_size=vocab_size
73 )
74
75class ModelContainer(torch.nn.Module):
76 def __init__(self, encoder, decoder_joiner):
77 super().__init__()
78 self.encoder = encoder
79 self.decoder = decoder_joiner.decoder
80 self.joiner = decoder_joiner.joiner
81
82# 2. Load weights & tokenizer
83device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
84bpe_model = hf_hub_download(repo_id="g-group-ai-lab/gipformer-65M-rnnt", filename="bpe.model")
85encoder_weights = hf_hub_download(repo_id="giangndm/gipformer-extract", filename="gipformer_encoder.safetensors")
86decoder_weights = hf_hub_download(repo_id="giangndm/gipformer-extract", filename="gipformer_decoder.safetensors")
87
88sp = spm.SentencePieceProcessor()
89sp.load(bpe_model)
90
91# 3. Instantiate and load split models
92encoder = PurePyTorchEncoder(encoder_dims=[192, 256, 384, 512, 384, 256]).to(device).eval()
93encoder.load_state_dict(load_file(encoder_weights), strict=True)
94
95decoder_joiner = PurePyTorchDecoder(vocab_size=2000).to(device).eval()
96decoder_joiner.load_state_dict(load_file(decoder_weights), strict=True)
97
98model = ModelContainer(encoder, decoder_joiner)
99
100# 4. Load audio and extract Fbank features
101speech, sr = sf.read("path_to_audio.wav", dtype="float32")
102speech_tensor = torch.from_numpy(speech).float().to(device)
103if sr != 16000:
104 speech_tensor = torchaudio.functional.resample(speech_tensor, sr, 16000)
105
106features = kaldi.fbank(
107 speech_tensor.unsqueeze(0),
108 num_mel_bins=80,
109 frame_shift=10.0,
110 frame_length=25.0,
111 dither=0.0,
112 sample_frequency=16000,
113 snip_edges=False,
114 high_freq=-400
115).unsqueeze(0).to(device)
116
117feature_lens = torch.tensor([features.size(1)], dtype=torch.int32, device=device)
118
119# 5. Decoupled Inference forward pass
120with torch.no_grad():
121 encoder_out, _ = encoder(features, feature_lens)
122 hyp_tokens = greedy_search(model=model, encoder_out=encoder_out, max_sym_per_frame=1)
123
124text = sp.decode(hyp_tokens)
125print("Transcription:", text)Whiten, Balancer, dropouts) have been converted to direct nn.Identity() bypasses, yielding massive speedups at evaluation time.k2 FST-based decoders.