File size: 6,756 Bytes
6986d0e 0149dd9 6986d0e 0149dd9 6986d0e 0149dd9 6986d0e 0149dd9 6986d0e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 | 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()
|