gipformer-extract / pytorch_split_asr.py
giangndm's picture
Upload pytorch_split_asr.py with huggingface_hub
0149dd9 verified
Raw
History Blame Contribute Delete
6.76 kB
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()