YAML Metadata Warning:empty or missing yaml metadata in repo card

Check out the documentation for more information.

Gipformer Pure PyTorch Split Model (65M Parameters)

This repository contains a 1000% pure PyTorch extraction and decoupling of the Gipformer-65M RNN-T model (originally from the k2-fsa/icefall Zipformer2 architecture). All external training and speech recognition library dependencies (such as k2, icefall, lhotse, and kaldifeat) have been stripped away.

We have split the pre-trained model into two independent state dictionaries saved in the standard safetensors format:

  1. Encoder (gipformer_encoder.safetensors): Contains the 2D front-end subsampling blocks and the multi-resolution zipformer layers.
  2. Decoder & Joiner (gipformer_decoder.safetensors): Contains the stateless predictor network and the RNN-T joint network.

This enables lightweight, dependency-free speech-to-text inference and makes it trivial to connect the encoder with large language models (LLMs) to build Speech-LLMs.


πŸ“‚ Repository Contents

  • 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.

πŸš€ Quick Start (ASR Inference)

1. Requirements

Ensure you have PyTorch, torchaudio, soundfile, sentencepiece, and safetensors installed:

pip install torch torchaudio soundfile sentencepiece huggingface_hub safetensors

2. Run Transcription

You can run the provided pytorch_split_asr.py script directly. It will automatically download the pre-trained weights and BPE tokenization model, and transcribe sample audio clips:

import torch
import torchaudio
import torchaudio.compliance.kaldi as kaldi
import soundfile as sf
import sentencepiece as spm
from huggingface_hub import hf_hub_download
from safetensors.torch import load_file

# Import architecture classes
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):
    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):
    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

# 2. Load weights & tokenizer
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
bpe_model = hf_hub_download(repo_id="g-group-ai-lab/gipformer-65M-rnnt", filename="bpe.model")
encoder_weights = hf_hub_download(repo_id="giangndm/gipformer-extract", filename="gipformer_encoder.safetensors")
decoder_weights = hf_hub_download(repo_id="giangndm/gipformer-extract", filename="gipformer_decoder.safetensors")

sp = spm.SentencePieceProcessor()
sp.load(bpe_model)

# 3. Instantiate and load split models
encoder = PurePyTorchEncoder(encoder_dims=[192, 256, 384, 512, 384, 256]).to(device).eval()
encoder.load_state_dict(load_file(encoder_weights), strict=True)

decoder_joiner = PurePyTorchDecoder(vocab_size=2000).to(device).eval()
decoder_joiner.load_state_dict(load_file(decoder_weights), strict=True)

model = ModelContainer(encoder, decoder_joiner)

# 4. Load audio and extract Fbank features
speech, sr = sf.read("path_to_audio.wav", dtype="float32")
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
).unsqueeze(0).to(device)

feature_lens = torch.tensor([features.size(1)], dtype=torch.int32, device=device)

# 5. Decoupled Inference forward pass
with torch.no_grad():
    encoder_out, _ = encoder(features, feature_lens)
    hyp_tokens = greedy_search(model=model, encoder_out=encoder_out, max_sym_per_frame=1)
    
text = sp.decode(hyp_tokens)
print("Transcription:", text)

πŸ› οΈ Key Architectural Enhancements

  1. Pure PyTorch Mathematical Activations: The original C++ Swoosh implementations have been replaced with standard, derivative-equivalent PyTorch expressions:
    • $\text{SwooshL}(x) = \text{log}(1 + e^{x-4}) - 0.08x - 0.035$
    • $\text{SwooshR}(x) = \text{log}(1 + e^{x-1}) - 0.08x - 0.31326$
  2. Simplified Layers: Dynamic training layers (Whiten, Balancer, dropouts) have been converted to direct nn.Identity() bypasses, yielding massive speedups at evaluation time.
  3. Standalone Greedy Search: Decodes token outputs entirely in PyTorch, removing the need for k2 FST-based decoders.
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. πŸ™‹ Ask for provider support