ashirbadsahu/shrutam-2-onnx-gguf
020
Shrutam-2 ONNX & GGUF Inference Guide
This repository provides optimized formats for Shrutam-2:
shrutam_encoder.onnx: Standalone speech encoder combining Log-Mel Spectrogram preprocessing, 17-layer Conformer Encoder, Conv1D Downsampler, and SMEAR MoE Projector into a single ONNX computational graph.shrutam_llm_f16.gguf/shrutam_llm_q4_k_m.gguf: GGUF quantized models for the LLM component, compatible withllama.cpp.
Model Architecture Overview
- Input Audio: 16 kHz Mono Audio Waveform
[batch_size, num_samples] - Audio Encoder (`shrutam_encoder.onnx`):
- MelSpectrogram Preprocessor (Log-Mel extraction)
- Conformer Encoder (1024 d_model)
- 1D Convolution Downsampler
- SMEAR MoE Router and Expert Projector
- Output: Projected embeddings
[batch_size, feature_length, 2048] - *LLM Decoder (`shrutam_llm_.gguf`)**:
- LlamaForCausalLM (2048 hidden size, 16 layers, 128k vocabulary)
Requirements
pip install onnxruntime torchaudio torch llama-cpp-pythonPython Usage Example
Step 1: Run the Audio Encoder (ONNX)
import torch
import torchaudio
import onnxruntime as ort
# 1. Configure ONNX Runtime Session
opts = ort.SessionOptions()
opts.intra_op_num_threads = 4
opts.log_severity_level = 3 # Suppress internal warnings
session = ort.InferenceSession("shrutam_encoder.onnx", sess_options=opts, providers=["CPUExecutionProvider"])
# 2. Load 16kHz Audio Waveform
wav, sr = torchaudio.load("audio.wav")
if sr != 16000:
resampler = torchaudio.transforms.Resample(orig_freq=sr, new_freq=16000)
wav = resampler(wav)
if wav.dim() == 1:
wav = wav.unsqueeze(0)
elif wav.shape[0] > 1:
wav = wav.mean(dim=0, keepdim=True)
# 3. Extract Audio Embeddings via ONNX
audio_embeds = session.run(None, {"audio": wav.numpy()})[0]
print("Audio Embeddings Shape:", audio_embeds.shape) # Shape: (1, seq_len, 2048)Step 2: Pass Audio Embeddings into GGUF LLM (llama.cpp)
from llama_cpp import Llama
import numpy as np
# Load GGUF LLM Model
llm = Llama(
model_path="shrutam_llm_q4_k_m.gguf",
n_ctx=4096,
n_threads=4,
verbose=False
)
# Prefix prompt formatting for Shrutam-2
prompt_text = "<|im_start|>user\nTranscribe speech to Hindi text.<|im_end|>\n<|im_start|>assistant\n"
prompt_tokens = llm.tokenize(prompt_text.encode("utf-8"), add_bos=True)
# Note: Combine audio_embeds with text prompt token embeddings using llama.cpp input embedding API.
print("Prompt Tokens Count:", len(prompt_tokens))File Details
Technical Features
- Dynamic Audio Length Support: Accepts variable length 16kHz audio inputs dynamically.
- Hardware Agnostic: Runs seamlessly on CPU (
CPUExecutionProvider) or GPU (CUDAExecutionProvider). - No PyTorch dependency required for inference: Pure ONNX Runtime + C++ / GGUF runner integration.
