Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
171 changes: 171 additions & 0 deletions .github/workflows/export-omnilingual-asr-to-onnx.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,171 @@
name: export-omnilingual-asr-to-onnx

on:
push:
branches:
- export-omnilingual-asr
workflow_dispatch:

concurrency:
group: export-omnilingual-asr-to-onnx-${{ github.ref }}
cancel-in-progress: true

jobs:
export-omnilingual-asr-to-onnx:
if: github.repository_owner == 'k2-fsa' || github.repository_owner == 'csukuangfj'
name: export omnilingual-asr
runs-on: ${{ matrix.os }}
strategy:
fail-fast: false
matrix:
os: [ubuntu-latest]
python-version: ["3.10"]

steps:
- uses: actions/checkout@v4

- name: Setup Python ${{ matrix.python-version }}
uses: actions/setup-python@v5
with:
python-version: ${{ matrix.python-version }}

- name: Install dependencies
shell: bash
run: |
sudo apt install libsndfile1

- name: Install Python dependencies
shell: bash
run: |
pip install fairseq2 \
--extra-index-url https://fair.pkg.atmeta.com/fairseq2/whl/pt2.8.0/cpu \
torch==2.8.0+cpu -f https://download.pytorch.org/whl/torch \
torchaudio==2.8.0+cpu -f https://download.pytorch.org/whl/torchaudio \
onnx==1.17.0 \
onnxruntime==1.17.1 \
soundfile \
librosa

pip install --no-deps omnilingual_asr

pip install retrying pandas polars pyarrow xxhash

- name: Setup tmate session
if: false
uses: mxschmitt/action-tmate@v3

- name: Run
shell: bash
run: |
cd scripts/omnilingual-asr
python3 ./export-onnx.py

ls -lh *.onnx

rm README.md

curl -SL -O https://raw.githubusercontent.com/facebookresearch/omnilingual-asr/refs/heads/main/README.md
curl -SL -O https://raw.githubusercontent.com/facebookresearch/omnilingual-asr/refs/heads/main/LICENSE
curl -SL -O https://raw.githubusercontent.com/facebookresearch/omnilingual-asr/refs/heads/main/LICENSE-CC-BY-4.0.md

curl -SL -O https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/en.wav
curl -SL -O https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/es.wav
curl -SL -O https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/fr.wav
curl -SL -O https://github.com/k2-fsa/sherpa-onnx/releases/download/asr-models/de.wav

echo "---test----"
python3 ./test.py

echo "---collect files----"

d=sherpa-onnx-omnilingual-asr-1600-languages-300M-ctc-2025-11-12

mkdir -p $d
mkdir -p $d/test_wavs

mv -v model.onnx $d
cp -v tokens.txt $d
cp -v README.md $d
cp -v LICENSE* $d
cp -v *.wav $d/test_wavs

ls -lh $d

tar cjfv $d.tar.bz2 $d
mv $d ../..

d=sherpa-onnx-omnilingual-asr-1600-languages-300M-ctc-int8-2025-11-12

mkdir -p $d
mkdir -p $d/test_wavs

mv -v model.int8.onnx $d
cp -v tokens.txt $d
cp -v README.md $d
cp -v LICENSE* $d
cp -v *.wav $d/test_wavs
ls -lh $d

tar cjfv $d.tar.bz2 $d

mv $d ../..

mv *.tar.bz2 ../../

cd ../..

ls -lh *.tar.bz2

- name: Publish to huggingface
env:
HF_TOKEN: ${{ secrets.HF_TOKEN }}
uses: nick-fields/retry@v3
with:
max_attempts: 20
timeout_seconds: 200
shell: bash
command: |
git config --global user.email "csukuangfj@gmail.com"
git config --global user.name "Fangjun Kuang"

export GIT_LFS_SKIP_SMUDGE=1
export GIT_CLONE_PROTECTION_ACTIVE=false

dirs=(
sherpa-onnx-omnilingual-asr-1600-languages-300M-ctc-2025-11-12
sherpa-onnx-omnilingual-asr-1600-languages-300M-ctc-int8-2025-11-12
)

for d in ${dirs[@]}; do
rm -rf huggingface
git clone https://csukuangfj:$HF_TOKEN@huggingface.co/csukuangfj/$d huggingface
pushd huggingface

git fetch
git pull
echo "pwd: $PWD"
cp -a ../$d/* .

git lfs track "*.onnx"
git lfs track "*.wav"
ls -lh
git add .

ls -lh

git status

git commit -m "add models"
git push https://csukuangfj:$HF_TOKEN@huggingface.co/csukuangfj/$d main || true
popd
done

- name: Release
uses: svenstaro/upload-release-action@v2
with:
file_glob: true
file: ./*.tar.bz2
overwrite: true
repo_name: k2-fsa/sherpa-onnx
repo_token: ${{ secrets.UPLOAD_GH_SHERPA_ONNX_TOKEN }}
tag: asr-models
23 changes: 23 additions & 0 deletions scripts/omnilingual-asr/README.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,23 @@
# Introduction

This folder contains script to export
https://github.com/facebookresearch/omnilingual-asr
to sherpa-onnx

See
https://github.com/k2-fsa/sherpa-onnx/blob/master/.github/workflows/export-omnilingual-asr-to-onnx.yaml
for usage.

```
num_frames = round(num_samples / 318 - 1.5)
num_samples = round(318 * num_frames + 477)

or
num_frames = round(num_samples / 320)

```

20ms per frame
Comment on lines +11 to +20

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

The formulas provided for num_frames and num_samples are a bit confusing. The formula num_frames = round(num_samples / 320) is consistent with a 20ms frame duration at a 16kHz sampling rate (since 16000 * 0.020 = 320). However, the other set of formulas using 318 and 477 seems to contradict this. Could you please clarify the relationship between these different formulas and explain when each should be used? Providing context on how these numbers are derived from the model architecture would be very helpful for users.




103 changes: 103 additions & 0 deletions scripts/omnilingual-asr/export-onnx.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,103 @@
#!/usr/bin/env python3
# Copyright 2025 Xiaomi Corp. (authors: Fangjun Kuang)

from typing import Dict

import onnx
import torch
from fairseq2.nn.batch_layout import BatchLayout
from omnilingual_asr.models.inference.pipeline import ASRInferencePipeline
from onnxruntime.quantization import QuantType, quantize_dynamic


def add_meta_data(filename: str, meta_data: Dict[str, str]):
"""Add meta data to an ONNX model. It is changed in-place.

Args:
filename:
Filename of the ONNX model to be changed.
meta_data:
Key-value pairs.
"""
model = onnx.load(filename)
while len(model.metadata_props):
model.metadata_props.pop()
Comment on lines +23 to +24

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

This while loop to clear the metadata properties is inefficient. The RepeatedCompositeContainer returned by model.metadata_props has a clear() method that is more efficient and idiomatic for this purpose.

    model.metadata_props.clear()


for key, value in meta_data.items():
meta = model.metadata_props.add()
meta.key = key
meta.value = str(value)


class ModelWrapper(torch.nn.Module):
def __init__(self, model):
super().__init__()
self.model = model

def forward(self, x):
"""
Args:
x: (N, num_samples), float32
"""
batch_layout = BatchLayout(shape=x.shape, seq_lens=[x.shape[1]])

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

The seq_lens argument for BatchLayout is incorrectly constructed for batch sizes greater than 1. It is currently [x.shape[1]], which means it's a list with a single element. This will fail if the batch size x.shape[0] is greater than 1. To support batching correctly, it should be a list of sequence lengths for each item in the batch.

        batch_layout = BatchLayout(shape=x.shape, seq_lens=[x.shape[1]] * x.shape[0])

logits, _ = self.model(x, batch_layout)
return logits


@torch.no_grad()
def main():
pipeline = ASRInferencePipeline(
model_card="omniASR_CTC_300M",
device="cpu",
dtype=torch.float32,
)

vocab_size = pipeline.tokenizer._model.vocabulary_size

with open("tokens.txt", "w") as f:
for i in range(pipeline.tokenizer._model.vocabulary_size):
f.write(f"{pipeline.tokenizer._model.index_to_token(i)} {i}\n")
Comment on lines +55 to +59

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

Accessing the protected member _model of pipeline.tokenizer is fragile and can lead to issues if the omnilingual-asr library is updated. If there is a public API to get the vocabulary size and map indices to tokens, it would be much safer to use that. If not, it would be good to add a comment here acknowledging the risk.


print("saved to tokens.txt")

wrapper = ModelWrapper(pipeline.model)
wrapper.eval()

x = torch.rand(1, 16000 * 10)
torch.onnx.export(
wrapper,
x,
"model.onnx",
opset_version=14,
input_names=["x"],
output_names=["logits"],
dynamic_axes={
"x": {0: "N", 1: "num_samples"},
"logits": {0: "N", 1: "num_frames"},
},
)

meta_data = {
"vocab_size": vocab_size,
"model_type": "omnilingual-asr",
"version": "1",
"sample_rate": 16000,
"model_author": "facebookresearch",
"url": "https://github.com/facebookresearch/omnilingual-asr",
"comment": "300M-CTC",
}

add_meta_data("model.onnx", meta_data)
print("saved to model.onnx")

quantize_dynamic(
model_input="./model.onnx",
model_output="./model.int8.onnx",
op_types_to_quantize=["MatMul"],
weight_type=QuantType.QUInt8,
)
print("saved to model.int8.onnx")


if __name__ == "__main__":
main()
Loading
Loading