From 8f295ec41b216e3596cdc217506b2cfa4dd0235d Mon Sep 17 00:00:00 2001 From: contentis Date: Wed, 23 Sep 2026 08:44:38 +0200 Subject: [PATCH 1/8] POC: add Qwen3 ASR and independent forced alignment --- .gitignore | 5 + CMakeLists.txt | 1 + README.md | 2 + THIRD_PARTY_NOTICES.md | 10 + asr/qwen3/CMakeLists.txt | 17 + asr/qwen3/README.md | 135 +++++ asr/qwen3/aligner_main.cpp | 61 ++ asr/qwen3/detail/runtime.h | 554 +++++++++++++++++++ asr/qwen3/forced_aligner.cpp | 231 ++++++++ asr/qwen3/forced_aligner.h | 44 ++ asr/qwen3/main.cpp | 60 ++ asr/qwen3/model_export/export_qwen3_asr.py | 412 ++++++++++++++ asr/qwen3/model_export/validate_qwen3_asr.py | 391 +++++++++++++ asr/qwen3/qwen3.cpp | 385 +++++++++++++ asr/qwen3/qwen3.h | 64 +++ asr/qwen3/requirements.txt | 13 + cmake/pcre2.cmake | 16 + cmake/utf8proc.cmake | 12 + common/io/CMakeLists.txt | 5 + common/io/tokenizer.cpp | 67 ++- common/io/tokenizer.h | 4 + common/io/unicode_regex.cpp | 51 ++ common/io/unicode_regex.h | 25 + common/ort_session.cpp | 73 ++- common/ort_session.h | 2 + common/progress.h | 26 + 26 files changed, 2624 insertions(+), 42 deletions(-) create mode 100644 asr/qwen3/CMakeLists.txt create mode 100644 asr/qwen3/README.md create mode 100644 asr/qwen3/aligner_main.cpp create mode 100644 asr/qwen3/detail/runtime.h create mode 100644 asr/qwen3/forced_aligner.cpp create mode 100644 asr/qwen3/forced_aligner.h create mode 100644 asr/qwen3/main.cpp create mode 100644 asr/qwen3/model_export/export_qwen3_asr.py create mode 100644 asr/qwen3/model_export/validate_qwen3_asr.py create mode 100644 asr/qwen3/qwen3.cpp create mode 100644 asr/qwen3/qwen3.h create mode 100644 asr/qwen3/requirements.txt create mode 100644 cmake/pcre2.cmake create mode 100644 cmake/utf8proc.cmake create mode 100644 common/io/unicode_regex.cpp create mode 100644 common/io/unicode_regex.h create mode 100644 common/progress.h diff --git a/.gitignore b/.gitignore index b8f79d9..18ca335 100644 --- a/.gitignore +++ b/.gitignore @@ -94,3 +94,8 @@ audio/ TODO frames/ .run/ + +# Local Qwen3 validation and development helpers +asr/qwen3/internal/ +asr/qwen3/model_export/_internal/ +asr/qwen3/tests/ diff --git a/CMakeLists.txt b/CMakeLists.txt index c70acfc..9392711 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -100,6 +100,7 @@ add_subdirectory(image_gen/flux2) add_subdirectory(asr/rnnt) add_subdirectory(vision/sam2) add_subdirectory(asr/whisper) +add_subdirectory(asr/qwen3) if(DIN_BUILD_TESTING) set(_din_cli_test_audio "${DIN_DEPLOY_ARTIFACTS_ROOT}/audio/de/schakaleundaraber_elli.mp3") diff --git a/README.md b/README.md index 4957d58..5f111ff 100644 --- a/README.md +++ b/README.md @@ -19,6 +19,8 @@ DIN Deploy is a collection of practical samples for exporting and running local | OpenAI Whisper | `openai/whisper-tiny`
`openai/whisper-base`
`openai/whisper-small`
`openai/whisper-medium`
`openai/whisper-large-v3`
`openai/whisper-large-v3-turbo` | [Whisper](asr/whisper/README.md) | | NVIDIA Parakeet TDT 0.6B v3 | `nvidia/parakeet-tdt-0.6b-v3` | [RNNT](asr/rnnt/README.md) | | NVIDIA Nemotron 3.5 ASR Streaming 0.6B | `nvidia/nemotron-3.5-asr-streaming-0.6b` | [RNNT](asr/rnnt/README.md) | +| Qwen3 ASR | `Qwen/Qwen3-ASR-0.6B-hf`
`Qwen/Qwen3-ASR-1.7B-hf` | [Qwen3](asr/qwen3/README.md) | +| Qwen3 Forced Aligner | `Qwen/Qwen3-ForcedAligner-0.6B-hf` | [Qwen3](asr/qwen3/README.md) | ### Computer Vision diff --git a/THIRD_PARTY_NOTICES.md b/THIRD_PARTY_NOTICES.md index f3ac284..e211c3b 100644 --- a/THIRD_PARTY_NOTICES.md +++ b/THIRD_PARTY_NOTICES.md @@ -8,12 +8,16 @@ DIN Deploy is distributed under the Apache License, Version 2.0. The project sou | miniaudio | MIT: | | lodepng | zlib: | | nlohmann/json | MIT: | +| PCRE2 | BSD-3-Clause WITH PCRE2-exception: | +| utf8proc | MIT and Unicode data license: | | NVIDIA NVTX | Apache-2.0 with LLVM exception: | | ONNX Runtime | MIT: | | ONNX Runtime TensorRT RTX EP ABI | Apache-2.0: | | Slang | Apache-2.0 with LLVM exception: | | Vulkan-Headers and Vulkan-Loader | Apache-2.0: and | | nanobind | BSD-3-Clause: | +| Qwen3-ASR inference utilities | Apache-2.0, copyright 2026 The Alibaba Qwen team: | +| Transformers Qwen3-ASR processing | Apache-2.0, copyright Hugging Face: | | Python dependencies | See the package license links and inventory below. | | Distributed models and model artifacts | See the model license inventory below. Model terms may differ from this project's license. | @@ -28,6 +32,8 @@ For components supplied through an SDK or binary package, the corresponding vend - `lodepng`: zlib - `nanobind` v2.9.2: BSD-3-Clause - `nlohmann/json` v3.11.3: MIT +- `PCRE2` 10.46: BSD-3-Clause WITH PCRE2-exception +- `utf8proc` 2.11.0: MIT and Unicode data license - NVIDIA NVTX v3.5.0 C/C++: Apache-2.0 WITH LLVM-exception - ONNX Runtime SDK 1.27.0: MIT - ONNX Runtime TensorRT RTX Execution Provider ABI: Apache-2.0 @@ -90,3 +96,7 @@ For components supplied through an SDK or binary package, the corresponding vend - `openai/whisper-large-v3-turbo`: Apache-2.0 - `openai/whisper-medium`: Apache-2.0 - `openai/whisper-small`: Apache-2.0 + +- `Qwen/Qwen3-ASR-0.6B-hf`: Apache-2.0 +- `Qwen/Qwen3-ASR-1.7B-hf`: Apache-2.0 +- `Qwen/Qwen3-ForcedAligner-0.6B-hf`: Apache-2.0 diff --git a/asr/qwen3/CMakeLists.txt b/asr/qwen3/CMakeLists.txt new file mode 100644 index 0000000..3426310 --- /dev/null +++ b/asr/qwen3/CMakeLists.txt @@ -0,0 +1,17 @@ +# Independent ASR and forced-alignment APIs, sharing the model runtime. +add_din_shared_library(din_asr_qwen3 STATIC qwen3.cpp forced_aligner.cpp) +target_include_directories(din_asr_qwen3 PUBLIC "${CMAKE_CURRENT_SOURCE_DIR}") +target_link_libraries(din_asr_qwen3 PUBLIC din_common_ort din_common_io nlohmann_json::nlohmann_json) +find_package(CUDAToolkit QUIET) +if(CUDAToolkit_FOUND) + target_link_libraries(din_asr_qwen3 PRIVATE CUDA::cudart) + target_compile_definitions(din_asr_qwen3 PRIVATE DIN_QWEN3_CUDA=1) +endif() +add_din_executable(din_asr_qwen3_cli main.cpp) +target_link_libraries(din_asr_qwen3_cli PRIVATE din_asr_qwen3 argparse) +add_din_executable(din_asr_qwen3_aligner_cli aligner_main.cpp) +target_link_libraries(din_asr_qwen3_aligner_cli PRIVATE din_asr_qwen3 argparse) + +if(BUILD_TESTING AND EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/tests/CMakeLists.txt") + add_subdirectory(tests) +endif() diff --git a/asr/qwen3/README.md b/asr/qwen3/README.md new file mode 100644 index 0000000..48342fe --- /dev/null +++ b/asr/qwen3/README.md @@ -0,0 +1,135 @@ +# Qwen3 ASR and forced alignment + +Offline C++ inference with CPU or TensorRT RTX. Use FP32 exports for CPU; TensorRT RTX supports BF16 (default), FP16 and FP32. + +## Supported models + +| Model | Hugging Face ID | Checkpoint | BF16 export | FP16 export | FP32 export | Recommended export | +| --- | --- | --- | --- | --- | --- | --- | +| ASR 0.6B | `Qwen/Qwen3-ASR-0.6B-hf` | BF16 | ✓ | ✓ | ✓ | BF16 | +| ASR 1.7B | `Qwen/Qwen3-ASR-1.7B-hf` | BF16 | ✓ | ✓ | ✓ | BF16 | +| Forced Aligner 0.6B | `Qwen/Qwen3-ForcedAligner-0.6B-hf` | BF16 | ✓ | ✓ | ✓ | BF16 | + +## Supported capabilities + +| Model / upstream toolkit capability | C++ sample | +|---|:---:| +| Offline, single stream | ✓ | +| Online / streaming | — | +| Batched inference | — | +| Long-form audio | ✓ | +| ASR with / without forced alignment | ✓ | +| Standalone alignment of supplied text | ✓ | +| Automatic language identification / language hint | ✓ | +| Multilingual ASR: 30 languages and 22 Chinese dialects | ✓ | +| Word timestamps: en, de, es, fr, it, pt, ru, ko | ✓ | +| Chinese / Cantonese character timestamps | ✓ | +| Japanese character timestamps | ✓ | +| All 11 upstream alignment languages | ✓ | + +ASR uses the [upstream model's language support](https://github.com/QwenLM/Qwen3-ASR). +ASR accepts all 30 upstream language codes/names and `auto`; dialects use automatic +recognition or the corresponding language hint, not separate dialect switches. +Alignment supports Chinese, Cantonese, English, German, Spanish, French, Italian, +Portuguese, Russian, Korean and Japanese. Japanese uses character timestamps; +upstream's Nagisa word boundaries differ. Latin words in CJK text stay together. +Language names/codes are case-insensitive. ASR's other languages require alignment +to be disabled. Language coverage is not an accuracy guarantee for every dialect. + +## Export + +Run Python commands from `asr/qwen3/model_export`, with the dependencies in +`../requirements.txt` and a CUDA-enabled PyTorch installation. + +```bash +python -X utf8 export_qwen3_asr.py --size 0.6B --output D:/models/qwen3-asr-0.6b-onnx-bf16 +python -X utf8 export_qwen3_asr.py --size 1.7B --output D:/models/qwen3-asr-1.7b-onnx-bf16 +python -X utf8 export_qwen3_asr.py --task aligner --output D:/models/qwen3-aligner-onnx-bf16 +``` + +Use `--dtype fp16` or `--dtype fp32` for converted exports; `--dtype original` +keeps BF16. This applies to both ASR sizes and the aligner. + +```bash +python -X utf8 export_qwen3_asr.py --dtype fp16 --output D:/models/qwen3-asr-0.6b-onnx-fp16 +python -X utf8 export_qwen3_asr.py --dtype fp32 --output D:/models/qwen3-asr-0.6b-onnx-fp32 +``` + +The C++ pipeline reads precision from each export; ASR and aligner can use different +precisions. FP32 uses decomposed attention for TensorRT RTX compatibility; BF16/FP16 +use fused attention. Log-mel stays FP32. Use separate output directories per precision. + +HF downloads checkpoints automatically. Use `--model` for a local checkpoint or +`--revision` to pin the source. Keep each export directory intact. Log-mel +processing reuses the shared Whisper frontend. + +Exports contain one encoder and one decoder, each with one weight file, plus the +shared log-mel graph. Prefill and token generation update one KV bank in place. +The decoder uses two fixed TensorRT profiles (512-token prefill and one-token +steps), compiled once and cached; audio length does not create more encoder/decoder +profiles. The two engines may each retain weights in GPU memory. Changed exports +get new cache keys; clear compiled caches when changing the GPU or runtime. + +`--cache-capacity` sets the token ceiling (default 8192; multiples of 512 up to +16384). Long audio uses upstream quiet-boundary splitting with enough room reserved +for `--max-new-tokens`. Reaching that generation limit returns `reached_eos=false`. +Attention still scans the allocated cache. BF16/FP16 KV uses 896 MiB at 8192 slots. +Re-export older ASR artifacts for format 3. + +## Verify + +```bash +python -X utf8 validate_qwen3_asr.py --onnx-dir D:/models/qwen3-asr-0.6b-onnx-bf16 --audio audio.mp3 +python -X utf8 validate_qwen3_asr.py --onnx-dir D:/models/qwen3-asr-1.7b-onnx-bf16 --audio audio.mp3 +python -X utf8 validate_qwen3_asr.py --task aligner --onnx-dir D:/models/qwen3-aligner-onnx-bf16 --audio audio.mp3 --transcript transcript.txt --language Chinese +``` + +The validator compares encoder, prefill and cached-token outputs against HF, +then uses HF `generate()` for a short end-to-end token/EOS check. Alignment uses +HF transcript preparation and span decoding. Strict BF16 numerical comparisons +can fail despite matching tokens/spans. Long-form specialized decoding can change +words; long-form alignment has small endpoint differences from HF. + +## Build + +Build from the repository root with the same TensorRT RTX setup as Whisper. +An NVIDIA GPU supporting the selected TensorRT RTX precision is required. + +```powershell +cmake --build out\build\windows-x64 --target din_asr_qwen3_cli din_asr_qwen3_aligner_cli +``` + +## Run + +```powershell +out\build\windows-x64\bin\din_asr_qwen3_cli.exe audio.mp3 --model-dir D:\models\qwen3-asr-1.7b-onnx-bf16 +out\build\windows-x64\bin\din_asr_qwen3_aligner_cli.exe audio.mp3 --model-dir D:\models\qwen3-aligner-onnx-bf16 --transcript transcript.txt --lang-id zh +``` + +Multi-configuration builds add the configuration (for example, `Release`) under `bin`. +ASR and alignment are independent APIs in the same library: + +| API / CLI | Input | Output | +|---|---|---| +| `Qwen3Pipeline` / `din_asr_qwen3_cli` | Audio | Text, tokens, language and audio chunk boundaries | +| `Qwen3ForcedAligner` / `din_asr_qwen3_aligner_cli` | Audio, supplied text and language | Word/character timestamps | + +Include `qwen3.h` for ASR or `forced_aligner.h` for alignment. The aligner loads +no ASR model; use text from Whisper, Parakeet, Nemotron, Qwen ASR or a text file. +Both CLIs accept `--provider cpu|trt-rtx`, `--model-dir` and cache options independently. +Reuse instances; each processes one synchronous call at a time. Both configs expose +`progress` callbacks for loading, compilation and processing. + +`Align` takes mono 16 kHz audio; `AlignFile` decodes it automatically. Alignment +accepts at most 180 seconds per call. Long recordings require matching audio/text +segments, with returned timestamps offset by each segment's start. + +ASR returns `segments` with text, detected language and half-open `start_sample` / +`end_sample` offsets at 16 kHz. Its upstream quiet-boundary splitter uses a +1200-second target and a ±5-second search, further limited by KV capacity. +`--max-chunk-seconds` overrides the target (6–1200; 0 uses the default). +For subsequent alignment, use a target of 175 seconds or less to reserve the +search margin within the 180-second limit. Chunk boundaries are not word timestamps. + +`--max-new-tokens` defaults to 1024 per chunk. If `reached_eos` is false, increase +the budget or reduce `--max-chunk-seconds`; the transcript is incomplete. diff --git a/asr/qwen3/aligner_main.cpp b/asr/qwen3/aligner_main.cpp new file mode 100644 index 0000000..1130e79 --- /dev/null +++ b/asr/qwen3/aligner_main.cpp @@ -0,0 +1,61 @@ +// SPDX-License-Identifier: Apache-2.0 +#include +#include +#include +#include +#include + +#include "forced_aligner.h" +#include +#include + +int main(int argc, char** argv) +{ + using namespace din::asr::qwen3; + try + { + ForcedAlignerConfig config; + argparse::ArgumentParser parser("din_asr_qwen3_aligner_cli"); + parser.add_description("Align supplied text with audio, independently of any ASR model (up to 180 seconds)."); + parser.add_argument("audiofile"); + parser.add_argument("--transcript").required().help("UTF-8 transcript file"); + parser.add_argument("--lang-id").default_value(std::string{"English"}).help("Language code or name"); + parser.add_argument("--provider").default_value(config.provider).choices("cpu", "trt-rtx"); + parser.add_argument("--model-dir").default_value(config.model_dir.string()); + parser.add_argument("--ep-cache").default_value(config.ep_cache_dir.string()); + parser.add_argument("--ep-context-dir").default_value(config.ep_context_dir.string()); + parser.parse_args(argc, argv); + config.provider = parser.get("--provider"); + config.model_dir = parser.get("--model-dir"); + config.ep_cache_dir = parser.get("--ep-cache"); + config.ep_context_dir = parser.get("--ep-context-dir"); + const auto transcript_path = parser.get("--transcript"); + std::ifstream file(transcript_path, std::ios::binary); + if (!file) + throw std::runtime_error("Cannot read transcript: " + transcript_path); + std::string transcript{std::istreambuf_iterator(file), std::istreambuf_iterator()}; + if (transcript.starts_with("\xef\xbb\xbf")) + transcript.erase(0, 3); + const auto language = parser.get("--lang-id"); + const auto audio = din::io::LoadAudio(parser.get("audiofile"), 16000); + Qwen3ForcedAligner aligner(std::move(config)); + const auto start = std::chrono::steady_clock::now(); + const auto timestamps = aligner.Align(audio, transcript, language); + const auto seconds = std::chrono::duration(std::chrono::steady_clock::now() - start).count(); + nlohmann::json output{{"text", transcript}, + {"language", language}, + {"audio_seconds", audio.Duration()}, + {"align_seconds", seconds}, + {"timestamps", nlohmann::json::array()}}; + for (const auto& word : timestamps) + output["timestamps"].push_back( + {{"text", word.text}, {"start_time", word.start_time}, {"end_time", word.end_time}}); + std::cout << output.dump(2) << '\n'; + return 0; + } + catch (const std::exception& error) + { + std::cerr << "Qwen3 aligner: " << error.what() << '\n'; + return 1; + } +} diff --git a/asr/qwen3/detail/runtime.h b/asr/qwen3/detail/runtime.h new file mode 100644 index 0000000..51058f3 --- /dev/null +++ b/asr/qwen3/detail/runtime.h @@ -0,0 +1,554 @@ +// SPDX-License-Identifier: Apache-2.0 +#pragma once +#ifdef DIN_QWEN3_CUDA +#include +#endif + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "audio.h" +#include "nvtx_helper.h" +#include "ort_session.h" +#include "tokenizer.h" +#include + +namespace din::asr::qwen3::detail +{ +using Json = nlohmann::json; +using BF16 = Ort::BFloat16_t; +using din::common::OrtRunner; +template +using Buffer = din::common::TensorBuffer; +class FloatBuffer +{ + using Storage = std::variant, Buffer, Buffer>; + Storage buffer_; + + static Storage Make(OrtRunner& runner, const std::vector& shape, ONNXTensorElementDataType dtype, + bool disable_uma) + { + switch (dtype) + { + case ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT: + return Storage(std::in_place_type>, runner, shape, runner.HasDeviceIo(), disable_uma); + case ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16: + return Storage(std::in_place_type>, runner, shape, runner.HasDeviceIo(), + disable_uma); + case ONNX_TENSOR_ELEMENT_DATA_TYPE_BFLOAT16: + return Storage(std::in_place_type>, runner, shape, runner.HasDeviceIo(), disable_uma); + default: + throw std::runtime_error("Unsupported Qwen3 tensor precision"); + } + } + +public: + FloatBuffer(OrtRunner& runner, const std::vector& shape, ONNXTensorElementDataType dtype, + bool disable_uma = false) + : buffer_(Make(runner, shape, dtype, disable_uma)) + { + } + + template + void WithHost(F&& fill) + { + std::visit( + [&](auto& buffer) + { + fill(buffer.HostData()); + }, + buffer_); + } + void Fill(float value) + { + std::visit( + [&](auto& buffer) + { + using T = std::remove_pointer_t; + buffer.Fill(T(value)); + }, + buffer_); + } + Ort::Value& BindingValue() + { + return std::visit( + [](auto& buffer) -> Ort::Value& + { + return buffer.BindingValue(); + }, + buffer_); + } + void CopyAsyncToDevice() + { + std::visit( + [](auto& buffer) + { + buffer.CopyAsyncToDevice(); + }, + buffer_); + } + void UploadAndWait() + { + std::visit( + [](auto& buffer) + { + buffer.CopyAsyncToDeviceWithNotification().Sync(); + }, + buffer_); + } +}; + +inline size_t ElementBytes(const Ort::Value& value) +{ + return value.GetTensorTypeAndShapeInfo().GetElementType() == ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT ? 4 : 2; +} + +constexpr int kRate = 16000; + +inline int64_t AudioTokens(int64_t frames) +{ + return frames / 100 * 13 + ((frames % 100) + 7) / 8; +} + +#ifdef DIN_QWEN3_CUDA +inline bool UsesCudaMemory(const OrtRunner& runner) +{ + if (!runner.HasDeviceIo()) + return false; + const auto hardware = runner.ep_device.Device(); + return hardware.Type() == OrtHardwareDeviceType_GPU && hardware.VendorId() == 0x10DE; +} +#endif + +inline void ZeroTensor(OrtRunner& runner, Ort::Value& value) +{ + const auto bytes = value.GetTensorTypeAndShapeInfo().GetElementCount() * ElementBytes(value); + if (!runner.HasDeviceIo()) + { + std::memset(value.GetTensorMutableRawData(), 0, bytes); + return; + } +#ifdef DIN_QWEN3_CUDA + if (UsesCudaMemory(runner)) + { + const auto status = cudaMemsetAsync(value.GetTensorMutableRawData(), 0, bytes, + reinterpret_cast(runner.compute_stream->GetHandle())); + if (status != cudaSuccess) + throw std::runtime_error(cudaGetErrorString(status)); + return; + } +#endif + std::vector zeros(bytes, 0); + const auto info = value.GetTensorTypeAndShapeInfo(); + const auto shape = info.GetShape(); + const auto memory = Ort::MemoryInfo::CreateCpu(OrtArenaAllocator, OrtMemTypeDefault); + auto source = + Ort::Value::CreateTensor(memory, zeros.data(), bytes, shape.data(), shape.size(), info.GetElementType()); + Ort::ThrowOnError(runner.env.CopyTensor(source, value, *runner.compute_stream)); + din::common::NotificationPtr done(*runner.compute_stream); + done.Record(); + done.Sync(); +} + +inline Json ReadJson(const std::filesystem::path& path) +{ + std::ifstream file(path); + if (!file) + throw std::runtime_error("Missing model asset: " + path.string() + "; re-export the model"); + return Json::parse(file); +} + +inline void Append(std::vector& dst, const std::vector& src) +{ + dst.insert(dst.end(), src.begin(), src.end()); +} + +inline Ort::Value TensorValue(OrtRunner& runner, const std::vector& shape, ONNXTensorElementDataType dtype) +{ + if (runner.HasDeviceIo()) + return Ort::Value::CreateTensor(runner.DeviceAllocator(), shape.data(), shape.size(), dtype); + Ort::AllocatorWithDefaultOptions allocator; + return Ort::Value::CreateTensor(allocator, shape.data(), shape.size(), dtype); +} + +struct EncodedAudio +{ + int64_t tokens; + Ort::Value values; +}; + +inline void CopyTensor(OrtRunner& runner, Ort::Value& dst, const Ort::Value& src, size_t count, size_t dst_offset = 0, + size_t src_offset = 0) +{ + const auto bytes = ElementBytes(dst); + auto* target = static_cast(dst.GetTensorMutableRawData()) + dst_offset * bytes; + const auto* source = static_cast(src.GetTensorRawData()) + src_offset * bytes; + if (!runner.HasDeviceIo()) + { + std::memcpy(target, source, count * bytes); + return; + } +#ifdef DIN_QWEN3_CUDA + if (UsesCudaMemory(runner)) + { + const auto status = cudaMemcpyAsync(target, source, count * bytes, cudaMemcpyDeviceToDevice, + reinterpret_cast(runner.compute_stream->GetHandle())); + if (status != cudaSuccess) + throw std::runtime_error(cudaGetErrorString(status)); + return; + } +#endif + const int64_t shape = static_cast(count); + const auto dtype = dst.GetTensorTypeAndShapeInfo().GetElementType(); + auto from = + Ort::Value::CreateTensor(src.GetTensorMemoryInfo(), const_cast(source), count * bytes, &shape, 1, dtype); + auto to = Ort::Value::CreateTensor(dst.GetTensorMemoryInfo(), target, count * bytes, &shape, 1, dtype); + Ort::ThrowOnError(runner.env.CopyTensor(from, to, *runner.compute_stream)); +} + +inline const din::io::Audio& NormalizeAudio(const din::io::Audio& audio, din::io::Audio& normalized) +{ + if (audio.sample_rate != kRate || audio.samples.empty()) + throw std::runtime_error("Expected nonempty mono 16 kHz audio"); + float peak = 0.f; + for (float sample : audio.samples) + { + if (!std::isfinite(sample)) + throw std::runtime_error("Audio contains non-finite samples"); + peak = std::max(peak, std::abs(sample)); + } + const auto* source = &audio; + if (peak > 1.f) + { + normalized.sample_rate = kRate; + normalized.samples.resize(audio.samples.size()); + std::transform(audio.samples.begin(), audio.samples.end(), normalized.samples.begin(), + [peak](float x) + { + return x / peak; + }); + source = &normalized; + } + return *source; +} + +struct TextInputs +{ + Buffer ids, positions, logits_index; + Ort::Value audio; + FloatBuffer bias; + Buffer mask; + int64_t sequence, hidden, capacity; + OrtRunner& runner; + bool mask_uploaded = false; + + TextInputs(OrtRunner& runner, int64_t seq, int64_t width, int64_t keys, ONNXTensorElementDataType dtype) + : ids(runner, {1, seq}, runner.HasDeviceIo()) + , positions(runner, {1, seq}, runner.HasDeviceIo()) + , logits_index(runner, {1}, runner.HasDeviceIo()) + , audio(TensorValue(runner, {1, seq, width}, dtype)) + , bias(runner, {1, 1, seq, keys}, dtype) + , mask(runner, {1, seq, 1}, runner.HasDeviceIo()) + , sequence(seq) + , hidden(width) + , capacity(keys) + , runner(runner) + { + ZeroTensor(runner, audio); + mask.Fill(false); + } + + void Fill(std::span tokens, int64_t start, int64_t audio_id, const EncodedAudio* embeddings, + int64_t audio_offset = 0) + { + if (tokens.empty() || tokens.size() > static_cast(sequence) || start + sequence > capacity) + throw std::runtime_error("Input exceeds the exported context capacity"); + int64_t audio_pos = audio_offset; + bool mask_changed = false; + for (int64_t i = 0; i < sequence; ++i) + { + ids.HostData()[i] = i < static_cast(tokens.size()) ? tokens[i] : 0; + positions.HostData()[i] = start + i; + const bool is_audio = embeddings && i < static_cast(tokens.size()) && tokens[i] == audio_id; + mask_changed |= mask.HostData()[i] != is_audio; + mask.HostData()[i] = is_audio; + if (is_audio) + { + if (audio_pos >= embeddings->tokens) + throw std::runtime_error("Too many audio placeholders"); + ++audio_pos; + } + } + bias.WithHost( + [&](auto* data) + { + using T = std::remove_pointer_t; + for (int64_t i = 0; i < sequence; ++i) + { + const int64_t visible = start + i + 1; + std::fill_n(data + i * capacity, visible, T(0.f)); + std::fill_n(data + i * capacity + visible, capacity - visible, T(-1e4f)); + } + }); + logits_index.HostData()[0] = tokens.size() - 1; + logits_index.CopyAsyncToDevice(); + ids.CopyAsyncToDevice(); + positions.CopyAsyncToDevice(); + // Audio placeholders form contiguous runs; assemble embeddings directly on + // the shared stream instead of downloading each encoder window to the CPU. + audio_pos = audio_offset; + for (int64_t i = 0; embeddings && i < static_cast(tokens.size());) + { + if (tokens[i] != audio_id) + { + ++i; + continue; + } + const int64_t begin = i; + while (i < static_cast(tokens.size()) && tokens[i] == audio_id) + ++i; + CopyTensor(runner, audio, embeddings->values, (i - begin) * hidden, begin * hidden, audio_pos * hidden); + audio_pos += i - begin; + } + if (!mask_uploaded || mask_changed) + { + mask.CopyAsyncToDevice(); + mask_uploaded = true; + } + bias.CopyAsyncToDevice(); + } + + void Bind(Ort::IoBinding& binding, bool decoder = false) + { + if (decoder) + binding.BindInput("logits_index", logits_index.BindingValue()); + binding.BindInput("input_ids", ids.BindingValue()); + binding.BindInput("position_ids", positions.BindingValue()); + binding.BindInput("audio_embeddings", audio); + binding.BindInput("audio_mask", mask.BindingValue()); + binding.BindInput("attention_bias", bias.BindingValue()); + } +}; +inline std::string LowerLanguage(const std::string& language) +{ + const auto first = language.find_first_not_of(" \r\n\t"); + auto result = first == std::string::npos ? std::string{} + : language.substr(first, language.find_last_not_of(" \r\n\t") - first + 1); + std::transform(result.begin(), result.end(), result.begin(), + [](unsigned char c) + { + return static_cast(std::tolower(c)); + }); + return result; +} + +struct Runtime +{ + std::string provider; + std::filesystem::path ep_cache_dir, ep_context_dir; + din::common::ProgressCallback progress; + Ort::Env env{ORT_LOGGING_LEVEL_WARNING, "din_asr_qwen3"}; + Ort::SyncStream stream{nullptr}; + std::unique_ptr mel; + + Runtime(const std::string& execution_provider, const std::filesystem::path& cache, + const std::filesystem::path& context, din::common::ProgressCallback callback) + : provider(execution_provider) + , ep_cache_dir(cache) + , ep_context_dir(context) + , progress(std::move(callback)) + { + if (provider == "trt-rtx") + { + din::common::RegisterTensorRTRTXProvider(env); + stream = din::common::CreateTensorRTRTXComputeStream(env); + } + } + void LoadMel(const std::filesystem::path& dir) + { + mel = Runner(dir, "mel", "samples:1x8000", "samples:1x2880000", "samples:1x19280000"); + } + std::unique_ptr Runner(const std::filesystem::path& dir, const std::string& name, const std::string& min, + const std::string& opt, const std::string& max, const std::string& variant = "") + { + din::common::ModelProfile profile; + profile.min_shapes = min; + profile.opt_shapes = opt; + profile.max_shapes = max; + const auto metadata = ReadJson(dir / "metadata.json"); + const auto identity = metadata.contains("graphs") && metadata["graphs"].contains(name) + ? metadata["graphs"][name].get().substr(0, 16) + : dir.filename().string(); + profile.cache_subpath = "qwen3_" + name + "_" + identity + "_fixed" + variant; + if (name == "decoder") + profile.cache_subpath += "_kv" + metadata["cache_capacity"].dump(); + if (name == "aligner") + profile.cache_subpath += "_bins"; + profile.enable_cuda_graph = name != "aligner"; + if (!profile.enable_cuda_graph) + { + // Alignment runs once per chunk with changing shapes. Keep graphs + // for repeated encoder windows and fixed-shape AR steps. + profile.extra_ep_options.emplace_back("enable_cuda_graph", "0"); + profile.cache_subpath += "_no_graph"; + } + profile.embed_ep_context = false; + // ORT names external engines by graph, so different profiles need separate directories. + return std::make_unique( + env, (dir / (name + ".onnx")).string(), provider, ep_cache_dir.string(), + din::common::EpContextOptions{(ep_context_dir / profile.cache_subpath).string(), progress}, profile, + stream ? &stream : nullptr); + } + + std::vector Features(std::span audio) + { + din::common::nvtx_scoped_range range{"qwen3.mel"}; + const int64_t samples = std::max(8000, audio.size()); + Buffer input(*mel, {1, samples}, mel->HasDeviceIo()), + output(*mel, {1, 128, samples / 160}, mel->HasDeviceIo()); + input.Fill(0.f); + std::copy(audio.begin(), audio.end(), input.HostData()); + input.CopyAsyncToDevice(); + Ort::IoBinding binding(mel->session); + binding.BindInput("samples", input.BindingValue()); + binding.BindOutput("features", output.BindingValue()); + mel->session.Run(Ort::RunOptions{}, binding); + output.CopyAsyncToHostWithNotification().Sync(); + return {output.HostData(), output.HostData() + 128 * (samples / 160)}; + } +}; +struct AudioModel +{ + Json metadata, native; + ONNXTensorElementDataType dtype = ONNX_TENSOR_ELEMENT_DATA_TYPE_BFLOAT16; + std::unique_ptr tokenizer; + std::unique_ptr encoder; + int64_t hidden = 0, audio_id = 0, window = 0; + + struct EncoderBuffers + { + int64_t frames, tokens; + FloatBuffer mel, bias; + Buffer indices; + Ort::Value output; + Ort::IoBinding binding; + + EncoderBuffers(OrtRunner& runner, int64_t count, int64_t hidden, ONNXTensorElementDataType dtype) + : frames(count) + , tokens(AudioTokens(count)) + , mel(runner, {(count + 99) / 100, 128, 100}, dtype, true) + , bias(runner, {1, 1, tokens, tokens}, dtype, true) + , indices(runner, {tokens}, runner.HasDeviceIo()) + , output(TensorValue(runner, {tokens, hidden}, dtype)) + , binding(runner.session) + { + // A call contains exactly one independent HF encoder window. + bias.Fill(0.f); + for (int64_t i = 0; i < tokens; ++i) + indices.HostData()[i] = i; + bias.CopyAsyncToDevice(); + indices.CopyAsyncToDevice(); + binding.BindInput("mel_chunks", mel.BindingValue()); + binding.BindInput("valid_indices", indices.BindingValue()); + binding.BindInput("attention_bias", bias.BindingValue()); + binding.BindOutput("audio_embeddings", output); + } + }; + + // Pad the last encoder window so all audio lengths reuse one GPU shape. + std::unique_ptr full_window; + + EncodedAudio Encode(const std::vector& features, int64_t frames) + { + const int64_t tokens = AudioTokens(frames); + EncodedAudio result{tokens, TensorValue(*encoder, {tokens, hidden}, dtype)}; + Ort::RunOptions options; + if (encoder->HasDeviceIo()) + options.AddConfigEntry("disable_synchronize_execution_providers", "1"); + int64_t token_offset = 0; + for (int64_t offset = 0; offset < frames; offset += window) + { + din::common::nvtx_scoped_range range{"qwen3.encoder_window"}; + const auto count = std::min(window, frames - offset); + if (!full_window) + full_window = std::make_unique(*encoder, window, hidden, dtype); + auto& input = *full_window; + const auto valid_tokens = AudioTokens(count); + input.bias.WithHost( + [&](auto* data) + { + using T = std::remove_pointer_t; + for (int64_t i = 0; i < input.tokens; ++i) + { + std::fill_n(data + i * input.tokens, valid_tokens, T(0.f)); + std::fill_n(data + i * input.tokens + valid_tokens, input.tokens - valid_tokens, T(-1e4f)); + } + }); + input.bias.CopyAsyncToDevice(); + input.mel.Fill(0.f); + input.mel.WithHost( + [&](auto* data) + { + using T = std::remove_pointer_t; + for (int64_t c = 0; c < (count + 99) / 100; ++c) + for (int64_t m = 0; m < 128; ++m) + for (int64_t f = 0; f < 100 && c * 100 + f < count; ++f) + data[(c * 128 + m) * 100 + f] = T(features[m * frames + offset + c * 100 + f]); + }); + // Complete only the upload before reusing host staging. The next + // window's CPU preparation can overlap this window's inference. + input.mel.UploadAndWait(); + encoder->session.Run(options, input.binding); + CopyTensor(*encoder, result.values, input.output, valid_tokens * hidden, token_offset * hidden); + token_offset += valid_tokens; + } + return result; + } + AudioModel(Runtime& runtime, const std::filesystem::path& dir, const std::string& task) + { + metadata = ReadJson(dir / "metadata.json"); + native = ReadJson(dir / "native.json"); + const auto& meta = metadata; + if (meta.at("format_version") != (task == "asr" ? 3 : 2) || meta.at("task") != task) + throw std::runtime_error("Re-export Qwen3 for the unified in-place decoder"); + const auto precision = meta.at("dtype").get(); + if (precision == "bfloat16") + dtype = ONNX_TENSOR_ELEMENT_DATA_TYPE_BFLOAT16; + else if (precision == "float16") + dtype = ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16; + else if (precision == "float32") + dtype = ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT; + else + throw std::runtime_error("Unsupported Qwen3 export precision: " + precision); + if (meta["audio_config"]["num_mel_bins"] != 128 || meta["audio_config"]["n_window"] != 50 || + meta["audio_config"]["n_window_infer"] != 800) + throw std::runtime_error("Unsupported audio geometry"); + window = meta["audio_config"]["n_window_infer"]; + hidden = meta["text_config"]["hidden_size"]; + audio_id = meta["audio_token_id"]; + tokenizer = std::make_unique((dir / "processor/tokenizer.json").string(), + din::io::TokenizerFormat::ByteBpeJson); + auto enc_shapes = [](int chunks, int tokens) + { + return "mel_chunks:" + std::to_string(chunks) + "x128x100,valid_indices:" + std::to_string(tokens) + + ",attention_bias:1x1x" + std::to_string(tokens) + "x" + std::to_string(tokens); + }; + encoder = runtime.Runner(dir, "encoder", enc_shapes(8, 104), enc_shapes(8, 104), enc_shapes(8, 104)); + } +}; +inline std::string TextShape(int64_t seq, int64_t hidden, int64_t keys) +{ + const auto s = std::to_string(seq); + return "input_ids:1x" + s + ",audio_embeddings:1x" + s + "x" + std::to_string(hidden) + ",audio_mask:1x" + s + + "x1,position_ids:1x" + s + ",attention_bias:1x1x" + s + "x" + std::to_string(keys); +} + +} // namespace din::asr::qwen3::detail diff --git a/asr/qwen3/forced_aligner.cpp b/asr/qwen3/forced_aligner.cpp new file mode 100644 index 0000000..bad8cb7 --- /dev/null +++ b/asr/qwen3/forced_aligner.cpp @@ -0,0 +1,231 @@ +// SPDX-License-Identifier: Apache-2.0 +#include "forced_aligner.h" + +#include "detail/runtime.h" +#include "unicode_regex.h" + +namespace din::asr::qwen3::detail +{ +std::vector AlignmentUnits(const std::string& text, const std::string& language) +{ + static constexpr std::pair languages[] = { + {"zh", "chinese"}, {"yue", "cantonese"}, {"en", "english"}, {"de", "german"}, + {"es", "spanish"}, {"fr", "french"}, {"it", "italian"}, {"pt", "portuguese"}, + {"ru", "russian"}, {"ko", "korean"}, {"ja", "japanese"}}; + const auto name = LowerLanguage(language); + const auto found = std::find_if(std::begin(languages), std::end(languages), + [&](const auto& entry) + { + return name == entry.first || name == entry.second; + }); + if (found == std::end(languages)) + throw std::invalid_argument("Forced alignment supports zh, yue, en, de, es, fr, it, pt, ru, ko and ja"); + + // HF keeps Unicode letters/numbers and ASCII apostrophes, dropping punctuation and marks. + static const din::io::UnicodeRegex kept(R"([\p{L}\p{N}'\s\x{1c}-\x{1f}]+)"); + std::string cleaned; + for (const auto& part : kept.FindAll(text)) + cleaned += part; + static const std::string cjk = R"(\x{4e00}-\x{9fff}\x{3400}-\x{4dbf}\x{20000}-\x{2a6df}\x{2a700}-\x{2b73f})" + R"(\x{2b740}-\x{2b81f}\x{2b820}-\x{2ceaf}\x{f900}-\x{faff}\x{2f800}-\x{2fa1f})"; + static const din::io::UnicodeRegex words("[" + cjk + "]|[^" + cjk + R"(\s\x{1c}-\x{1f}]+)"); + // HF's unscored Korean LTokenizer splits on whitespace before cleaning each unit. + static const din::io::UnicodeRegex korean(R"([^\s\x{1c}-\x{1f}]+)"); + // Japanese character timestamps avoid a separate Nagisa word-segmentation runtime. + static const std::string kana = R"(\x{3040}-\x{30ff}\x{31f0}-\x{31ff}\x{ff66}-\x{ff9f})"; + static const din::io::UnicodeRegex japanese("[" + cjk + kana + "]|[^" + cjk + kana + R"(\s\x{1c}-\x{1f}]+)"); + if (found->first == "ja") + return japanese.FindAll(cleaned); + if (found->first == "ko") + return korean.FindAll(cleaned); + return words.FindAll(cleaned); +} + +// HF's nondecreasing subsequence repair, including its tie rules. +std::vector FixTimestamps(const std::vector& data) +{ + const int n = static_cast(data.size()); + if (n == 0) + return {}; + std::vector dp(n, 1), parent(n, -1), result = data; + std::vector normal(n, false); + for (int i = 1; i < n; ++i) + for (int j = 0; j < i; ++j) + if (data[j] <= data[i] && dp[j] + 1 > dp[i]) + { + dp[i] = dp[j] + 1; + parent[i] = j; + } + int i = static_cast(std::max_element(dp.begin(), dp.end()) - dp.begin()); + for (; i >= 0; i = parent[i]) + normal[i] = true; + for (int begin = 0; begin < n;) + { + if (normal[begin]) + { + ++begin; + continue; + } + int end = begin; + while (end < n && !normal[end]) + ++end; + for (int k = begin; k < end; ++k) + { + if (begin == 0) + result[k] = result[end]; + else if (end == n) + result[k] = result[begin - 1]; + else if (end - begin <= 2) + result[k] = k - begin + 1 <= end - k ? result[begin - 1] : result[end]; + else + result[k] = static_cast(result[begin - 1] + (result[end] - result[begin - 1]) * + float(k - begin + 1) / (end - begin + 1)); + } + begin = end; + } + return result; +} + +struct AlignmentEngine +{ + AudioModel model; + std::unique_ptr text; + AlignmentEngine(Runtime& runtime, const std::filesystem::path& dir); + std::vector Align(const std::vector& features, const std::string& transcript, + const std::string& language); +}; + +AlignmentEngine::AlignmentEngine(Runtime& runtime, const std::filesystem::path& dir) + : model(runtime, dir, "aligner") +{ + if (!model.metadata.value("timestamp_bins", false)) + throw std::runtime_error("Re-export the aligner with --only aligner for GPU timestamp selection"); + const auto shape = [&](int64_t seq, int slots) + { + return TextShape(seq, model.hidden, seq) + ",timestamp_indices:" + std::to_string(slots); + }; + text = runtime.Runner(dir, "aligner", shape(4, 2), shape(128, 32), shape(8192, 4096)); +} +std::vector AlignmentEngine::Align(const std::vector& features, const std::string& transcript, + const std::string& language) +{ + din::common::nvtx_scoped_range range{"qwen3.align"}; + if (transcript.empty()) + return {}; + const int64_t frames = features.size() / 128; + std::vector result; + const auto words = detail::AlignmentUnits(transcript, language); + if (words.empty()) + return {}; + if (words.size() > 2048) + throw std::runtime_error("Native alignment is limited to 2048 alignment units per chunk"); + auto audio = [&] + { + din::common::nvtx_scoped_range range{"qwen3.align_encoder"}; + return model.Encode(features, frames); + }(); + std::vector ids = model.native["audio_start"]; + ids.insert(ids.end(), audio.tokens, model.audio_id); + Append(ids, model.native["audio_end"].get>()); + std::vector slots; + const int64_t timestamp_id = model.metadata["timestamp_token_id"]; + for (const auto& word : words) + { + Append(ids, model.tokenizer->Encode(word, false)); + for (int i = 0; i < 2; ++i) + { + slots.push_back(ids.size()); + ids.push_back(timestamp_id); + } + } + const auto seq = static_cast(ids.size()); + if (seq > 8192) + throw std::runtime_error("Alignment exceeds the native context limit"); + TextInputs input(*text, seq, model.hidden, seq, model.dtype); + input.Fill(ids, 0, model.audio_id, &audio); + const int64_t labels = model.metadata["num_labels"]; + Buffer indices(*text, {static_cast(slots.size())}, text->HasDeviceIo()); + auto logits = TensorValue(*text, {1, static_cast(slots.size()), labels}, model.dtype); + Buffer output(*text, {1, static_cast(slots.size())}, text->HasDeviceIo()); + std::copy(slots.begin(), slots.end(), indices.HostData()); + indices.CopyAsyncToDevice(); + Ort::IoBinding binding(text->session); + input.Bind(binding); + binding.BindInput("timestamp_indices", indices.BindingValue()); + binding.BindOutput("timestamp_logits", logits); + binding.BindOutput("timestamp_bins", output.BindingValue()); + { + din::common::nvtx_scoped_range range{"qwen3.align_inference"}; + text->session.Run(Ort::RunOptions{}, binding); + } + { + din::common::nvtx_scoped_range range{"qwen3.align_download"}; + output.CopyAsyncToHostWithNotification().Sync(); + } + din::common::nvtx_scoped_range postprocess{"qwen3.align_postprocess"}; + const float timestamp_scale = model.metadata["timestamp_segment_time"]; + std::vector raw; + for (size_t i = 0; i < slots.size(); ++i) + { + const auto best = output.HostData()[i]; + if (best < 0 || best >= labels) + throw std::runtime_error("Invalid timestamp bin"); + raw.push_back(static_cast(best * timestamp_scale)); + } + const auto times = detail::FixTimestamps(raw); + for (size_t i = 0; i < words.size(); ++i) + result.push_back({words[i], times[2 * i] / 1000.f, times[2 * i + 1] / 1000.f}); + return result; +} +} // namespace din::asr::qwen3::detail + +namespace din::asr::qwen3 +{ +using namespace detail; +struct Qwen3ForcedAligner::Impl +{ + ForcedAlignerConfig config; + Runtime runtime; + AlignmentEngine engine; + explicit Impl(ForcedAlignerConfig cfg) + : config(std::move(cfg)) + , runtime(config.provider, config.ep_cache_dir, config.ep_context_dir, config.progress) + , engine(runtime, config.model_dir) + { + runtime.LoadMel(config.model_dir); + } + std::vector AlignAudio(const din::io::Audio& audio, const std::string& text, + const std::string& language) + { + din::io::Audio normalized; + const auto& source = NormalizeAudio(audio, normalized); + if (source.samples.size() > 180 * kRate) + throw std::invalid_argument("Standalone alignment accepts up to 180 seconds; supply audio/text segments"); + if (config.progress) + config.progress({din::common::ProgressStage::Aligning, "Aligning supplied text", 0, audio.Duration()}); + const auto features = runtime.Features(source.samples); + auto result = engine.Align(features, text, language); + if (config.progress) + config.progress( + {din::common::ProgressStage::Aligning, "Alignment complete", audio.Duration(), audio.Duration()}); + return result; + } +}; +Qwen3ForcedAligner::Qwen3ForcedAligner(ForcedAlignerConfig config) +{ + impl_ = std::make_unique(std::move(config)); +} +Qwen3ForcedAligner::~Qwen3ForcedAligner() = default; +std::vector Qwen3ForcedAligner::Align(const din::io::Audio& audio, const std::string& text, + const std::string& language) +{ + return impl_->AlignAudio(audio, text, language); +} +std::vector Qwen3ForcedAligner::AlignFile(const std::filesystem::path& path, const std::string& text, + const std::string& language) +{ + if (impl_->config.progress) + impl_->config.progress({din::common::ProgressStage::DecodingAudio, path.filename().string()}); + return Align(din::io::LoadAudio(path, kRate), text, language); +} +} // namespace din::asr::qwen3 diff --git a/asr/qwen3/forced_aligner.h b/asr/qwen3/forced_aligner.h new file mode 100644 index 0000000..48cc644 --- /dev/null +++ b/asr/qwen3/forced_aligner.h @@ -0,0 +1,44 @@ +// SPDX-License-Identifier: Apache-2.0 +#pragma once +#include +#include +#include +#include + +#include "audio.h" +#include "progress.h" + +namespace din::asr::qwen3 +{ +struct ForcedAlignerConfig +{ + std::string provider = "trt-rtx"; + std::filesystem::path model_dir = "artifacts/qwen3/aligner-onnx-bf16"; + std::filesystem::path ep_cache_dir = "artifacts/qwen3/rt_cache"; + std::filesystem::path ep_context_dir = "artifacts/qwen3/ep_context"; + din::common::ProgressCallback progress; +}; + +struct WordTimestamp +{ + std::string text; + float start_time = 0; + float end_time = 0; +}; + +// Accepts text from any recognizer; no ASR model is loaded. One synchronous call per instance. +class Qwen3ForcedAligner +{ +public: + explicit Qwen3ForcedAligner(ForcedAlignerConfig config = {}); + ~Qwen3ForcedAligner(); + std::vector Align(const din::io::Audio& audio, const std::string& text, + const std::string& language = "English"); + std::vector AlignFile(const std::filesystem::path& path, const std::string& text, + const std::string& language = "English"); + +private: + struct Impl; + std::unique_ptr impl_; +}; +} // namespace din::asr::qwen3 diff --git a/asr/qwen3/main.cpp b/asr/qwen3/main.cpp new file mode 100644 index 0000000..9b2287a --- /dev/null +++ b/asr/qwen3/main.cpp @@ -0,0 +1,60 @@ +// SPDX-License-Identifier: Apache-2.0 +#include +#include + +#include "qwen3.h" +#include +#include + +int main(int argc, char** argv) +{ + using namespace din::asr::qwen3; + try + { + Qwen3Config config; + argparse::ArgumentParser parser("din_asr_qwen3_cli"); + parser.add_description("Offline Qwen3 ASR."); + parser.add_argument("audiofile"); + parser.add_argument("--provider").default_value(config.provider).choices("cpu", "trt-rtx"); + parser.add_argument("--model-dir").default_value(config.model_dir.string()); + parser.add_argument("--lang-id") + .default_value(config.lang_id) + .help("Language code or name (case-insensitive); auto detects ASR language"); + parser.add_argument("--ep-cache").default_value(config.ep_cache_dir.string()); + parser.add_argument("--ep-context-dir").default_value(config.ep_context_dir.string()); + parser.add_argument("--max-new-tokens").default_value(config.max_new_tokens).scan<'i', int>(); + parser.add_argument("--max-chunk-seconds") + .default_value(config.max_chunk_seconds) + .scan<'i', int>() + .help("Chunk target: 0 = auto (1200 s), limited by KV capacity; boundaries may add 5 s"); + parser.parse_args(argc, argv); + config.provider = parser.get("--provider"); + config.model_dir = parser.get("--model-dir"); + config.lang_id = parser.get("--lang-id"); + config.ep_cache_dir = parser.get("--ep-cache"); + config.ep_context_dir = parser.get("--ep-context-dir"); + config.max_new_tokens = parser.get("--max-new-tokens"); + config.max_chunk_seconds = parser.get("--max-chunk-seconds"); + Qwen3Pipeline pipeline(std::move(config)); + const auto result = pipeline.TranscribeFile(parser.get("audiofile")); + nlohmann::json output{ + {"transcription", result.text}, {"language", result.language}, + {"reached_eos", result.reached_eos}, {"tokens", result.tokens}, + {"audio_seconds", result.audio_seconds}, {"transcribe_seconds", result.transcribe_seconds}}; + output["segments"] = nlohmann::json::array(); + output["sample_rate"] = 16000; + output["chunks_processed"] = result.chunks_processed; + for (const auto& segment : result.segments) + output["segments"].push_back({{"text", segment.text}, + {"language", segment.language}, + {"start_sample", segment.start_sample}, + {"end_sample", segment.end_sample}}); + std::cout << output.dump(2) << '\n'; + return result.reached_eos ? 0 : 1; + } + catch (const std::exception& error) + { + std::cerr << "Qwen3: " << error.what() << '\n'; + return 1; + } +} diff --git a/asr/qwen3/model_export/export_qwen3_asr.py b/asr/qwen3/model_export/export_qwen3_asr.py new file mode 100644 index 0000000..a46fefc --- /dev/null +++ b/asr/qwen3/model_export/export_qwen3_asr.py @@ -0,0 +1,412 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Export Qwen3-ASR with explicit tensor interfaces; checkpoint BF16 is the default.""" + +import argparse +import hashlib +import json +import sys +from pathlib import Path +from types import MethodType + +import onnx +import torch +import transformers +from torch import nn +from torch.nn import functional as F +from transformers import ( + AttentionInterface, + AutoProcessor, + Qwen3ASRForConditionalGeneration, + Qwen3ASRForTokenClassification, +) +from transformers.models.qwen3_asr.processing_qwen3_asr import LANGUAGE_CODE_TO_NAME + +sys.path.insert(0, str(Path(__file__).resolve().parents[3])) +from common.model_export.kv_cache import FixedKVCache # noqa: E402 +from common.model_export.log_mel import LogMel as WhisperMel # noqa: E402 + + +class LogMel(WhisperMel): + def __init__(self, processor): + fe = processor.feature_extractor + if (fe.n_fft, fe.hop_length, fe.feature_size, fe.sampling_rate, fe.dither) != (400, 160, 128, 16000, 0): + raise ValueError("Native frontend requires the standard Qwen3 16 kHz configuration") + super().__init__(fe.mel_filters.T.copy(), torch.float32, frames=None) + + +def save_native_assets(processor, output, task): + tokenizer = processor.tokenizer + + def encode(text): + return tokenizer.encode(text, add_special_tokens=False) + + data = { + "audio_start": encode(processor.audio_bos_token), + "audio_end": encode(processor.audio_eos_token), + } + if task == "asr": + prefixes = {} + suffixes = {} + languages = {} + for language in [None, *LANGUAGE_CODE_TO_NAME.values()]: + messages = [{"role": "user", "content": [{"type": "audio"}]}] + rendered = tokenizer.apply_chat_template( + messages, chat_template=processor.chat_template, tokenize=False, add_generation_prompt=True + ) + prefix, suffix = rendered.split(processor.audio_token) + prefixes[language or "auto"] = encode(prefix) + suffixes[language or "auto"] = encode(suffix + (f"language {language}" if language else "")) + languages[language or "auto"] = language or "" + for code, language in LANGUAGE_CODE_TO_NAME.items(): + prefixes[code] = prefixes[language] + suffixes[code] = suffixes[language] + languages[code] = language + data["prefixes"] = prefixes + data["suffixes"] = suffixes + data["languages"] = languages + data["suffix"] = suffixes["auto"] + (output / "native.json").write_text(json.dumps(data, indent=2), encoding="utf-8") + + +def decode_attention(module, query, key, value, attention_mask, scaling=None, **kwargs): + if query.dtype == torch.float32: + return export_attention(module, query, key, value, attention_mask, scaling, **kwargs) + output = F.scaled_dot_product_attention( + query, + key, + value, + attn_mask=attention_mask, + scale=scaling, + enable_gqa=query.shape[1] != key.shape[1], + ) + return output.transpose(1, 2), None + + +def attention(q, k, v, bias, scale): + scores = q @ k.transpose(-1, -2) * scale + bias + return torch.softmax(scores.float(), dim=-1).to(v.dtype) @ v + + +def audio_attention(module, hidden_states, attention_bias, **kwargs): + q, k, v = [ + p(hidden_states).reshape(1, -1, module.num_heads, module.head_dim).transpose(1, 2) + for p in (module.q_proj, module.k_proj, module.v_proj) + ] + y = attention(q, k, v, attention_bias, module.scaling) + return module.out_proj(y.transpose(1, 2).flatten(2)) + + +class AudioEncoder(nn.Module): + def __init__(self, model): + super().__init__() + if model.config.audio_config.n_window != 50 or model.config.audio_config.num_mel_bins != 128: + raise ValueError("This export contract requires n_window=50 and 128 mel bins") + self.encoder = model.model.audio_tower + self.projector = model.model.multi_modal_projector + # Replace only HF's data-dependent attention splits; keep its encoder layers. + for layer in self.encoder.layers: + layer.self_attn.forward = MethodType(audio_attention, layer.self_attn) + + def forward(self, mel_chunks, valid_indices, attention_bias): + enc = self.encoder + x = F.gelu(enc.conv2d1(mel_chunks.unsqueeze(1))) + x = F.gelu(enc.conv2d2(x)) + x = F.gelu(enc.conv2d3(x)) + x = x.permute(0, 3, 1, 2).flatten(2) + x = enc.conv_out(x) + x = x + enc.positional_embedding.positional_embedding[: x.shape[1]].to(x.dtype) + x = x.flatten(0, 1).index_select(0, valid_indices).unsqueeze(0) + for layer in enc.layers: + x = layer(x, cu_seqlens=None, attention_bias=attention_bias)[0] + return self.projector(enc.ln_post(x)).squeeze(0) + + +def export_attention(module, query, key, value, attention_mask, scaling=None, **kwargs): + # Fold query groups into the query axis. TensorRT RTX otherwise lowers GQA + # to full-size KV-head replication. This is ordinary MHA with identical math + # and the original smaller KV tensors; only the small query/mask is repeated. + batch, heads, sequence, width = query.shape + kv_heads = key.shape[1] + groups = heads // kv_heads + query = query.reshape(batch, kv_heads, groups * sequence, width) + if attention_mask is not None: + attention_mask = attention_mask.repeat(1, 1, groups, 1) + # TensorRT RTX 1.6 has no fused FP32 Attention kernel for this graph. + if query.dtype == torch.float32: + output = attention(query, key, value, attention_mask, scaling) + else: + output = F.scaled_dot_product_attention(query, key, value, attn_mask=attention_mask, scale=scaling) + output = output.reshape(batch, heads, sequence, width) + return output.transpose(1, 2), None + + +AttentionInterface.register("qwen3_onnx", export_attention) + + +class TextBackbone(nn.Module): + def __init__(self, model): + super().__init__() + if model.config.text_config.rope_parameters["rope_type"] != "default": + raise ValueError("Only default RoPE is supported") + self.decoder = model.model.language_model + self.decoder.config._attn_implementation = "qwen3_onnx" + + def hidden(self, input_ids, audio_embeddings, audio_mask, position_ids, attention_bias, *past): + dec = self.decoder + x = torch.where(audio_mask, audio_embeddings, dec.embed_tokens(input_ids)) + rotary = dec.rotary_emb(x, position_ids) + cache = FixedKVCache(past, position_ids, tensor_scatter=True) if past else None + for layer in dec.layers: + x = layer(x, attention_mask=attention_bias, position_embeddings=rotary, past_key_values=cache) + return dec.norm(x), cache.present if cache else [] + + +class TextDecoder(TextBackbone): + """Prompt prefill and single-token decoding with fixed-capacity KV tensors.""" + + def __init__(self, model): + super().__init__(model) + self.lm_head = model.lm_head + AttentionInterface.register("qwen3_onnx_decode", decode_attention) + self.decoder.config._attn_implementation = "qwen3_onnx_decode" + + def forward(self, input_ids, audio_embeddings, audio_mask, position_ids, attention_bias, logits_index, *past): + hidden, present = self.hidden(input_ids, audio_embeddings, audio_mask, position_ids, attention_bias, *past) + logits = self.lm_head(hidden.index_select(1, logits_index).squeeze(1)) + return logits, logits.argmax(-1), *present + + +class ForcedAligner(TextBackbone): + """Predict all timestamp slots together, without an autoregressive cache.""" + + def __init__(self, model): + super().__init__(model) + self.score = model.score + + def forward(self, input_ids, audio_embeddings, audio_mask, position_ids, attention_bias, timestamp_indices): + hidden, _ = self.hidden(input_ids, audio_embeddings, audio_mask, position_ids, attention_bias) + logits = self.score(hidden.index_select(1, timestamp_indices)) + return logits, logits.argmax(-1) + + +def export(module, args, path, names, outputs, shapes): + print(f"Exporting {path.name}", flush=True) + path.unlink(missing_ok=True) + path.with_suffix(".onnx.data").unlink(missing_ok=True) + torch.onnx.export( + module.eval(), + args, + str(path), + input_names=names, + output_names=outputs, + dynamo=True, + dynamic_shapes=shapes, + opset_version=24, + external_data=True, + ) + onnx.checker.check_model(str(path)) + print(f"Checked {path}", flush=True) + + +def save_metadata(output, metadata): + metadata["graphs"] = {} + for name in ("mel", "encoder", "decoder" if metadata["task"] == "asr" else "aligner"): + if not (output / (name + ".onnx")).is_file(): + continue + digest = hashlib.sha256() + for suffix in (".onnx", ".onnx.data"): + path = output / (name + suffix) + if path.exists(): + with path.open("rb") as file: + while data := file.read(8 * 1024 * 1024): + digest.update(data) + metadata["graphs"][name] = digest.hexdigest() + (output / "metadata.json").write_text(json.dumps(metadata, indent=2), encoding="utf-8") + + +def export_mel(processor, output): + export( + LogMel(processor), + (torch.zeros(1, 16000),), + output / "mel.onnx", + ["samples"], + ["features"], + {"samples": {1: torch.export.Dim("samples", min=8000, max=1205 * 16000)}}, + ) + + +@torch.inference_mode() +def export_encoder(model, output): + dtype = model.dtype + export( + AudioEncoder(model), + (torch.zeros(2, 128, 100, dtype=dtype), torch.arange(26), torch.zeros(1, 1, 26, 26, dtype=dtype)), + output / "encoder.onnx", + ["mel_chunks", "valid_indices", "attention_bias"], + ["audio_embeddings"], + { + "mel_chunks": {0: torch.export.Dim("chunks", min=1)}, + "valid_indices": {0: torch.export.Dim("audio_tokens", min=1)}, + "attention_bias": {2: torch.export.Dim("audio_tokens", min=1), 3: torch.export.Dim("audio_tokens", min=1)}, + }, + ) + + +@torch.inference_mode() +def export_decoder(model, output, cache_capacity): + dtype = model.dtype + c = model.config.text_config + past = tuple( + torch.zeros(1, c.num_key_value_heads, cache_capacity, c.head_dim, dtype=dtype) + for _ in range(2 * c.num_hidden_layers) + ) + seq = torch.export.Dim("sequence", min=1, max=512) + inputs = ( + torch.ones(1, 3, dtype=torch.int64), + torch.zeros(1, 3, c.hidden_size, dtype=dtype), + torch.zeros(1, 3, 1, dtype=torch.bool), + torch.arange(3)[None], + torch.zeros(1, 1, 3, cache_capacity, dtype=dtype), + torch.tensor([2]), + *past, + ) + names = ["input_ids", "audio_embeddings", "audio_mask", "position_ids", "attention_bias", "logits_index"] + export( + TextDecoder(model), + inputs, + output / "decoder.onnx", + names + [f"past_{i}" for i in range(len(past))], + ["logits", "next_token"] + [f"present_{i}" for i in range(len(past))], + { + "input_ids": {1: seq}, + "audio_embeddings": {1: seq}, + "audio_mask": {1: seq}, + "position_ids": {1: seq}, + "attention_bias": {2: seq}, + "logits_index": {}, + "past": tuple({} for _ in past), + }, + ) + + +@torch.inference_mode() +def export_aligner(model, output): + dtype = model.dtype + seq = torch.export.Dim("sequence", min=1) + inputs = ( + torch.ones(1, 3, dtype=torch.int64), + torch.zeros(1, 3, model.config.text_config.hidden_size, dtype=dtype), + torch.zeros(1, 3, 1, dtype=torch.bool), + torch.arange(3)[None], + torch.zeros(1, 1, 3, 3, dtype=dtype), + torch.tensor([1, 2]), + ) + export( + ForcedAligner(model), + inputs, + output / "aligner.onnx", + ["input_ids", "audio_embeddings", "audio_mask", "position_ids", "attention_bias", "timestamp_indices"], + ["timestamp_logits", "timestamp_bins"], + { + "input_ids": {1: seq}, + "audio_embeddings": {1: seq}, + "audio_mask": {1: seq}, + "position_ids": {1: seq}, + "attention_bias": {2: seq, 3: seq}, + "timestamp_indices": {0: torch.export.Dim("timestamp_slots", min=1)}, + }, + ) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--model", "--checkpoint", dest="model", help="HF model ID or local checkpoint directory") + parser.add_argument("--size", choices=["0.6B", "1.7B"], default="0.6B", help="ASR model size") + parser.add_argument("--revision", help="Optional HF revision or commit") + parser.add_argument("--output", type=Path, help="Defaults to the ONNX artifact directory for --task") + parser.add_argument("--task", choices=["asr", "aligner"], default="asr") + parser.add_argument("--dtype", choices=["original", "fp16", "fp32"], default="original") + parser.add_argument("--only", choices=["mel", "encoder", "decoder", "aligner"]) + parser.add_argument("--threads", type=int, default=4) + parser.add_argument("--cache-capacity", type=int, choices=range(512, 16385, 512), default=8192, metavar="TOKENS") + args = parser.parse_args() + prefix = "aligner-" if args.task == "aligner" else "" + args.model = args.model or ( + "Qwen/Qwen3-ForcedAligner-0.6B-hf" if args.task == "aligner" else f"Qwen/Qwen3-ASR-{args.size}-hf" + ) + precision = "bf16" if args.dtype == "original" else args.dtype + size_suffix = "-1.7b" if args.task == "asr" and "1.7b" in args.model.lower() else "" + args.output = args.output or Path(f"artifacts/qwen3/{prefix}onnx-{precision}{size_suffix}") + if args.only in ("decoder", "aligner") and args.only != ("decoder" if args.task == "asr" else "aligner"): + parser.error("--only must match --task") + metadata_path = args.output / "metadata.json" + existing = json.loads(metadata_path.read_text(encoding="utf-8")) if args.only and metadata_path.exists() else {} + requested = {"original": "bfloat16", "fp16": "float16", "fp32": "float32"}[args.dtype] + if ( + existing + and args.only != "mel" + and (existing["dtype"] != requested or existing["task"] != args.task or existing.get("quantization")) + ): + parser.error("Partial export requires the same task and precision, without quantization; use a new directory") + torch.set_num_threads(args.threads) + args.output.mkdir(parents=True, exist_ok=True) + revision = args.revision + if args.only != "mel": + model_class = Qwen3ASRForConditionalGeneration if args.task == "asr" else Qwen3ASRForTokenClassification + model = model_class.from_pretrained( + args.model, + dtype={"original": "auto", "fp16": torch.float16, "fp32": torch.float32}[args.dtype], + attn_implementation="eager", + revision=args.revision, + ).eval() + revision = model.config._commit_hash or args.revision + processor = AutoProcessor.from_pretrained(args.model, revision=revision) + processor.save_pretrained(args.output / "processor") + save_native_assets(processor, args.output, args.task) + if args.only in (None, "mel"): + export_mel(processor, args.output) + if args.only == "mel": + if existing: + save_metadata(args.output, existing) + return + cfg = model.config + dtype = model.dtype + if args.only in (None, "encoder"): + export_encoder(model, args.output) + if args.task == "asr" and args.only in (None, "decoder"): + export_decoder(model, args.output, args.cache_capacity) + if args.task == "aligner" and args.only in (None, "aligner"): + export_aligner(model, args.output) + eos = cfg.eos_token_id + metadata = { + "format_version": 3 if args.task == "asr" else 2, + "cache_capacity": args.cache_capacity if args.task == "asr" else None, + "prefill_block": 512 if args.task == "asr" else None, + "task": args.task, + "dtype": str(dtype).removeprefix("torch."), + "opset": 24, + "batch_size": 1, + "audio_config": cfg.audio_config.to_dict(), + "text_config": cfg.text_config.to_dict(), + "audio_token_id": cfg.audio_token_id, + "eos_token_ids": list(eos) if isinstance(eos, (list, tuple)) else [eos], + "torch": torch.__version__, + "transformers": transformers.__version__, + } + if args.task == "aligner": + metadata.update( + timestamp_bins=True, + timestamp_token_id=cfg.timestamp_token_id, + num_labels=cfg.num_labels, + timestamp_segment_time=processor.timestamp_segment_time, + ) + metadata["source"] = {"model": args.model, "revision": revision} + if args.only == "encoder" and existing: + for key in ("cache_capacity", "format_version", "prefill_block", "timestamp_bins"): + if key in existing: + metadata[key] = existing[key] + save_metadata(args.output, metadata) + + +if __name__ == "__main__": + main() diff --git a/asr/qwen3/model_export/validate_qwen3_asr.py b/asr/qwen3/model_export/validate_qwen3_asr.py new file mode 100644 index 0000000..5f3f7ec --- /dev/null +++ b/asr/qwen3/model_export/validate_qwen3_asr.py @@ -0,0 +1,391 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Validate Qwen3 ASR or forced alignment against HF in checkpoint precision.""" + +import argparse +import json +from pathlib import Path +from unittest.mock import patch + +import onnxruntime as ort +import torch +from transformers import AutoProcessor, Qwen3ASRForConditionalGeneration, Qwen3ASRForTokenClassification +from transformers.modeling_outputs import CausalLMOutputWithPast + +_REGISTERED_EP = None + + +def session_options(provider, threads, extra_options=None): + global _REGISTERED_EP + options = ort.SessionOptions() + options.intra_op_num_threads = threads + if provider == "cpu": + return options, ["CPUExecutionProvider"] + if provider != "trt-rtx": + raise ValueError(f"Unknown provider: {provider}") + import onnxruntime_ep_nv_tensorrt_rtx as ep + + if _REGISTERED_EP is None: + _REGISTERED_EP = ep.get_ep_name() + ort.register_execution_provider_library(_REGISTERED_EP, ep.get_library_path()) + devices = [device for device in ort.get_ep_devices() if device.ep_name == _REGISTERED_EP] + if not devices: + raise RuntimeError("TensorRT RTX registered but no GPU was discovered") + options.add_provider_for_devices(devices, {"enable_cuda_graph": "0", **(extra_options or {})}) + return options, None + + +def pack_audio(features, mask, config): + chunk_size = config["n_window"] * 2 + if features.shape[0] != 1 or features.shape[-1] % chunk_size: + raise ValueError("Expected batch-one features padded to the encoder chunk size") + chunks = features.reshape(1, features.shape[1], -1, chunk_size)[0].permute(1, 0, 2) + lengths = mask.reshape(-1, chunk_size).sum(1) + for _ in range(3): + lengths = (lengths + 1) // 2 + indices = (torch.arange((chunk_size + 7) // 8)[None] < lengths[:, None]).flatten().nonzero().flatten() + if not len(indices): + raise ValueError("Audio has no valid feature frames") + window = int(lengths.max()) * (config["n_window_infer"] // chunk_size) + groups = torch.arange(len(indices)) // window + bias = torch.zeros(len(indices), len(indices)).masked_fill(groups[:, None] != groups[None], -1e4) + return chunks, indices, bias[None, None] + + +def decoder_inputs(ids, embeddings, audio_token_id, hidden_size, past_length, capacity=None): + ids = ids.cpu().long().reshape(1, -1) + seq = ids.shape[1] + if not seq or past_length < 0 or (capacity is not None and past_length + seq > capacity): + raise ValueError("Input exceeds KV cache capacity or has invalid length") + mask = (ids == audio_token_id)[..., None] if embeddings is not None else torch.zeros(1, seq, 1, dtype=torch.bool) + padded = torch.zeros(1, seq, hidden_size) + if embeddings is not None: + if int(mask.sum()) != len(embeddings): + raise ValueError("Audio placeholder count differs from encoder output length") + padded[mask[..., 0]] = embeddings.float().cpu() + positions = torch.arange(past_length, past_length + seq)[None] + allowed = torch.arange(capacity or past_length + seq)[None] <= positions.T + return { + "input_ids": ids, + "audio_embeddings": padded, + "audio_mask": mask, + "position_ids": positions, + "attention_bias": torch.zeros_like(allowed, dtype=torch.float32).masked_fill(~allowed, -1e4)[None, None], + } + + +class OnnxAudioModel: + def __init__(self, directory, task, threads=4, provider=None): + self.text_graph = "decoder.onnx" if task == "asr" else "aligner.onnx" + directory = Path(directory) + self.directory = directory + self.threads = threads + self.metadata = json.loads((directory / "metadata.json").read_text(encoding="utf-8")) + if self.metadata["task"] != task or self.metadata["format_version"] != (3 if task == "asr" else 2): + raise ValueError("Incompatible export; re-export the model with the current exporter") + self.processor = AutoProcessor.from_pretrained(directory / "processor", local_files_only=True) + self.dtype = getattr(torch, self.metadata["dtype"]) + self.provider = provider or ("trt-rtx" if self.dtype in (torch.float16, torch.bfloat16) else "cpu") + encoder_shape = "mel_chunks:8x128x100,valid_indices:104,attention_bias:1x1x104x104" + options, providers = session_options( + self.provider, threads, {f"nv_profile_{bound}_shapes": encoder_shape for bound in ("min", "opt", "max")} + ) + self.encoder = ort.InferenceSession(str(directory / "encoder.onnx"), options, providers=providers) + if self.text_graph == "decoder.onnx": + c = self.metadata["text_config"] + capacity = self.metadata["cache_capacity"] + + def shapes(sequence): + return ( + f"input_ids:1x{sequence},audio_embeddings:1x{sequence}x{c['hidden_size']}," + f"audio_mask:1x{sequence}x1,position_ids:1x{sequence}," + f"attention_bias:1x1x{sequence}x{capacity},logits_index:1" + ) + + sessions = [] + for sequence in (1, self.metadata["prefill_block"]): + options, providers = session_options( + self.provider, + threads, + {f"nv_profile_{bound}_shapes": shapes(sequence) for bound in ("min", "opt", "max")}, + ) + sessions.append(ort.InferenceSession(str(directory / self.text_graph), options, providers=providers)) + self.decoder, self.prefill = sessions + else: + options, providers = session_options(self.provider, threads) + self.decoder = ort.InferenceSession(str(directory / self.text_graph), options, providers=providers) + + def run(self, session, feed, inplace=False): + values = {} + for name, tensor in feed.items(): + tensor = tensor.detach().to("cuda" if inplace else "cpu").contiguous() + if tensor.is_floating_point(): + tensor = tensor.to(self.dtype) + # DLPack preserves BF16; ORT expects boolean capsules encoded as uint8. + if tensor.dtype == torch.bool: + capsule = torch.utils.dlpack.to_dlpack(tensor.view(torch.uint8)) + values[name] = ort.OrtValue(ort.capi._pybind_state.OrtValue.from_dlpack(capsule, True)) + else: + values[name] = ort.OrtValue.from_dlpack(tensor) + if inplace: + # TensorRT requires TensorScatter past/present to alias, even for validation. + binding = session.io_binding() + for name, value in values.items(): + binding.bind_ortvalue_input(name, value) + for output in session.get_outputs(): + if output.name.startswith("present_"): + binding.bind_ortvalue_output(output.name, values[output.name.replace("present_", "past_")]) + else: + binding.bind_output(output.name) + torch.cuda.synchronize() + session.run_with_iobinding(binding) + binding.synchronize_outputs() + outputs = binding.get_outputs() + else: + outputs = session.run_with_ort_values(None, values) + return [torch.from_dlpack(value).cpu().clone() for value in outputs] + + def encode(self, inputs): + config = self.metadata["audio_config"] + window = config["n_window_infer"] + features, mask = inputs["input_features"], inputs["input_features_mask"] + outputs = [] + for start in range(0, int(mask.sum()), window): + chunks, indices, _ = pack_audio( + features[..., start : start + window], mask[..., start : start + window], config + ) + tokens = window // 100 * 13 + bias = torch.zeros(1, 1, tokens, tokens) + bias[..., len(indices) :] = -1e4 + packed = torch.zeros(window // 100, 128, 100) + packed[: len(chunks)] = chunks + outputs.append( + self.run( + self.encoder, {"mel_chunks": packed, "valid_indices": torch.arange(tokens), "attention_bias": bias} + )[0][: len(indices)] + ) + return torch.cat(outputs) + + def empty_cache(self, capacity): + c = self.metadata["text_config"] + return [ + torch.zeros(1, c["num_key_value_heads"], capacity, c["head_dim"], dtype=self.dtype) + for _ in range(2 * c["num_hidden_layers"]) + ] + + def step(self, ids, embeddings, cache, length): + ids = ids.reshape(1, -1) + block = self.metadata["prefill_block"] if ids.numel() > 1 else 1 + audio_offset = 0 + for offset in range(0, ids.numel(), block): + chunk = ids[:, offset : offset + block] + valid = chunk.numel() + padded = torch.zeros(1, block, dtype=torch.int64) + padded[:, :valid] = chunk + count = int((chunk == self.metadata["audio_token_id"]).sum()) + audio = embeddings[audio_offset : audio_offset + count] if embeddings is not None else None + audio_offset += count + feed = decoder_inputs( + padded, + audio, + self.metadata["audio_token_id"], + self.metadata["text_config"]["hidden_size"], + length + offset, + cache[0].shape[2], + ) + feed["logits_index"] = torch.tensor([valid - 1]) + feed.update({f"past_{i}": value for i, value in enumerate(cache)}) + logits, _, *cache = self.run( + self.prefill if block > 1 else self.decoder, feed, inplace=self.provider == "trt-rtx" + ) + return logits, cache + + +def compare(actual, expected, atol): + actual, expected = actual.float().cpu(), expected.float().cpu() + error = (actual - expected).abs().max().item() + return { + "passed": bool(torch.isfinite(actual).all() and torch.isfinite(expected).all() and error <= atol), + "max_abs_error": error, + } + + +def compare_cache(actual, expected, length, atol): + values = [value for layer in expected.layers for value in (layer.keys, layer.values)] + checks = [compare(a[:, :, :length], b, atol) for a, b in zip(actual, values, strict=True)] + return {"passed": all(c["passed"] for c in checks), "max_abs_error": max(c["max_abs_error"] for c in checks)} + + +def generate_onnx(reference, runner, inputs, embeddings, capacity, max_new_tokens): + cache, length = runner.empty_cache(capacity), 0 + + def forward(input_ids, input_features=None, input_features_mask=None, attention_mask=None, **kwargs): + nonlocal cache, length + ids = input_ids[:, length:] + logits, cache = runner.step(ids, embeddings if length == 0 else None, cache, length) + length += ids.shape[1] + return CausalLMOutputWithPast(logits=logits[:, None].to(input_ids.device)) + + # HF owns token selection and stopping. Only graph execution is replaced. + with patch.object(reference, "forward", forward): + return reference.generate(**inputs, use_cache=False, do_sample=False, max_new_tokens=max_new_tokens) + + +def validate_asr(args, runner, reference, inputs, ref_inputs): + length = inputs["input_ids"].shape[1] + required = length + args.max_new_tokens + metadata = runner.metadata + capacity = metadata["cache_capacity"] + if required > capacity: + raise ValueError("Prompt and generation budget exceed exported cache capacity") + audio = runner.encode(inputs) + expected_audio = reference.get_audio_features( + ref_inputs["input_features"], ref_inputs["input_features_mask"] + ).pooler_output + expected = reference(**ref_inputs, use_cache=True, logits_to_keep=1) + logits, cache = runner.step(inputs["input_ids"], audio, runner.empty_cache(capacity), 0) + checks = { + "encoder": compare(audio, expected_audio, args.atol), + "prefill_logits": compare(logits, expected.logits[:, -1], args.atol), + "prefill_kv": compare_cache(cache, expected.past_key_values, length, args.atol), + } + token = expected.logits[:, -1].argmax(-1, keepdim=True) + next_expected = reference(input_ids=token, past_key_values=expected.past_key_values, use_cache=True) + next_logits, present = runner.step(token, None, cache, length) + checks["cached_logits"] = compare(next_logits, next_expected.logits[:, -1], args.atol) + checks["cached_kv"] = compare_cache(present, next_expected.past_key_values, length + 1, args.atol) + expected_ids = reference.generate(**ref_inputs, do_sample=False, max_new_tokens=args.max_new_tokens) + actual_ids = generate_onnx(reference, runner, ref_inputs, audio, capacity, args.max_new_tokens) + expected_tokens, actual_tokens = expected_ids[0, length:].tolist(), actual_ids[0, length:].tolist() + eos = metadata["eos_token_ids"] + exact = expected_tokens == actual_tokens + stopped = bool(expected_tokens and actual_tokens and expected_tokens[-1] in eos and actual_tokens[-1] in eos) + return { + "checks": checks, + "exact_token_match": exact, + "reached_eos": stopped, + "reference_tokens": expected_tokens, + "onnx_tokens": actual_tokens, + "passed": all(c["passed"] for c in checks.values()) and exact and stopped, + } + + +def validate_aligner(args, runner, reference, inputs, ref_inputs, words): + audio = runner.encode(inputs) + expected_audio = reference.model.get_audio_features( + ref_inputs["input_features"], ref_inputs["input_features_mask"] + ).pooler_output + expected = reference(**ref_inputs, use_cache=False).logits.cpu() + timestamp_id = reference.config.timestamp_token_id + indices = (inputs["input_ids"][0] == timestamp_id).nonzero().flatten() + feed = decoder_inputs( + inputs["input_ids"], audio, runner.metadata["audio_token_id"], runner.metadata["text_config"]["hidden_size"], 0 + ) + feed["timestamp_indices"] = indices + logits = runner.run(runner.decoder, feed)[0] + processor = runner.processor + expected_spans = processor.decode_forced_alignment(expected.float(), inputs["input_ids"], words, timestamp_id)[0] + actual_spans = processor.decode_forced_alignment( + logits.float(), torch.full((1, len(indices)), timestamp_id), words, timestamp_id + )[0] + checks = { + "encoder": compare(audio, expected_audio, args.atol), + "timestamp_logits": compare(logits, expected[:, indices], args.atol), + } + exact = actual_spans == expected_spans + return { + "checks": checks, + "exact_spans": exact, + "reference_spans": expected_spans, + "onnx_spans": actual_spans, + "passed": all(c["passed"] for c in checks.values()) and exact, + } + + +@torch.inference_mode() +def validate(args): + runner = OnnxAudioModel(args.onnx_dir, args.task, args.threads, args.provider) + model_type = Qwen3ASRForConditionalGeneration if args.task == "asr" else Qwen3ASRForTokenClassification + reference = ( + model_type.from_pretrained(args.model, revision=args.revision, dtype=runner.dtype, attn_implementation="eager") + .to(args.reference_device) + .eval() + ) + reports = [] + for path in args.audio: + if args.task == "asr": + inputs = runner.processor.apply_transcription_request( + audio=str(path), language=args.language, return_tensors="pt" + ) + else: + inputs, words = runner.processor.prepare_forced_aligner_inputs( + audio=str(path), + transcript=args.transcript.read_text(encoding="utf-8-sig").strip(), + language=args.language, + return_tensors="pt", + ) + if not words[0]: + raise ValueError("Transcript must contain alignable words") + ref_inputs = { + k: v.to(device=args.reference_device, dtype=runner.dtype if v.is_floating_point() else v.dtype) + for k, v in inputs.items() + } + report = ( + validate_asr(args, runner, reference, inputs, ref_inputs) + if args.task == "asr" + else validate_aligner(args, runner, reference, inputs, ref_inputs, words) + ) + reports.append({"audio": str(path), **report}) + payload = { + "passed": all(r["passed"] for r in reports), + "atol": args.atol, + "provider": runner.provider, + "dtype": str(runner.dtype), + "source": runner.metadata["source"], + "cases": reports, + } + args.report.parent.mkdir(parents=True, exist_ok=True) + args.report.write_text(json.dumps(payload, indent=2, ensure_ascii=False), encoding="utf-8") + print(json.dumps(payload, indent=2, ensure_ascii=False)) + raise SystemExit(0 if payload["passed"] else 1) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--task", choices=["asr", "aligner"], default="asr") + parser.add_argument("--onnx-dir", type=Path) + parser.add_argument( + "--model", "--checkpoint", dest="model", help="HF model ID or local checkpoint; defaults to export metadata" + ) + parser.add_argument("--revision", help="Optional HF revision; defaults to export metadata") + parser.add_argument("--audio", type=Path, nargs="+", default=[Path("assets/sample.wav")]) + parser.add_argument("--transcript", type=Path, help="UTF-8 transcript file, required for alignment") + parser.add_argument("--language") + parser.add_argument("--provider", choices=["cpu", "trt-rtx"]) + parser.add_argument("--reference-device", choices=["cpu", "cuda"], default="cuda") + parser.add_argument("--max-new-tokens", type=int, default=128) + parser.add_argument("--threads", type=int, default=4) + parser.add_argument("--atol", type=float, default=0.005) + parser.add_argument("--report", type=Path) + args = parser.parse_args() + if args.threads < 1 or args.max_new_tokens < 1 or args.atol <= 0: + parser.error("threads, max-new-tokens and atol must be positive") + if args.task == "aligner" and (args.transcript is None or len(args.audio) != 1): + parser.error("Alignment requires one --audio and a --transcript") + prefix = "aligner-" if args.task == "aligner" else "" + args.onnx_dir = args.onnx_dir or Path(f"artifacts/qwen3/{prefix}onnx-bf16") + args.report = args.report or Path(f"artifacts/qwen3/validation-{prefix}bf16.json") + metadata = json.loads((args.onnx_dir / "metadata.json").read_text(encoding="utf-8")) + if metadata["task"] != args.task: + parser.error("--task does not match the exported model") + source = metadata["source"] + if args.model is None: + args.model = source["model"] + args.revision = args.revision or source.get("revision") + torch.set_num_threads(args.threads) + if args.task == "aligner": + args.language = args.language or "English" + validate(args) + + +if __name__ == "__main__": + main() diff --git a/asr/qwen3/qwen3.cpp b/asr/qwen3/qwen3.cpp new file mode 100644 index 0000000..f36ff91 --- /dev/null +++ b/asr/qwen3/qwen3.cpp @@ -0,0 +1,385 @@ +// SPDX-License-Identifier: Apache-2.0 +#include "qwen3.h" + +#include +#include +#include +#include +#include + +#include "detail/runtime.h" + +namespace din::asr::qwen3::detail +{ +constexpr int kMaxSamples = 1205 * kRate; +constexpr int kPrefillBlock = 512; + +inline int ResolveChunkSeconds(int seconds) +{ + if (seconds == 0) + return 1200; + // A target inside the search radius can repeatedly split silence into single samples. + if (seconds <= 5 || seconds > 1200) + throw std::invalid_argument("max-chunk-seconds must be 0 (auto) or 6..1200"); + return seconds; +} + +struct AudioChunk +{ + size_t begin; + size_t end; +}; + +inline float WindowEnergy(std::span samples, size_t begin) +{ + // Independent FP32 reductions avoid cancellation drift across quiet windows. + std::array sums{}; + for (size_t i = 0; i < 1600; i += 8) + for (size_t j = 0; j < 8; ++j) + sums[j] += std::abs(samples[begin + i + j]); + return ((sums[0] + sums[1]) + (sums[2] + sums[3])) + ((sums[4] + sums[5]) + (sums[6] + sums[7])); +} + +// Qwen3-ASR split_audio_into_chunks: quietest 100 ms within +/-5 seconds, +// then the quietest sample inside that window. Keep the first minimum on ties. +inline std::vector SplitAudio(std::span samples, size_t target_seconds = 1200) +{ + const size_t target = target_seconds * 16000; + constexpr size_t expand = 5 * 16000, window = 1600; + std::vector chunks; + size_t start = 0; + while (samples.size() - start > target) + { + const auto cut = start + target; + const auto left = cut > expand ? std::max(start, cut - expand) : start; + const auto right = std::min(samples.size(), cut + expand); + size_t boundary = cut; + if (right - left > window) + { + float best = WindowEnergy(samples, left); + size_t minimum = left; + for (size_t i = left + 1; i + window <= right; ++i) + { + const float sum = WindowEnergy(samples, i); + if (sum < best) + { + best = sum; + minimum = i; + } + } + boundary = minimum; + for (size_t i = minimum + 1; i < minimum + window; ++i) + if (std::abs(samples[i]) < std::abs(samples[boundary])) + boundary = i; + } + boundary = std::clamp(boundary, start + 1, samples.size()); + chunks.push_back({start, boundary}); + start = boundary; + } + if (start < samples.size()) + chunks.push_back({start, samples.size()}); + return chunks; +} +} // namespace din::asr::qwen3::detail + +namespace din::asr::qwen3::detail +{ +std::string Trim(const std::string& text) +{ + const auto first = text.find_first_not_of(" \r\n\t"); + if (first == std::string::npos) + return {}; + return text.substr(first, text.find_last_not_of(" \r\n\t") - first + 1); +} + +// Official detect_and_fix_repetitions operates on Unicode characters, not UTF-8 bytes. +static std::string FixRepetitions(const std::string& text) +{ + std::vector chars, filtered; + for (size_t i = 0; i < text.size();) + { + const auto c = static_cast(text[i]); + const size_t length = c < 0x80 ? 1 : c < 0xe0 ? 2 : c < 0xf0 ? 3 : 4; + chars.emplace_back(text.data() + i, std::min(length, text.size() - i)); + i += length; + } + for (size_t i = 0; i < chars.size();) + { + size_t end = i + 1; + while (end < chars.size() && chars[end] == chars[i]) + ++end; + filtered.insert(filtered.end(), chars.begin() + i, chars.begin() + (end - i > 20 ? i + 1 : end)); + i = end; + } + std::string result; + for (size_t i = 0; i < filtered.size();) + { + size_t advance = 1, keep = 1; + if (filtered.size() - i >= 40) + for (size_t length = 1; length <= 20 && i + length * 20 <= filtered.size(); ++length) + { + size_t end = i + length; + while (end + length <= filtered.size() && + std::equal(filtered.begin() + i, filtered.begin() + i + length, filtered.begin() + end)) + end += length; + if ((end - i) / length >= 20) + { + keep = length; + advance = end - i; + break; + } + } + for (size_t j = 0; j < keep; ++j) + result += filtered[i + j]; + i += advance; + } + return result; +} + +std::pair ParseOutput(const std::string& raw, const std::string& forced_language = {}) +{ + const auto text = FixRepetitions(Trim(raw)); + if (text.empty()) + return {}; + if (!forced_language.empty()) + return {forced_language, text}; + const auto marker = text.find(""); + if (marker == std::string::npos) + return {{}, Trim(text)}; + auto meta = text.substr(0, marker); + std::transform(meta.begin(), meta.end(), meta.begin(), + [](unsigned char c) + { + return std::tolower(c); + }); + const auto transcript = Trim(text.substr(marker + 10)); + if (meta.find("language none") != std::string::npos) + return {{}, transcript}; + std::istringstream lines(meta); + std::string line; + while (std::getline(lines, line)) + { + line = Trim(line); + if (line.starts_with("language ")) + { + auto language = Trim(line.substr(9)); + if (!language.empty()) + language[0] = static_cast(std::toupper(static_cast(language[0]))); + return {language, transcript}; + } + } + return {{}, transcript}; +} + +std::string AsrLanguage(const std::string& language, const nlohmann::json& languages) +{ + const auto hint = LowerLanguage(language); + for (const auto& [key, value] : languages.items()) + if (LowerLanguage(key) == hint) + return key; + throw std::invalid_argument("Unknown ASR language hint; use an upstream language code/name or auto"); +} + +} // namespace din::asr::qwen3::detail + +namespace din::asr::qwen3 +{ +using namespace detail; +struct Qwen3Pipeline::Impl +{ + Qwen3Config config; + Runtime runtime; + AudioModel asr; + std::unique_ptr text, prefill_session; + int64_t capacity = 0; + std::vector eos; + std::vector cache; + std::unique_ptr step, prefill; + Ort::Value logits{nullptr}; + std::unique_ptr> next_token; + std::unique_ptr decode_binding, prefill_binding; + Ort::RunOptions decode_options; + + explicit Impl(Qwen3Config cfg) + : config(std::move(cfg)) + , runtime(config.provider, config.ep_cache_dir, config.ep_context_dir, config.progress) + , asr(runtime, config.model_dir, "asr") + { + if (config.max_new_tokens <= 0) + throw std::runtime_error("max-new-tokens must be positive"); + config.max_chunk_seconds = detail::ResolveChunkSeconds(config.max_chunk_seconds); + const auto& meta = asr.metadata; + if (meta.at("prefill_block") != kPrefillBlock || meta.at("cache_capacity").get() < kPrefillBlock || + meta.at("cache_capacity").get() % kPrefillBlock) + throw std::runtime_error("Unsupported Qwen3 cache geometry"); + eos = asr.metadata["eos_token_ids"].get>(); + config.lang_id = detail::AsrLanguage(config.lang_id, asr.native["languages"]); + const auto cap = meta["cache_capacity"].get(); + const auto decode_shape = TextShape(1, asr.hidden, cap) + ",logits_index:1"; + const auto prefill_shape = TextShape(kPrefillBlock, asr.hidden, cap) + ",logits_index:1"; + text = runtime.Runner(config.model_dir, "decoder", decode_shape, decode_shape, decode_shape); + prefill_session = + runtime.Runner(config.model_dir, "decoder", prefill_shape, prefill_shape, prefill_shape, "_prefill"); + if (text->HasDeviceIo()) + decode_options.AddConfigEntry("disable_synchronize_execution_providers", "1"); + runtime.LoadMel(config.model_dir); + } + + void PrepareCache() + { + if (!cache.empty()) + return; + capacity = asr.metadata.at("cache_capacity"); + const auto& c = asr.metadata["text_config"]; + const std::vector shape{1, c["num_key_value_heads"], capacity, c["head_dim"]}; + for (int i = 0; i < 2 * c["num_hidden_layers"].get(); ++i) + cache.push_back(TensorValue(*text, shape, asr.dtype)); + step = std::make_unique(*text, 1, asr.hidden, capacity, asr.dtype); + prefill = std::make_unique(*prefill_session, kPrefillBlock, asr.hidden, capacity, asr.dtype); + logits = TensorValue(*text, {1, c["vocab_size"]}, asr.dtype); + next_token = std::make_unique>(*text, std::vector{1}, text->HasDeviceIo()); + decode_binding = std::make_unique(text->session); + prefill_binding = std::make_unique(prefill_session->session); + step->Bind(*decode_binding, true); + prefill->Bind(*prefill_binding, true); + BindDecode(*decode_binding); + BindDecode(*prefill_binding); + } + + void BindDecode(Ort::IoBinding& binding) + { + for (size_t i = 0; i < cache.size(); ++i) + { + binding.BindInput(("past_" + std::to_string(i)).c_str(), cache[i]); + binding.BindOutput(("present_" + std::to_string(i)).c_str(), cache[i]); + } + binding.BindOutput("logits", logits); + binding.BindOutput("next_token", next_token->BindingValue()); + } + + TranscriptionResult TranscribeChunk(std::span audio) + { + if (audio.empty() || audio.size() > kMaxSamples) + throw std::runtime_error("Expected nonempty mono 16 kHz audio, at most 1205 seconds per chunk"); + auto features = runtime.Features(audio); + const int64_t frames = features.size() / 128; + auto encoded = asr.Encode(features, frames); + std::vector ids = asr.native["prefixes"][config.lang_id]; + ids.insert(ids.end(), encoded.tokens, asr.audio_id); + if (!asr.native.contains("suffixes")) + throw std::runtime_error("Re-export native prompt assets with --only mel for official language forcing"); + Append(ids, asr.native["suffixes"][config.lang_id].get>()); + PrepareCache(); + if (ids.size() + config.max_new_tokens > static_cast(capacity)) + throw std::runtime_error("Prompt and generation budget exceed the exported KV capacity"); + // Clear on the shared CUDA stream. Unused NaN cache values can poison attention even when masked. + for (auto& value : cache) + ZeroTensor(*text, value); + int64_t audio_offset = 0; + for (size_t offset = 0; offset < ids.size(); offset += kPrefillBlock) + { + din::common::nvtx_scoped_range range{"qwen3.prefill"}; + const auto block = + std::span(ids).subspan(offset, std::min(kPrefillBlock, ids.size() - offset)); + prefill->Fill(block, offset, asr.audio_id, &encoded, audio_offset); + audio_offset += std::count(block.begin(), block.end(), asr.audio_id); + prefill_session->session.Run(Ort::RunOptions{}, *prefill_binding); + prefill_binding->SynchronizeOutputs(); + } + int64_t position = ids.size(); + TranscriptionResult result; + for (int count = 0; count < config.max_new_tokens; ++count) + { + din::common::nvtx_scoped_range range{"qwen3.decode_step"}; + next_token->CopyAsyncToHostWithNotification().Sync(); + const auto token = next_token->HostData()[0]; + result.tokens.push_back(token); + if (std::find(eos.begin(), eos.end(), token) != eos.end()) + { + result.reached_eos = true; + break; + } + if (position >= capacity || count + 1 == config.max_new_tokens) + break; + step->Fill(std::span(&token, 1), position++, asr.audio_id, nullptr); + text->session.Run(decode_options, *decode_binding); + } + auto text_tokens = result.tokens; + if (result.reached_eos) + text_tokens.pop_back(); + const auto parsed = detail::ParseOutput(asr.tokenizer->Decode(text_tokens, false), + asr.native["languages"][config.lang_id].get()); + result.language = parsed.first; + result.text = parsed.second; + return result; + } + + TranscriptionResult Transcribe(const din::io::Audio& audio) + { + din::common::nvtx_scoped_range range{"qwen3.transcribe"}; + const auto start = std::chrono::steady_clock::now(); + din::io::Audio normalized; + const auto* source = &NormalizeAudio(audio, normalized); + if (config.progress) + config.progress({din::common::ProgressStage::Transcribing, "Transcribing", 0, audio.Duration()}); + TranscriptionResult result; + result.reached_eos = true; + std::string previous_language; + const int64_t prompt = + asr.native["prefixes"][config.lang_id].size() + asr.native["suffixes"][config.lang_id].size(); + const int64_t available = asr.metadata["cache_capacity"].get() - prompt - config.max_new_tokens; + const int seconds = static_cast(available / 13); + if (seconds < 1) + throw std::invalid_argument( + "Generation budget leaves no room for audio; increase cache-capacity or reduce max-new-tokens"); + // Preserve upstream boundaries when possible; reserve the +5s quiet search margin. + const int target = std::min(config.max_chunk_seconds, seconds - 5); + if (source->samples.size() > static_cast(seconds) * kRate && target < 6) + throw std::invalid_argument("KV capacity is too small for long-form quiet-boundary splitting"); + const auto chunks = source->samples.size() <= static_cast(seconds) * kRate && target < 6 + ? std::vector{{0, source->samples.size()}} + : detail::SplitAudio(source->samples, std::max(6, target)); + for (const auto chunk : chunks) + { + auto part = TranscribeChunk(std::span(source->samples).subspan(chunk.begin, chunk.end - chunk.begin)); + result.text += part.text; // Upstream joins literally, without overlap or inserted separators. + Append(result.tokens, part.tokens); + if (!part.language.empty() && part.language != previous_language) + { + if (!result.language.empty()) + result.language += ","; + result.language += part.language; + previous_language = part.language; + } + result.segments.push_back({std::move(part.text), std::move(part.language), chunk.begin, chunk.end}); + ++result.chunks_processed; + result.reached_eos = result.reached_eos && part.reached_eos; + if (config.progress) + config.progress( + {din::common::ProgressStage::Transcribing, + "Completed chunk " + std::to_string(result.chunks_processed), + chunk.end == source->samples.size() ? audio.Duration() : static_cast(chunk.end) / kRate, + audio.Duration()}); + } + result.audio_seconds = static_cast(audio.Duration()); + result.transcribe_seconds = std::chrono::duration(std::chrono::steady_clock::now() - start).count(); + return result; + } +}; + +Qwen3Pipeline::Qwen3Pipeline(Qwen3Config config) + : impl_(std::make_unique(std::move(config))) +{ +} +Qwen3Pipeline::~Qwen3Pipeline() = default; +TranscriptionResult Qwen3Pipeline::Transcribe(const din::io::Audio& audio) +{ + return impl_->Transcribe(audio); +} +TranscriptionResult Qwen3Pipeline::TranscribeFile(const std::filesystem::path& path) +{ + if (impl_->config.progress) + impl_->config.progress({din::common::ProgressStage::DecodingAudio, path.filename().string()}); + return Transcribe(din::io::LoadAudio(path, kRate)); +} +} // namespace din::asr::qwen3 diff --git a/asr/qwen3/qwen3.h b/asr/qwen3/qwen3.h new file mode 100644 index 0000000..e4a8a1c --- /dev/null +++ b/asr/qwen3/qwen3.h @@ -0,0 +1,64 @@ +// SPDX-License-Identifier: Apache-2.0 +#pragma once + +#include +#include +#include +#include +#include + +#include "audio.h" +#include "progress.h" + +namespace din::asr::qwen3 +{ +struct Qwen3Config +{ + std::string provider = "trt-rtx"; + std::filesystem::path model_dir = "artifacts/qwen3/onnx-bf16"; + std::filesystem::path ep_cache_dir = "artifacts/qwen3/rt_cache"; + std::filesystem::path ep_context_dir = "artifacts/qwen3/ep_context"; + std::string lang_id = "auto"; + int max_new_tokens = 1024; + int max_chunk_seconds = 0; // Auto: 1200 s, limited by KV capacity; boundaries add up to 5 s. + din::common::ProgressCallback progress; +}; + +struct TranscriptionSegment +{ + std::string text; + std::string language; + // Half-open sample offsets in normalized 16 kHz audio, not word timestamps. + size_t start_sample = 0; + size_t end_sample = 0; +}; + +struct TranscriptionResult +{ + std::string text; + std::string language; + std::vector tokens; + std::vector segments; + bool reached_eos = false; + size_t chunks_processed = 0; + float audio_seconds = 0; + float transcribe_seconds = 0; +}; + +// One synchronous stream per instance. Reuse the instance across files. +class Qwen3Pipeline +{ +public: + explicit Qwen3Pipeline(Qwen3Config config); + ~Qwen3Pipeline(); + Qwen3Pipeline(const Qwen3Pipeline&) = delete; + Qwen3Pipeline& operator=(const Qwen3Pipeline&) = delete; + TranscriptionResult Transcribe(const din::io::Audio& audio); + TranscriptionResult TranscribeFile(const std::filesystem::path& path); + +private: + struct Impl; + std::unique_ptr impl_; +}; + +} // namespace din::asr::qwen3 diff --git a/asr/qwen3/requirements.txt b/asr/qwen3/requirements.txt new file mode 100644 index 0000000..23c9a8a --- /dev/null +++ b/asr/qwen3/requirements.txt @@ -0,0 +1,13 @@ +# Tested export/validation environment (Python 3.12). +# Install this directly: pip install -r asr/qwen3/requirements.txt +# For BF16 GPU validation, install the CUDA PyTorch wheel as described in README. +torch==2.14.0 +transformers==5.13.0 +onnx==1.22.0 +onnxscript==0.7.2 +onnxruntime==1.30.0 +onnxruntime-ep-nv-tensorrt-rtx==0.4.0 +numpy==2.5.3 +scipy==1.18.1 +soundfile==0.14.0 +librosa==1.0.0 diff --git a/cmake/pcre2.cmake b/cmake/pcre2.cmake new file mode 100644 index 0000000..215d138 --- /dev/null +++ b/cmake/pcre2.cmake @@ -0,0 +1,16 @@ +# SPDX-License-Identifier: Apache-2.0 +include_guard(GLOBAL) +include(FetchContent) + +set(PCRE2_BUILD_PCRE2_8 ON CACHE BOOL "" FORCE) +set(PCRE2_BUILD_PCRE2_16 OFF CACHE BOOL "" FORCE) +set(PCRE2_BUILD_PCRE2_32 OFF CACHE BOOL "" FORCE) +set(PCRE2_BUILD_PCRE2GREP OFF CACHE BOOL "" FORCE) +set(PCRE2_BUILD_TESTS OFF CACHE BOOL "" FORCE) +set(PCRE2_STATIC_PIC ON CACHE BOOL "" FORCE) +FetchContent_Declare(pcre2 + GIT_REPOSITORY https://github.com/PCRE2Project/pcre2.git + GIT_TAG pcre2-10.46 + GIT_SHALLOW TRUE +) +FetchContent_MakeAvailable(pcre2) diff --git a/cmake/utf8proc.cmake b/cmake/utf8proc.cmake new file mode 100644 index 0000000..c4addbd --- /dev/null +++ b/cmake/utf8proc.cmake @@ -0,0 +1,12 @@ +# SPDX-License-Identifier: Apache-2.0 +include_guard(GLOBAL) +include(FetchContent) + +set(UTF8PROC_ENABLE_TESTING OFF CACHE BOOL "" FORCE) +FetchContent_Declare(utf8proc + GIT_REPOSITORY https://github.com/JuliaStrings/utf8proc.git + GIT_TAG v2.11.0 + GIT_SHALLOW TRUE +) +FetchContent_MakeAvailable(utf8proc) +set_target_properties(utf8proc PROPERTIES POSITION_INDEPENDENT_CODE ON) diff --git a/common/io/CMakeLists.txt b/common/io/CMakeLists.txt index e84b9a9..e613f5a 100644 --- a/common/io/CMakeLists.txt +++ b/common/io/CMakeLists.txt @@ -1,9 +1,12 @@ include(${PROJECT_SOURCE_DIR}/cmake/lodepng.cmake) include(${PROJECT_SOURCE_DIR}/cmake/miniaudio.cmake) +include(${PROJECT_SOURCE_DIR}/cmake/pcre2.cmake) +include(${PROJECT_SOURCE_DIR}/cmake/utf8proc.cmake) add_library(din_common_io STATIC image.cpp tokenizer.cpp + unicode_regex.cpp ) target_include_directories(din_common_io @@ -16,6 +19,8 @@ target_link_libraries(din_common_io PRIVATE lodepng nlohmann_json::nlohmann_json + pcre2-8-static + utf8proc ) set_target_properties(din_common_io PROPERTIES diff --git a/common/io/tokenizer.cpp b/common/io/tokenizer.cpp index 86dcc97..46c7671 100644 --- a/common/io/tokenizer.cpp +++ b/common/io/tokenizer.cpp @@ -7,15 +7,19 @@ #include #include #include +#include #include #include +#include #include #include #include #include #include +#include "unicode_regex.h" #include +#include namespace din::io { @@ -103,6 +107,8 @@ struct JsonTokenizerData std::unordered_map token_to_id; std::unordered_map bpe_ranks; std::vector special_tokens; + std::vector added_tokens; + bool normalize_nfc = false; }; void AddToken(JsonTokenizerData& data, const std::string& token, int64_t id) @@ -133,6 +139,8 @@ JsonTokenizerData LoadJsonTokenizerData(const std::string& path) } JsonTokenizerData tokenizer_data; + tokenizer_data.normalize_nfc = + data.contains("normalizer") && data["normalizer"].is_object() && data["normalizer"].value("type", "") == "NFC"; const auto& vocab = data["model"]["vocab"]; for (const auto& [token, id_val] : vocab.items()) @@ -150,6 +158,7 @@ JsonTokenizerData LoadJsonTokenizerData(const std::string& path) } const auto content = token["content"].get(); AddToken(tokenizer_data, content, token["id"].get()); + tokenizer_data.added_tokens.push_back(content); if (token.value("special", false)) { tokenizer_data.special_tokens.push_back(content); @@ -194,6 +203,11 @@ JsonTokenizerData LoadJsonTokenizerData(const std::string& path) { return left.size() > right.size(); }); + std::sort(tokenizer_data.added_tokens.begin(), tokenizer_data.added_tokens.end(), + [](const auto& left, const auto& right) + { + return left.size() > right.size(); + }); return tokenizer_data; } @@ -347,14 +361,17 @@ Tokenizer::Tokenizer(const std::string& path, TokenizerFormat format) switch (format) { case TokenizerFormat::Json: + case TokenizerFormat::ByteBpeJson: { auto tokenizer_data = LoadJsonTokenizerData(path); id_to_token_ = std::move(tokenizer_data.id_to_token); token_to_id_ = std::move(tokenizer_data.token_to_id); bpe_ranks_ = std::move(tokenizer_data.bpe_ranks); special_tokens_ = std::move(tokenizer_data.special_tokens); + added_tokens_ = std::move(tokenizer_data.added_tokens); + normalize_nfc_ = format == TokenizerFormat::ByteBpeJson && tokenizer_data.normalize_nfc; byte_encoder_ = BuildByteEncoder(); - decode_mode_ = DecodeMode::Pieces; + decode_mode_ = format == TokenizerFormat::ByteBpeJson ? DecodeMode::ByteBpe : DecodeMode::Pieces; break; } case TokenizerFormat::Vocab: @@ -434,8 +451,14 @@ bool IsContractionAt(std::string_view text, size_t pos, size_t* length) return false; } -std::vector SplitForByteLevelBpe(std::string_view text) +std::vector SplitForByteLevelBpe(std::string_view text, bool qwen = false) { + if (qwen) + { + static const UnicodeRegex pattern( + R"((?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\r\n\p{L}\p{N}]?\p{L}+|\p{N}| ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+)"); + return pattern.FindAll(text); + } std::vector pieces; size_t pos = 0; while (pos < text.size()) @@ -593,7 +616,7 @@ int64_t Tokenizer::TokenId(const std::string& token) const return it->second; } -std::vector Tokenizer::Encode(const std::string& text, bool add_special_tokens) const +std::vector Tokenizer::Encode(const std::string& input, bool add_special_tokens) const { if (bpe_ranks_.empty()) { @@ -601,10 +624,23 @@ std::vector Tokenizer::Encode(const std::string& text, bool add_special } (void)add_special_tokens; + std::string normalized; + if (normalize_nfc_) + { + utf8proc_uint8_t* data = nullptr; + const auto length = utf8proc_map(reinterpret_cast(input.data()), input.size(), &data, + static_cast(UTF8PROC_STABLE | UTF8PROC_COMPOSE)); + const std::unique_ptr owner(data, std::free); + if (length < 0) + throw std::invalid_argument(utf8proc_errmsg(length)); + normalized.assign(reinterpret_cast(data), static_cast(length)); + } + const auto& text = normalize_nfc_ ? normalized : input; + std::vector ids; const auto encode_segment = [&](std::string_view segment) { - for (const auto& piece : SplitForByteLevelBpe(segment)) + for (const auto& piece : SplitForByteLevelBpe(segment, decode_mode_ == DecodeMode::ByteBpe)) { for (const auto& token : ApplyByteLevelBpe(piece, byte_encoder_, bpe_ranks_)) { @@ -623,7 +659,7 @@ std::vector Tokenizer::Encode(const std::string& text, bool add_special while (pos < text.size()) { const std::string* matched_special = nullptr; - for (const auto& special : special_tokens_) + for (const auto& special : decode_mode_ == DecodeMode::ByteBpe ? added_tokens_ : special_tokens_) { if (StartsWith(std::string_view(text).substr(pos), special)) { @@ -667,6 +703,27 @@ std::string Tokenizer::CleanToken(int64_t id) const std::string Tokenizer::Decode(const std::vector& ids, bool skip_special_tokens, bool strip_lang_tags) const { + if (decode_mode_ == DecodeMode::ByteBpe) + { + const auto byte_decoder = BuildByteDecoder(); + std::string text; + for (auto id : ids) + { + const auto& token = Token(id); + if (skip_special_tokens && + std::find(special_tokens_.begin(), special_tokens_.end(), token) != special_tokens_.end()) + continue; + for (size_t i = 0; i < token.size();) + { + const auto code = NextCodePoint(token, i); + const auto byte = byte_decoder.find(code); + if (byte == byte_decoder.end()) + throw std::runtime_error("Invalid byte-level vocabulary entry"); + text.push_back(static_cast(byte->second)); + } + } + return text; + } if (decode_mode_ == DecodeMode::WhisperByteBpe) { std::string text; diff --git a/common/io/tokenizer.h b/common/io/tokenizer.h index eec5b97..c79ac5d 100644 --- a/common/io/tokenizer.h +++ b/common/io/tokenizer.h @@ -17,6 +17,7 @@ enum class TokenizerFormat Json, Vocab, WhisperJson, + ByteBpeJson, // Qwen byte-level decoding and Unicode encoding. }; inline constexpr int64_t kWhisperEndOfText = 50257; @@ -48,6 +49,7 @@ class Tokenizer enum class DecodeMode { Pieces, + ByteBpe, WhisperByteBpe, }; @@ -56,6 +58,8 @@ class Tokenizer std::unordered_map bpe_ranks_; std::array byte_encoder_; std::vector special_tokens_; + std::vector added_tokens_; + bool normalize_nfc_ = false; std::vector lang_codes_; DecodeMode decode_mode_ = DecodeMode::Pieces; }; diff --git a/common/io/unicode_regex.cpp b/common/io/unicode_regex.cpp new file mode 100644 index 0000000..f0adb66 --- /dev/null +++ b/common/io/unicode_regex.cpp @@ -0,0 +1,51 @@ +// SPDX-License-Identifier: Apache-2.0 +#include "unicode_regex.h" + +#include +#include + +#define PCRE2_CODE_UNIT_WIDTH 8 +#include + +namespace din::io +{ +UnicodeRegex::UnicodeRegex(std::string_view pattern) +{ + int error; + PCRE2_SIZE offset; + code_ = pcre2_compile(reinterpret_cast(pattern.data()), pattern.size(), PCRE2_UTF | PCRE2_UCP, &error, + &offset, nullptr); + if (!code_) + throw std::runtime_error("Invalid Unicode regex at byte " + std::to_string(offset)); +} + +UnicodeRegex::~UnicodeRegex() +{ + pcre2_code_free(code_); +} + +std::vector UnicodeRegex::FindAll(std::string_view text) const +{ + std::vector matches; + const std::unique_ptr data( + pcre2_match_data_create_from_pattern(code_, nullptr), pcre2_match_data_free); + if (!data) + throw std::bad_alloc(); + size_t offset = 0; + while (offset < text.size()) + { + const int result = pcre2_match(code_, reinterpret_cast(text.data()), text.size(), offset, + offset ? PCRE2_NO_UTF_CHECK : 0, data.get(), nullptr); + if (result == PCRE2_ERROR_NOMATCH) + break; + if (result < 0) + throw std::invalid_argument("Invalid UTF-8 text or Unicode regex match failure: " + std::to_string(result)); + const auto* span = pcre2_get_ovector_pointer(data.get()); + if (span[1] <= offset) + throw std::runtime_error("Unicode regex must consume input"); + matches.emplace_back(text.substr(span[0], span[1] - span[0])); + offset = span[1]; + } + return matches; +} +} // namespace din::io diff --git a/common/io/unicode_regex.h b/common/io/unicode_regex.h new file mode 100644 index 0000000..7bbc463 --- /dev/null +++ b/common/io/unicode_regex.h @@ -0,0 +1,25 @@ +// SPDX-License-Identifier: Apache-2.0 +#pragma once + +#include +#include +#include + +struct pcre2_real_code_8; + +namespace din::io +{ +// Compiled once; matching uses per-call state and validates UTF-8. +class UnicodeRegex +{ +public: + explicit UnicodeRegex(std::string_view pattern); + ~UnicodeRegex(); + UnicodeRegex(const UnicodeRegex&) = delete; + UnicodeRegex& operator=(const UnicodeRegex&) = delete; + std::vector FindAll(std::string_view text) const; + +private: + pcre2_real_code_8* code_; +}; +} // namespace din::io diff --git a/common/ort_session.cpp b/common/ort_session.cpp index 43a3a26..fbcce7c 100644 --- a/common/ort_session.cpp +++ b/common/ort_session.cpp @@ -399,51 +399,43 @@ bool IsCudaUnifiedMemoryDevice(Ort::ConstEpDevice ep_device) bool RegisterTensorRTRTXExecutionProvider(Ort::Env& env) { - static std::once_flag registration_once; - static Ort::Env* registered_env = nullptr; - std::call_once(registration_once, - [&env] - { - auto provider_library = std::filesystem::path{ONNXRUNTIME_TRT_RTX_EP_LIBRARY_PATH}; - if (!std::filesystem::is_regular_file(provider_library)) - { + static std::mutex registration_mutex; + const std::lock_guard lock(registration_mutex); + // Check the native environment, not the address of its C++ wrapper. + for (const auto& device : env.GetEpDevices()) + if (std::string_view{device.EpName()} == kDinNvTensorRTRTXExecutionProvider) + return true; + auto provider_library = std::filesystem::path{ONNXRUNTIME_TRT_RTX_EP_LIBRARY_PATH}; + if (!std::filesystem::is_regular_file(provider_library)) + { #ifdef _WIN32 - provider_library = ExecutableDirectory() / "onnxruntime_providers_nv_tensorrt_rtx.dll"; + provider_library = ExecutableDirectory() / "onnxruntime_providers_nv_tensorrt_rtx.dll"; #else - provider_library = ExecutableDirectory() / "libonnxruntime_providers_nv_tensorrt_rtx.so"; + provider_library = ExecutableDirectory() / "libonnxruntime_providers_nv_tensorrt_rtx.so"; #endif - if (!std::filesystem::is_regular_file(provider_library)) - { - throw std::runtime_error("TensorRT RTX execution provider library not found: " + - provider_library.string()); - } - } + if (!std::filesystem::is_regular_file(provider_library)) + { + throw std::runtime_error("TensorRT RTX execution provider library not found: " + provider_library.string()); + } + } #ifdef _WIN32 - const auto provider_directory = provider_library.parent_path().wstring(); - if (SetDllDirectoryW(provider_directory.c_str()) == 0) - { - throw std::runtime_error("Failed to add TensorRT RTX EP directory to the DLL search path."); - } + const auto provider_directory = provider_library.parent_path().wstring(); + if (SetDllDirectoryW(provider_directory.c_str()) == 0) + { + throw std::runtime_error("Failed to add TensorRT RTX EP directory to the DLL search path."); + } #endif - const auto provider_library_path = ToOrtPathString(provider_library); - env.RegisterExecutionProviderLibrary(kDinNvTensorRTRTXExecutionProvider, - provider_library_path.c_str()); - - const auto ep_devices = env.GetEpDevices(); - std::cout << "Execution provider devices after TRT RTX registration:\n"; - for (const auto& device : ep_devices) - { - std::cout << " " << device.EpName() << " vendor=" << device.EpVendor() - << " device_id=" << device.Device().DeviceId() << '\n'; - } - registered_env = &env; - }); + const auto provider_library_path = ToOrtPathString(provider_library); + env.RegisterExecutionProviderLibrary(kDinNvTensorRTRTXExecutionProvider, provider_library_path.c_str()); - if (registered_env != &env) + const auto ep_devices = env.GetEpDevices(); + std::cout << "Execution provider devices after TRT RTX registration:\n"; + for (const auto& device : ep_devices) { - throw std::logic_error("TensorRT RTX was already registered on a different Ort::Env in this process."); + std::cout << " " << device.EpName() << " vendor=" << device.EpVendor() + << " device_id=" << device.Device().DeviceId() << '\n'; } return true; } @@ -564,7 +556,8 @@ std::string CompileEpContextModel(Ort::Env& env, const std::string& model_path, const auto output_model_path = CompiledModelPath(model_path, ep_context, profile); if (fs::exists(output_model_path)) { - if (IsCompatibleEpContext(env, output_model_path)) + if (fs::last_write_time(output_model_path) >= fs::last_write_time(model_path) && + IsCompatibleEpContext(env, output_model_path)) { return output_model_path; } @@ -595,6 +588,8 @@ std::string CompileEpContextModel(Ort::Env& env, const std::string& model_path, const size_t size_threshold_external_init = 1024; compile_options.SetOutputModelExternalInitializersFile(output_external_wide.c_str(), size_threshold_external_init); + if (ep_context.progress) + ep_context.progress({ProgressStage::CompilingModel, fs::path(model_path).filename().string()}); const Ort::Status status = Ort::CompileModel(env, compile_options); if (!status.IsOK()) { @@ -788,6 +783,8 @@ OrtRunner::OrtRunner(Ort::Env& env_in, const std::string& model_path, const std: { DIN_NVTX_FUNC_RANGE(); auto session_model_path = model_path; + if (ep_context.progress) + ep_context.progress({ProgressStage::LoadingModel, fs::path(model_path).filename().string()}); if (provider == "trt-rtx") { @@ -825,6 +822,8 @@ OrtRunner::OrtRunner(Ort::Env& env_in, const std::string& model_path, const std: throw std::runtime_error("unsupported provider: " + provider); } + if (ep_context.progress) + ep_context.progress({ProgressStage::LoadingModel, fs::path(session_model_path).filename().string()}); #ifdef _WIN32 const auto wide = ToOrtPathString(session_model_path); { diff --git a/common/ort_session.h b/common/ort_session.h index 127a599..2fca6c1 100644 --- a/common/ort_session.h +++ b/common/ort_session.h @@ -13,6 +13,7 @@ #include #include +#include "progress.h" #include #include #include @@ -68,6 +69,7 @@ QueryCudaGraphicsInteropSharedMemoryInfo(int cuda_device_ordinal, bool high_prio struct EpContextOptions { std::string output_dir; + ProgressCallback progress; }; struct ModelProfile diff --git a/common/progress.h b/common/progress.h new file mode 100644 index 0000000..7c65d66 --- /dev/null +++ b/common/progress.h @@ -0,0 +1,26 @@ +// SPDX-License-Identifier: Apache-2.0 +#pragma once +#include +#include + +namespace din::common +{ +enum class ProgressStage +{ + DecodingAudio, + LoadingModel, + CompilingModel, + Transcribing, + Aligning +}; +struct InferenceProgress +{ + ProgressStage stage; + std::string detail; + double completed_audio_seconds = 0; + double total_audio_seconds = 0; +}; +// Invoked synchronously on the pipeline's calling thread. Keep callbacks short. +// Optional: existing CLI callers need not install a callback. +using ProgressCallback = std::function; +} // namespace din::common From 4541203102cded26a70f68b3221c5538c57dd04a Mon Sep 17 00:00:00 2001 From: contentis Date: Wed, 23 Sep 2026 10:46:17 +0200 Subject: [PATCH 2/8] POC: extend Qwen streaming and alignment; defer progress callbacks --- THIRD_PARTY_NOTICES.md | 3 + asr/qwen3/CMakeLists.txt | 1 + asr/qwen3/README.md | 52 ++++- asr/qwen3/aligner_main.cpp | 22 +- asr/qwen3/detail/japanese.h | 192 ++++++++++++++++++ asr/qwen3/detail/runtime.h | 6 +- asr/qwen3/detail/text.h | 22 ++ asr/qwen3/forced_aligner.cpp | 109 +++++++--- asr/qwen3/forced_aligner.h | 21 +- asr/qwen3/main.cpp | 64 +++++- asr/qwen3/model_export/detail/japanese.py | 149 ++++++++++++++ .../model_export/detail/nagisa.LICENSE.txt | 21 ++ asr/qwen3/model_export/export_qwen3_asr.py | 10 +- asr/qwen3/qwen3.cpp | 108 ++++++++-- asr/qwen3/qwen3.h | 24 ++- common/ort_session.cpp | 6 - common/ort_session.h | 2 - common/progress.h | 26 --- 18 files changed, 736 insertions(+), 102 deletions(-) create mode 100644 asr/qwen3/detail/japanese.h create mode 100644 asr/qwen3/detail/text.h create mode 100644 asr/qwen3/model_export/detail/japanese.py create mode 100644 asr/qwen3/model_export/detail/nagisa.LICENSE.txt delete mode 100644 common/progress.h diff --git a/THIRD_PARTY_NOTICES.md b/THIRD_PARTY_NOTICES.md index e211c3b..0a48691 100644 --- a/THIRD_PARTY_NOTICES.md +++ b/THIRD_PARTY_NOTICES.md @@ -6,6 +6,7 @@ DIN Deploy is distributed under the Apache License, Version 2.0. The project sou |---|---| | argparse | MIT: | | miniaudio | MIT: | +| Nagisa word segmentation | MIT: | | lodepng | zlib: | | nlohmann/json | MIT: | | PCRE2 | BSD-3-Clause WITH PCRE2-exception: | @@ -80,6 +81,8 @@ For components supplied through an SDK or binary package, the corresponding vend ## Model/Artifact +- Nagisa `nagisa_v001` word segmenter: MIT: + - `black-forest-labs/FLUX.2-klein-4b`: Apache-2.0 - `black-forest-labs/FLUX.2-klein-4b-fp8`: Apache-2.0 - `black-forest-labs/FLUX.2-klein-4b-nvfp4`: Apache-2.0 diff --git a/asr/qwen3/CMakeLists.txt b/asr/qwen3/CMakeLists.txt index 3426310..c084c7b 100644 --- a/asr/qwen3/CMakeLists.txt +++ b/asr/qwen3/CMakeLists.txt @@ -2,6 +2,7 @@ add_din_shared_library(din_asr_qwen3 STATIC qwen3.cpp forced_aligner.cpp) target_include_directories(din_asr_qwen3 PUBLIC "${CMAKE_CURRENT_SOURCE_DIR}") target_link_libraries(din_asr_qwen3 PUBLIC din_common_ort din_common_io nlohmann_json::nlohmann_json) +target_link_libraries(din_asr_qwen3 PRIVATE utf8proc) find_package(CUDAToolkit QUIET) if(CUDAToolkit_FOUND) target_link_libraries(din_asr_qwen3 PRIVATE CUDA::cudart) diff --git a/asr/qwen3/README.md b/asr/qwen3/README.md index 48342fe..6229c14 100644 --- a/asr/qwen3/README.md +++ b/asr/qwen3/README.md @@ -1,6 +1,6 @@ # Qwen3 ASR and forced alignment -Offline C++ inference with CPU or TensorRT RTX. Use FP32 exports for CPU; TensorRT RTX supports BF16 (default), FP16 and FP32. +C++ inference with CPU or TensorRT RTX. Use FP32 exports for CPU; TensorRT RTX supports BF16 (default), FP16 and FP32. ## Supported models @@ -15,24 +15,28 @@ Offline C++ inference with CPU or TensorRT RTX. Use FP32 exports for CPU; Tensor | Model / upstream toolkit capability | C++ sample | |---|:---:| | Offline, single stream | ✓ | -| Online / streaming | — | +| Online / streaming, single utterance | ✓ | | Batched inference | — | | Long-form audio | ✓ | | ASR with / without forced alignment | ✓ | | Standalone alignment of supplied text | ✓ | +| Long-form alignment with timed transcript segments (C++ API) | ✓ | +| Long-form alignment of unsegmented text | — | | Automatic language identification / language hint | ✓ | | Multilingual ASR: 30 languages and 22 Chinese dialects | ✓ | | Word timestamps: en, de, es, fr, it, pt, ru, ko | ✓ | | Chinese / Cantonese character timestamps | ✓ | -| Japanese character timestamps | ✓ | +| Japanese word timestamps (Nagisa) | ✓ | +| Character alignment units: all 11 languages (sample extension) | ✓ | +| Caller-supplied alignment units | ✓ | | All 11 upstream alignment languages | ✓ | ASR uses the [upstream model's language support](https://github.com/QwenLM/Qwen3-ASR). ASR accepts all 30 upstream language codes/names and `auto`; dialects use automatic recognition or the corresponding language hint, not separate dialect switches. Alignment supports Chinese, Cantonese, English, German, Spanish, French, Italian, -Portuguese, Russian, Korean and Japanese. Japanese uses character timestamps; -upstream's Nagisa word boundaries differ. Latin words in CJK text stay together. +Portuguese, Russian, Korean and Japanese. Japanese word boundaries use upstream +Nagisa; Chinese/Cantonese default to characters, keeping Latin words together. Language names/codes are case-insensitive. ASR's other languages require alignment to be disabled. Language coverage is not an accuracy guarantee for every dialect. @@ -63,7 +67,12 @@ HF downloads checkpoints automatically. Use `--model` for a local checkpoint or `--revision` to pin the source. Keep each export directory intact. Log-mel processing reuses the shared Whisper frontend. -Exports contain one encoder and one decoder, each with one weight file, plus the +Aligner exports also include Nagisa's small FP32 word segmenter and vocabulary. +It runs on CPU without Python. Add it to an existing aligner export with +`--task aligner --only japanese --output `; no ASR/aligner +weights or GPU engines need rebuilding. + +ASR exports contain one encoder and one decoder, each with one weight file, plus the shared log-mel graph. Prefill and token generation update one KV bank in place. The decoder uses two fixed TensorRT profiles (512-token prefill and one-token steps), compiled once and cached; audio length does not create more encoder/decoder @@ -104,6 +113,8 @@ cmake --build out\build\windows-x64 --target din_asr_qwen3_cli din_asr_qwen3_ali ```powershell out\build\windows-x64\bin\din_asr_qwen3_cli.exe audio.mp3 --model-dir D:\models\qwen3-asr-1.7b-onnx-bf16 out\build\windows-x64\bin\din_asr_qwen3_aligner_cli.exe audio.mp3 --model-dir D:\models\qwen3-aligner-onnx-bf16 --transcript transcript.txt --lang-id zh +out\build\windows-x64\bin\din_asr_qwen3_aligner_cli.exe audio.mp3 --model-dir D:\models\qwen3-aligner-onnx-bf16 --transcript transcript.txt --lang-id ja --granularity characters +out\build\windows-x64\bin\din_asr_qwen3_cli.exe audio.mp3 --model-dir D:\models\qwen3-asr-0.6b-onnx-bf16 --stream ``` Multi-configuration builds add the configuration (for example, `Release`) under `bin`. @@ -117,12 +128,33 @@ ASR and alignment are independent APIs in the same library: Include `qwen3.h` for ASR or `forced_aligner.h` for alignment. The aligner loads no ASR model; use text from Whisper, Parakeet, Nemotron, Qwen ASR or a text file. Both CLIs accept `--provider cpu|trt-rtx`, `--model-dir` and cache options independently. -Reuse instances; each processes one synchronous call at a time. Both configs expose -`progress` callbacks for loading, compilation and processing. +Reuse instances; each processes one synchronous call at a time. `Align` takes mono 16 kHz audio; `AlignFile` decodes it automatically. Alignment -accepts at most 180 seconds per call. Long recordings require matching audio/text -segments, with returned timestamps offset by each segment's start. +accepts at most 180 seconds per call. For long recordings, use +`AlignSegments(audio, segments)` with text from any ASR. Each `AlignmentSegment` +contains `text`, half-open `start_sample` / `end_sample` offsets at 16 kHz, and +`language` (default English). The method reuses the aligner and returns timestamps +relative to the full recording, in segment order. Bounds and the 180-second limit +are checked before inference; overlapping intervals are preserved without +deduplication. Each interval must contain all speech for its text. + +`--granularity characters` / `AlignmentGranularity::Characters` aligns Unicode +graphemes, retaining combining marks and omitting spaces/punctuation except apostrophes. +`AlignUnits(audio, units)` bypasses text splitting; CLI `--units` reads one unit +per transcript line. Limits: 2048 units and 8192 context tokens per call. Character +alignment is a sample extension: the model's 80 ms timestamp bins can give adjacent +characters identical times; sub-word accuracy is not guaranteed. + +Streaming uses `StartStream()`, `PushAudio(mono16k)` and `FinishStream()`. +The CLI emits replacement hypotheses as JSON lines; use `- --stream` to read +little-endian float32 mono 16 kHz PCM from stdin. `--chunk-seconds` defaults to 2 +(minimum 0.5); `--unfixed-chunks 2 --unfixed-tokens 5` matches upstream rollback. +Each update reprocesses accumulated audio using the existing engines, so latency +grows with utterance length. The exported KV ceiling still applies and overflow +raises an error; start a new stream for the next utterance. Streaming has no live +timestamps: align the final text separately. `FinishStream()` flushes the tail; +UTF-8-safe rollback also applies to that tail. ASR returns `segments` with text, detected language and half-open `start_sample` / `end_sample` offsets at 16 kHz. Its upstream quiet-boundary splitter uses a diff --git a/asr/qwen3/aligner_main.cpp b/asr/qwen3/aligner_main.cpp index 1130e79..f3e74e3 100644 --- a/asr/qwen3/aligner_main.cpp +++ b/asr/qwen3/aligner_main.cpp @@ -3,6 +3,7 @@ #include #include #include +#include #include #include "forced_aligner.h" @@ -20,6 +21,8 @@ int main(int argc, char** argv) parser.add_argument("audiofile"); parser.add_argument("--transcript").required().help("UTF-8 transcript file"); parser.add_argument("--lang-id").default_value(std::string{"English"}).help("Language code or name"); + parser.add_argument("--granularity").default_value(std::string{"words"}).choices("words", "characters"); + parser.add_argument("--units").flag().help("Treat each transcript line as one supplied alignment unit"); parser.add_argument("--provider").default_value(config.provider).choices("cpu", "trt-rtx"); parser.add_argument("--model-dir").default_value(config.model_dir.string()); parser.add_argument("--ep-cache").default_value(config.ep_cache_dir.string()); @@ -29,6 +32,8 @@ int main(int argc, char** argv) config.model_dir = parser.get("--model-dir"); config.ep_cache_dir = parser.get("--ep-cache"); config.ep_context_dir = parser.get("--ep-context-dir"); + config.granularity = parser.get("--granularity") == "characters" ? AlignmentGranularity::Characters + : AlignmentGranularity::Words; const auto transcript_path = parser.get("--transcript"); std::ifstream file(transcript_path, std::ios::binary); if (!file) @@ -40,7 +45,22 @@ int main(int argc, char** argv) const auto audio = din::io::LoadAudio(parser.get("audiofile"), 16000); Qwen3ForcedAligner aligner(std::move(config)); const auto start = std::chrono::steady_clock::now(); - const auto timestamps = aligner.Align(audio, transcript, language); + std::vector timestamps; + if (parser.get("--units")) + { + std::vector units; + std::istringstream lines(transcript); + for (std::string line; std::getline(lines, line);) + { + if (!line.empty() && line.back() == '\r') + line.pop_back(); + if (!line.empty()) + units.push_back(std::move(line)); + } + timestamps = aligner.AlignUnits(audio, units); + } + else + timestamps = aligner.Align(audio, transcript, language); const auto seconds = std::chrono::duration(std::chrono::steady_clock::now() - start).count(); nlohmann::json output{{"text", transcript}, {"language", language}, diff --git a/asr/qwen3/detail/japanese.h b/asr/qwen3/detail/japanese.h new file mode 100644 index 0000000..5345d73 --- /dev/null +++ b/asr/qwen3/detail/japanese.h @@ -0,0 +1,192 @@ +// SPDX-License-Identifier: Apache-2.0 +#pragma once +#include +#include +#include + +#include "runtime.h" +#include "text.h" +#include "unicode_regex.h" + +namespace din::asr::qwen3::detail +{ +// Nagisa preprocessing, dictionary features and BMES decoding; neural inference stays in ORT. +class JapaneseTokenizer +{ + using Vocabulary = std::unordered_map; + Vocabulary unigrams_, bigrams_, words_; + int window_; + int64_t padding_word_; + std::array, 6> transitions_; + Ort::Session session_{nullptr}; + + static std::string Utf8(int32_t code) + { + utf8proc_uint8_t data[4]; + const auto size = utf8proc_encode_char(code, data); + return {reinterpret_cast(data), static_cast(size)}; + } + static int64_t Lookup(const Vocabulary& vocabulary, const std::string& text) + { + const auto found = vocabulary.find(text); + return found == vocabulary.end() ? vocabulary.at("oov") : found->second; + } + +public: + JapaneseTokenizer(Ort::Env& env, const std::filesystem::path& dir) + { + if (!std::filesystem::exists(dir / "japanese.json")) + throw std::runtime_error("Japanese word alignment requires --task aligner --only japanese export"); + const auto meta = ReadJson(dir / "japanese.json"); + unigrams_ = meta.at("unigrams").get(); + bigrams_ = meta.at("bigrams").get(); + words_ = meta.at("words").get(); + window_ = meta.at("window"); + padding_word_ = meta.at("padding_word"); + transitions_ = meta.at("transitions").get(); + Ort::SessionOptions options; + options.SetIntraOpNumThreads(1); + session_ = Ort::Session(env, (dir / "japanese.onnx").c_str(), options); + } + + std::vector Words(std::string text) + { + if (!ValidUtf8(text)) + throw std::invalid_argument("Japanese alignment requires valid UTF-8 text"); + static const din::io::UnicodeRegex leading(R"(^[\s\x{1c}-\x{1f}]+)"), trailing(R"([\s\x{1c}-\x{1f}]+$)"); + const auto head = leading.FindAll(text); + if (!head.empty()) + text.erase(0, head[0].size()); + const auto tail = trailing.FindAll(text); + if (!tail.empty()) + text.resize(text.size() - tail[0].size()); + utf8proc_uint8_t* normalized = nullptr; + const auto length = + utf8proc_map(reinterpret_cast(text.data()), text.size(), &normalized, + static_cast(UTF8PROC_STABLE | UTF8PROC_COMPAT | UTF8PROC_COMPOSE)); + const std::unique_ptr owner(normalized, std::free); + if (length < 0) + throw std::invalid_argument(utf8proc_errmsg(length)); + std::vector characters, lower; + std::vector types; + for (size_t i = 0; i < static_cast(length);) + { + int32_t code; + i += utf8proc_iterate(normalized + i, length - i, &code); + if (code == 0x130) + code = 'I'; + if (code == ' ') + code = 0x3000; + characters.push_back(Utf8(code)); + code = utf8proc_tolower(code); + lower.push_back(Utf8(code)); + types.push_back(code >= 0x3040 && code <= 0x309f ? 0 + : code >= 0x30a1 && code <= 0x30fa ? 1 + : code >= 0x4e00 && code <= 0x9fa5 ? 2 + : code >= 'a' && code <= 'z' ? 3 + : code >= '0' && code <= '9' ? 4 + : 5); + } + const int64_t count = characters.size(); + if (!count) + return {}; + std::array, 5> features; + for (size_t i = 0; i < 3; ++i) + features[i].resize(count * window_, i == 2 ? 6 : 1); + for (size_t i = 3; i < 5; ++i) + features[i].resize(count * 8, padding_word_); + for (int64_t i = 0; i < count; ++i) + { + for (int j = 0; j < window_; ++j) + { + const auto at = i + j - window_ / 2; + if (at < 0 || at >= count) + continue; + features[0][i * window_ + j] = Lookup(unigrams_, lower[at]); + features[1][i * window_ + j] = Lookup(bigrams_, lower[at] + (at + 1 < count ? lower[at + 1] : "")); + features[2][i * window_ + j] = types[at]; + } + for (int direction = 0; direction < 2; ++direction) + { + std::string word; + int matches = 0; + for (int j = 0; j < 8; ++j) + { + const auto at = direction ? i - j : i + j; + if (at < 0 || at >= count) + break; + word = direction ? lower[at] + word : word + lower[at]; + if (const auto found = words_.find(word); found != words_.end()) + features[3 + direction][i * 8 + matches++] = found->second; + } + if (!matches) + features[3 + direction][i * 8] = words_.at("oov"); + } + } + const auto memory = Ort::MemoryInfo::CreateCpu(OrtArenaAllocator, OrtMemTypeDefault); + std::vector inputs; + for (size_t i = 0; i < features.size(); ++i) + { + const int64_t shape[]{count, i < 3 ? window_ : 8}; + inputs.push_back(Ort::Value::CreateTensor(memory, features[i].data(), features[i].size(), shape, 2)); + } + const char* names[]{"unigrams", "bigrams", "types", "word_starts", "word_ends"}; + const char* output_name = "emissions"; + auto outputs = session_.Run(Ort::RunOptions{}, names, inputs.data(), inputs.size(), &output_name, 1); + const auto* emissions = outputs[0].GetTensorData(); + std::vector> parents(count); + std::array scores; + scores.fill(-1e10f); + scores[4] = 0; + for (int64_t i = 0; i < count; ++i) + { + std::array next; + for (int to = 0; to < 6; ++to) + { + int best = 0; + for (int from = 1; from < 6; ++from) + if (scores[from] + transitions_[to][from] > scores[best] + transitions_[to][best]) + best = from; + parents[i][to] = best; + next[to] = scores[best] + transitions_[to][best] + emissions[i * 6 + to]; + } + // A common offset preserves the best path while keeping long sequences in FP32 range. + const float maximum = *std::max_element(next.begin(), next.end()); + for (int j = 0; j < 6; ++j) + scores[j] = next[j] - maximum; + } + int tag = 0; + for (int i = 1; i < 6; ++i) + if (scores[i] + transitions_[5][i] > scores[tag] + transitions_[5][tag]) + tag = i; + std::vector tags(count); + for (int64_t i = count; i-- > 0;) + { + tags[i] = tag; + tag = parents[i][tag]; + } + static const din::io::UnicodeRegex kept(R"([\p{L}\p{N}']+)"); + std::vector result; + std::string word; + auto emit = [&] + { + std::string cleaned; + for (const auto& part : kept.FindAll(word)) + cleaned += part; + if (!cleaned.empty()) + result.push_back(std::move(cleaned)); + word.clear(); + }; + for (int64_t i = 0; i < count; ++i) + { + if (tags[i] == 3) + emit(); + word += characters[i]; + if (tags[i] == 2 || tags[i] == 3) + emit(); + } + emit(); + return result; + } +}; +} // namespace din::asr::qwen3::detail diff --git a/asr/qwen3/detail/runtime.h b/asr/qwen3/detail/runtime.h index 51058f3..e539dee 100644 --- a/asr/qwen3/detail/runtime.h +++ b/asr/qwen3/detail/runtime.h @@ -354,17 +354,15 @@ struct Runtime { std::string provider; std::filesystem::path ep_cache_dir, ep_context_dir; - din::common::ProgressCallback progress; Ort::Env env{ORT_LOGGING_LEVEL_WARNING, "din_asr_qwen3"}; Ort::SyncStream stream{nullptr}; std::unique_ptr mel; Runtime(const std::string& execution_provider, const std::filesystem::path& cache, - const std::filesystem::path& context, din::common::ProgressCallback callback) + const std::filesystem::path& context) : provider(execution_provider) , ep_cache_dir(cache) , ep_context_dir(context) - , progress(std::move(callback)) { if (provider == "trt-rtx") { @@ -404,7 +402,7 @@ struct Runtime // ORT names external engines by graph, so different profiles need separate directories. return std::make_unique( env, (dir / (name + ".onnx")).string(), provider, ep_cache_dir.string(), - din::common::EpContextOptions{(ep_context_dir / profile.cache_subpath).string(), progress}, profile, + din::common::EpContextOptions{(ep_context_dir / profile.cache_subpath).string()}, profile, stream ? &stream : nullptr); } diff --git a/asr/qwen3/detail/text.h b/asr/qwen3/detail/text.h new file mode 100644 index 0000000..63c0e51 --- /dev/null +++ b/asr/qwen3/detail/text.h @@ -0,0 +1,22 @@ +// SPDX-License-Identifier: Apache-2.0 +#pragma once +#include + +#include + +namespace din::asr::qwen3::detail +{ +inline bool ValidUtf8(const std::string& text) +{ + for (size_t i = 0; i < text.size();) + { + utf8proc_int32_t code; + const auto n = + utf8proc_iterate(reinterpret_cast(text.data() + i), text.size() - i, &code); + if (n < 0) + return false; + i += n; + } + return true; +} +} // namespace din::asr::qwen3::detail diff --git a/asr/qwen3/forced_aligner.cpp b/asr/qwen3/forced_aligner.cpp index bad8cb7..958287c 100644 --- a/asr/qwen3/forced_aligner.cpp +++ b/asr/qwen3/forced_aligner.cpp @@ -1,12 +1,15 @@ // SPDX-License-Identifier: Apache-2.0 #include "forced_aligner.h" +#include "detail/japanese.h" #include "detail/runtime.h" +#include "detail/text.h" #include "unicode_regex.h" namespace din::asr::qwen3::detail { -std::vector AlignmentUnits(const std::string& text, const std::string& language) +std::vector AlignmentUnits(const std::string& text, const std::string& language, + AlignmentGranularity granularity, JapaneseTokenizer* japanese_tokenizer) { static constexpr std::pair languages[] = { {"zh", "chinese"}, {"yue", "cantonese"}, {"en", "english"}, {"de", "german"}, @@ -21,6 +24,20 @@ std::vector AlignmentUnits(const std::string& text, const std::stri if (found == std::end(languages)) throw std::invalid_argument("Forced alignment supports zh, yue, en, de, es, fr, it, pt, ru, ko and ja"); + if (!ValidUtf8(text)) + throw std::invalid_argument("Alignment requires valid UTF-8 text"); + if (granularity == AlignmentGranularity::Characters) + { + static const din::io::UnicodeRegex characters(R"(\X)"), spoken(R"([\p{L}\p{N}'])"); + std::vector result; + for (auto unit : characters.FindAll(text)) + if (!spoken.FindAll(unit).empty()) + result.push_back(std::move(unit)); + return result; + } + if (found->first == "ja") + return japanese_tokenizer->Words(text); + // HF keeps Unicode letters/numbers and ASCII apostrophes, dropping punctuation and marks. static const din::io::UnicodeRegex kept(R"([\p{L}\p{N}'\s\x{1c}-\x{1f}]+)"); std::string cleaned; @@ -31,11 +48,6 @@ std::vector AlignmentUnits(const std::string& text, const std::stri static const din::io::UnicodeRegex words("[" + cjk + "]|[^" + cjk + R"(\s\x{1c}-\x{1f}]+)"); // HF's unscored Korean LTokenizer splits on whitespace before cleaning each unit. static const din::io::UnicodeRegex korean(R"([^\s\x{1c}-\x{1f}]+)"); - // Japanese character timestamps avoid a separate Nagisa word-segmentation runtime. - static const std::string kana = R"(\x{3040}-\x{30ff}\x{31f0}-\x{31ff}\x{ff66}-\x{ff9f})"; - static const din::io::UnicodeRegex japanese("[" + cjk + kana + "]|[^" + cjk + kana + R"(\s\x{1c}-\x{1f}]+)"); - if (found->first == "ja") - return japanese.FindAll(cleaned); if (found->first == "ko") return korean.FindAll(cleaned); return words.FindAll(cleaned); @@ -91,8 +103,7 @@ struct AlignmentEngine AudioModel model; std::unique_ptr text; AlignmentEngine(Runtime& runtime, const std::filesystem::path& dir); - std::vector Align(const std::vector& features, const std::string& transcript, - const std::string& language); + std::vector Align(const std::vector& features, const std::vector& words); }; AlignmentEngine::AlignmentEngine(Runtime& runtime, const std::filesystem::path& dir) @@ -106,15 +117,12 @@ AlignmentEngine::AlignmentEngine(Runtime& runtime, const std::filesystem::path& }; text = runtime.Runner(dir, "aligner", shape(4, 2), shape(128, 32), shape(8192, 4096)); } -std::vector AlignmentEngine::Align(const std::vector& features, const std::string& transcript, - const std::string& language) +std::vector AlignmentEngine::Align(const std::vector& features, + const std::vector& words) { din::common::nvtx_scoped_range range{"qwen3.align"}; - if (transcript.empty()) - return {}; const int64_t frames = features.size() / 128; std::vector result; - const auto words = detail::AlignmentUnits(transcript, language); if (words.empty()) return {}; if (words.size() > 2048) @@ -131,7 +139,11 @@ std::vector AlignmentEngine::Align(const std::vector& feat const int64_t timestamp_id = model.metadata["timestamp_token_id"]; for (const auto& word : words) { - Append(ids, model.tokenizer->Encode(word, false)); + const auto tokens = model.tokenizer->Encode(word, false); + if (std::find(tokens.begin(), tokens.end(), model.audio_id) != tokens.end() || + std::find(tokens.begin(), tokens.end(), timestamp_id) != tokens.end()) + throw std::invalid_argument("Alignment units must not contain audio or timestamp markers"); + Append(ids, tokens); for (int i = 0; i < 2; ++i) { slots.push_back(ids.size()); @@ -187,28 +199,34 @@ struct Qwen3ForcedAligner::Impl ForcedAlignerConfig config; Runtime runtime; AlignmentEngine engine; + std::unique_ptr japanese; explicit Impl(ForcedAlignerConfig cfg) : config(std::move(cfg)) - , runtime(config.provider, config.ep_cache_dir, config.ep_context_dir, config.progress) + , runtime(config.provider, config.ep_cache_dir, config.ep_context_dir) , engine(runtime, config.model_dir) { runtime.LoadMel(config.model_dir); } - std::vector AlignAudio(const din::io::Audio& audio, const std::string& text, - const std::string& language) + std::vector Units(const std::string& text, const std::string& language) + { + const auto lang = LowerLanguage(language); + if ((lang == "ja" || lang == "japanese") && config.granularity == AlignmentGranularity::Words && !japanese) + japanese = std::make_unique(runtime.env, config.model_dir); + return AlignmentUnits(text, language, config.granularity, japanese.get()); + } + std::vector AlignAudio(const din::io::Audio& audio, const std::vector& words) { + for (const auto& word : words) + if (word.empty() || !ValidUtf8(word)) + throw std::invalid_argument("Alignment units must be nonempty UTF-8 text"); + if (words.empty()) + return {}; din::io::Audio normalized; const auto& source = NormalizeAudio(audio, normalized); if (source.samples.size() > 180 * kRate) throw std::invalid_argument("Standalone alignment accepts up to 180 seconds; supply audio/text segments"); - if (config.progress) - config.progress({din::common::ProgressStage::Aligning, "Aligning supplied text", 0, audio.Duration()}); const auto features = runtime.Features(source.samples); - auto result = engine.Align(features, text, language); - if (config.progress) - config.progress( - {din::common::ProgressStage::Aligning, "Alignment complete", audio.Duration(), audio.Duration()}); - return result; + return engine.Align(features, words); } }; Qwen3ForcedAligner::Qwen3ForcedAligner(ForcedAlignerConfig config) @@ -219,13 +237,50 @@ Qwen3ForcedAligner::~Qwen3ForcedAligner() = default; std::vector Qwen3ForcedAligner::Align(const din::io::Audio& audio, const std::string& text, const std::string& language) { - return impl_->AlignAudio(audio, text, language); + return impl_->AlignAudio(audio, impl_->Units(text, language)); +} +std::vector Qwen3ForcedAligner::AlignUnits(const din::io::Audio& audio, + const std::vector& units) +{ + return impl_->AlignAudio(audio, units); +} +std::vector Qwen3ForcedAligner::AlignSegments(const din::io::Audio& audio, + std::span segments) +{ + if (segments.empty()) + return {}; + if (audio.sample_rate != kRate) + throw std::invalid_argument("Segment alignment expects mono 16 kHz audio"); + for (size_t i = 0; i < segments.size(); ++i) + { + const auto& segment = segments[i]; + if (segment.start_sample >= segment.end_sample || segment.end_sample > audio.samples.size()) + throw std::invalid_argument("Alignment segment " + std::to_string(i) + " has invalid audio bounds"); + const auto count = segment.end_sample - segment.start_sample; + if (count > 180 * kRate) + throw std::invalid_argument("Alignment segment " + std::to_string(i) + + " exceeds 180 seconds; supply shorter audio/text segments"); + } + din::io::Audio clip; + clip.sample_rate = kRate; + std::vector result; + for (const auto& segment : segments) + { + clip.samples.assign(audio.samples.begin() + segment.start_sample, audio.samples.begin() + segment.end_sample); + auto aligned = impl_->AlignAudio(clip, impl_->Units(segment.text, segment.language)); + const float offset = static_cast(segment.start_sample) / kRate; + for (auto& word : aligned) + { + word.start_time += offset; + word.end_time += offset; + result.push_back(std::move(word)); + } + } + return result; } std::vector Qwen3ForcedAligner::AlignFile(const std::filesystem::path& path, const std::string& text, const std::string& language) { - if (impl_->config.progress) - impl_->config.progress({din::common::ProgressStage::DecodingAudio, path.filename().string()}); return Align(din::io::LoadAudio(path, kRate), text, language); } } // namespace din::asr::qwen3 diff --git a/asr/qwen3/forced_aligner.h b/asr/qwen3/forced_aligner.h index 48cc644..586979b 100644 --- a/asr/qwen3/forced_aligner.h +++ b/asr/qwen3/forced_aligner.h @@ -2,21 +2,26 @@ #pragma once #include #include +#include #include #include #include "audio.h" -#include "progress.h" namespace din::asr::qwen3 { +enum class AlignmentGranularity +{ + Words, + Characters +}; struct ForcedAlignerConfig { std::string provider = "trt-rtx"; std::filesystem::path model_dir = "artifacts/qwen3/aligner-onnx-bf16"; std::filesystem::path ep_cache_dir = "artifacts/qwen3/rt_cache"; std::filesystem::path ep_context_dir = "artifacts/qwen3/ep_context"; - din::common::ProgressCallback progress; + AlignmentGranularity granularity = AlignmentGranularity::Words; }; struct WordTimestamp @@ -26,6 +31,15 @@ struct WordTimestamp float end_time = 0; }; +struct AlignmentSegment +{ + std::string text; + // Half-open offsets in the supplied mono 16 kHz audio. + size_t start_sample = 0; + size_t end_sample = 0; + std::string language = "English"; +}; + // Accepts text from any recognizer; no ASR model is loaded. One synchronous call per instance. class Qwen3ForcedAligner { @@ -36,6 +50,9 @@ class Qwen3ForcedAligner const std::string& language = "English"); std::vector AlignFile(const std::filesystem::path& path, const std::string& text, const std::string& language = "English"); + std::vector AlignUnits(const din::io::Audio& audio, const std::vector& units); + // Returns recording-relative timestamps in segment order; overlaps are preserved. + std::vector AlignSegments(const din::io::Audio& audio, std::span segments); private: struct Impl; diff --git a/asr/qwen3/main.cpp b/asr/qwen3/main.cpp index 9b2287a..676384a 100644 --- a/asr/qwen3/main.cpp +++ b/asr/qwen3/main.cpp @@ -1,6 +1,11 @@ // SPDX-License-Identifier: Apache-2.0 +#include #include #include +#ifdef _WIN32 +#include +#include +#endif #include "qwen3.h" #include @@ -13,7 +18,7 @@ int main(int argc, char** argv) { Qwen3Config config; argparse::ArgumentParser parser("din_asr_qwen3_cli"); - parser.add_description("Offline Qwen3 ASR."); + parser.add_description("Offline or streaming Qwen3 ASR."); parser.add_argument("audiofile"); parser.add_argument("--provider").default_value(config.provider).choices("cpu", "trt-rtx"); parser.add_argument("--model-dir").default_value(config.model_dir.string()); @@ -23,6 +28,12 @@ int main(int argc, char** argv) parser.add_argument("--ep-cache").default_value(config.ep_cache_dir.string()); parser.add_argument("--ep-context-dir").default_value(config.ep_context_dir.string()); parser.add_argument("--max-new-tokens").default_value(config.max_new_tokens).scan<'i', int>(); + parser.add_argument("--stream") + .flag() + .help("Emit replacement hypotheses; audiofile '-' reads mono 16 kHz float32 PCM from stdin"); + parser.add_argument("--chunk-seconds").default_value(2.f).scan<'g', float>(); + parser.add_argument("--unfixed-chunks").default_value(2).scan<'i', int>(); + parser.add_argument("--unfixed-tokens").default_value(5).scan<'i', int>(); parser.add_argument("--max-chunk-seconds") .default_value(config.max_chunk_seconds) .scan<'i', int>() @@ -35,7 +46,58 @@ int main(int argc, char** argv) config.ep_context_dir = parser.get("--ep-context-dir"); config.max_new_tokens = parser.get("--max-new-tokens"); config.max_chunk_seconds = parser.get("--max-chunk-seconds"); + std::ostream results(std::cout.rdbuf()); + // Keep runtime diagnostics off the streaming JSON output. + if (parser.get("--stream")) + std::cout.rdbuf(std::cerr.rdbuf()); Qwen3Pipeline pipeline(std::move(config)); + if (parser.get("--stream")) + { + pipeline.StartStream({parser.get("--chunk-seconds"), parser.get("--unfixed-chunks"), + parser.get("--unfixed-tokens")}); + auto print = [&results](const StreamingResult& result) + { + results << nlohmann::json{{"transcription", result.text}, + {"language", result.language}, + {"samples_processed", result.samples_processed}, + {"sample_rate", 16000}, + {"updates", result.updates}, + {"final", result.final}, + {"reached_eos", result.reached_eos}} + .dump() + << std::endl; + }; + auto feed = [&](std::span samples) + { + for (const auto& result : pipeline.PushAudio(samples)) + print(result); + }; + if (parser.get("audiofile") == "-") + { +#ifdef _WIN32 + _setmode(_fileno(stdin), _O_BINARY); +#endif + float block[4096]; + while (std::cin.read(reinterpret_cast(block), sizeof(block)) || std::cin.gcount()) + { + if (std::cin.gcount() % sizeof(float)) + throw std::invalid_argument("Incomplete float32 PCM sample on stdin"); + feed(std::span(block, static_cast(std::cin.gcount()) / sizeof(float))); + } + if (std::cin.bad()) + throw std::runtime_error("Failed to read streaming PCM"); + } + else + { + const auto audio = din::io::LoadAudio(parser.get("audiofile"), 16000); + const auto samples = std::span(audio.samples); + for (size_t offset = 0; offset < samples.size(); offset += 4096) + feed(samples.subspan(offset, std::min(4096, samples.size() - offset))); + } + const auto result = pipeline.FinishStream(); + print(result); + return result.reached_eos ? 0 : 1; + } const auto result = pipeline.TranscribeFile(parser.get("audiofile")); nlohmann::json output{ {"transcription", result.text}, {"language", result.language}, diff --git a/asr/qwen3/model_export/detail/japanese.py b/asr/qwen3/model_export/detail/japanese.py new file mode 100644 index 0000000..c48afe8 --- /dev/null +++ b/asr/qwen3/model_export/detail/japanese.py @@ -0,0 +1,149 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Export Nagisa's word segmenter as an FP32 ONNX LSTM; no POS model is needed.""" + +import gzip +import hashlib +import json +import pickle +import re +from array import array +from pathlib import Path + +import onnx +import torch +from onnx import TensorProto, helper + +REVISION = "3c4bb48d3ba7451e3314b35337c79f1256ade0cf" +HASHES = { + "dict": "968ac9e6c7a53051ef24d8561673dd31de81b2feb9b5bff01b1d3b6b2473113c", + "hp": "6737f76b588315fe2fe05d05c99939d3142a0f1c6468f49f4c23e7003192f204", + "model": "9db9abc06a927c56e18af8d485e20a14908138752be459a83f0dc7ac85368c1b", +} + + +def export_japanese(output, source=None): + source = Path(source or Path(torch.hub.get_dir()) / "nagisa" / REVISION) + source.mkdir(parents=True, exist_ok=True) + for suffix, digest in HASHES.items(): + path = source / f"nagisa_v001.{suffix}" + if not path.exists(): + torch.hub.download_url_to_file( + f"https://raw.githubusercontent.com/taishi-i/nagisa/{REVISION}/nagisa/data/{path.name}", + str(path), + hash_prefix=digest, + ) + if hashlib.sha256(path.read_bytes()).hexdigest() != digest: + raise ValueError(f"Unexpected Nagisa asset: {path}") + # Only load the authenticated, pinned upstream dictionaries above. + vocabs = pickle.loads(gzip.decompress((source / "nagisa_v001.dict").read_bytes())) + hp = pickle.loads(gzip.decompress((source / "nagisa_v001.hp").read_bytes())) + parameters, lookups = [], [] + with (source / "nagisa_v001.model").open("rb") as file: + while header := file.readline(): + match = re.fullmatch(rb"#(Parameter|LookupParameter)# (\S+) \{([\d,]+)\} (\d+) ZERO_GRAD\s*", header) + if not match: + raise ValueError("Unsupported Nagisa parameter layout") + kind, name, dimensions, count = match.groups() + shape = [int(n) for n in dimensions.split(b",")] + data = torch.tensor(array("f", (float(n) for n in file.read(int(count)).split())), dtype=torch.float32) + if len(shape) == 2: + data = data.reshape(shape[::-1]) + if kind == b"Parameter": + data = data.T + (parameters if kind == b"Parameter" else lookups).append((name.decode(), data.contiguous())) + + nodes, constants = [], [] + + def const(name, data): + data = data.contiguous().clone() + tensor = TensorProto(name=name, data_type=TensorProto.INT64 if data.dtype == torch.int64 else TensorProto.FLOAT) + tensor.dims.extend(data.shape) + tensor.raw_data = bytes(data.untyped_storage()) + constants.append(tensor) + return name + + def node(op, inputs, name, **attributes): + nodes.append(helper.make_node(op, inputs, [name], **attributes)) + return name + + width = hp["WINDOW_SIZE"] + inputs, pieces = [], [] + word_table = torch.cat([lookups[2][1], torch.zeros(1, hp["DIM_WORD"])]) + const("word_table", word_table) + const("axis1", torch.tensor([1])) + for name, table in [("unigrams", lookups[0][1]), ("bigrams", lookups[1][1]), ("types", lookups[3][1])]: + inputs.append(helper.make_tensor_value_info(name, TensorProto.INT64, ["characters", width])) + gathered = node("Gather", [const(name + "_table", table), name], name + "_vectors", axis=0) + pieces.append(node("Flatten", [gathered], name + "_flat", axis=1)) + for name in ("word_starts", "word_ends"): + inputs.append(helper.make_tensor_value_info(name, TensorProto.INT64, ["characters", 8])) + gathered = node("Gather", ["word_table", name], name + "_vectors", axis=0) + pieces.append(node("ReduceSum", [gathered, "axis1"], name + "_sum", keepdims=0)) + x = node("Concat", pieces, "features", axis=1) + x = node("Unsqueeze", [x, "axis1"], "sequence") + hidden = hp["DIM_HIDDEN"] // 2 + # DyNet gates i,f,o,g -> ONNX gates i,o,f,g; DyNet adds +1 to the forget bias. + order = torch.cat([torch.arange(i * hidden, (i + 1) * hidden) for i in (0, 2, 1, 3)]) + for layer in range(hp["LAYERS"]): + weights, recurrent, biases = [], [], [] + for direction in range(2): + offset = (2 * layer + direction) * 3 + wx, wh, bias = [p[1] for p in parameters[offset : offset + 3]] + bias = bias.clone() + bias[hidden : 2 * hidden] += 1 + weights.append(wx[order]) + recurrent.append(wh[order]) + biases.append(torch.cat([bias[order], torch.zeros(4 * hidden)])) + name = f"lstm_{layer}" + x = node( + "LSTM", + [ + x, + const(name + "_w", torch.stack(weights)), + const(name + "_r", torch.stack(recurrent)), + const(name + "_b", torch.stack(biases)), + ], + name, + hidden_size=hidden, + direction="bidirectional", + ) + x = node("Transpose", [x], name + "_ordered", perm=[0, 2, 1, 3]) + x = node("Reshape", [x, const(name + "_shape", torch.tensor([-1, 1, 2 * hidden]))], name + "_flat") + plain = [p for p in parameters if p[0].count("/") == 1] + x = node("Squeeze", [x, "axis1"], "hidden") + x = node("MatMul", [x, const("projection", plain[0][1].T)], "projected") + node("Add", [x, const("bias", plain[1][1])], "emissions") + graph = helper.make_graph( + nodes, + "nagisa_word_segmentation", + inputs, + [helper.make_tensor_value_info("emissions", TensorProto.FLOAT, ["characters", 6])], + constants, + ) + model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 17)], ir_version=8) + output = Path(output) + output.mkdir(parents=True, exist_ok=True) + path = output / "japanese.onnx" + path.with_suffix(".onnx.data").unlink(missing_ok=True) + onnx.save_model( + model, + path, + save_as_external_data=True, + all_tensors_to_one_file=True, + location="japanese.onnx.data", + size_threshold=1024, + ) + onnx.checker.check_model(str(path)) + metadata = { + "revision": REVISION, + "window": width, + "unigrams": vocabs[0], + "bigrams": vocabs[1], + "words": vocabs[2], + "padding_word": len(word_table) - 1, + "transitions": lookups[5][1].tolist(), + } + (output / "japanese.json").write_text(json.dumps(metadata, ensure_ascii=False), encoding="utf-8") + (output / "japanese.LICENSE.txt").write_text( + Path(__file__).with_name("nagisa.LICENSE.txt").read_text(encoding="utf-8"), encoding="utf-8" + ) diff --git a/asr/qwen3/model_export/detail/nagisa.LICENSE.txt b/asr/qwen3/model_export/detail/nagisa.LICENSE.txt new file mode 100644 index 0000000..52360b0 --- /dev/null +++ b/asr/qwen3/model_export/detail/nagisa.LICENSE.txt @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2018 taishi-i + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/asr/qwen3/model_export/export_qwen3_asr.py b/asr/qwen3/model_export/export_qwen3_asr.py index a46fefc..dee041d 100644 --- a/asr/qwen3/model_export/export_qwen3_asr.py +++ b/asr/qwen3/model_export/export_qwen3_asr.py @@ -326,7 +326,7 @@ def main(): parser.add_argument("--output", type=Path, help="Defaults to the ONNX artifact directory for --task") parser.add_argument("--task", choices=["asr", "aligner"], default="asr") parser.add_argument("--dtype", choices=["original", "fp16", "fp32"], default="original") - parser.add_argument("--only", choices=["mel", "encoder", "decoder", "aligner"]) + parser.add_argument("--only", choices=["mel", "encoder", "decoder", "aligner", "japanese"]) parser.add_argument("--threads", type=int, default=4) parser.add_argument("--cache-capacity", type=int, choices=range(512, 16385, 512), default=8192, metavar="TOKENS") args = parser.parse_args() @@ -337,6 +337,14 @@ def main(): precision = "bf16" if args.dtype == "original" else args.dtype size_suffix = "-1.7b" if args.task == "asr" and "1.7b" in args.model.lower() else "" args.output = args.output or Path(f"artifacts/qwen3/{prefix}onnx-{precision}{size_suffix}") + if args.only == "japanese" and args.task != "aligner": + parser.error("--only japanese requires --task aligner") + if args.task == "aligner" and args.only in (None, "japanese"): + from asr.qwen3.model_export.detail.japanese import export_japanese + + export_japanese(args.output) + if args.only == "japanese": + return if args.only in ("decoder", "aligner") and args.only != ("decoder" if args.task == "asr" else "aligner"): parser.error("--only must match --task") metadata_path = args.output / "metadata.json" diff --git a/asr/qwen3/qwen3.cpp b/asr/qwen3/qwen3.cpp index f36ff91..85e488d 100644 --- a/asr/qwen3/qwen3.cpp +++ b/asr/qwen3/qwen3.cpp @@ -5,9 +5,11 @@ #include #include #include +#include #include #include "detail/runtime.h" +#include "detail/text.h" namespace din::asr::qwen3::detail { @@ -199,10 +201,19 @@ struct Qwen3Pipeline::Impl std::unique_ptr> next_token; std::unique_ptr decode_binding, prefill_binding; Ort::RunOptions decode_options; + struct Stream + { + StreamingConfig config; + size_t chunk_samples; + std::vector audio; + std::string raw; + StreamingResult result; + }; + std::optional stream; explicit Impl(Qwen3Config cfg) : config(std::move(cfg)) - , runtime(config.provider, config.ep_cache_dir, config.ep_context_dir, config.progress) + , runtime(config.provider, config.ep_cache_dir, config.ep_context_dir) , asr(runtime, config.model_dir, "asr") { if (config.max_new_tokens <= 0) @@ -257,21 +268,26 @@ struct Qwen3Pipeline::Impl binding.BindOutput("next_token", next_token->BindingValue()); } - TranscriptionResult TranscribeChunk(std::span audio) + TranscriptionResult TranscribeChunk(std::span audio, const std::string& prefix = {}, + std::string* raw_output = nullptr) { if (audio.empty() || audio.size() > kMaxSamples) throw std::runtime_error("Expected nonempty mono 16 kHz audio, at most 1205 seconds per chunk"); - auto features = runtime.Features(audio); - const int64_t frames = features.size() / 128; - auto encoded = asr.Encode(features, frames); std::vector ids = asr.native["prefixes"][config.lang_id]; - ids.insert(ids.end(), encoded.tokens, asr.audio_id); + const auto audio_tokens = AudioTokens(std::max(8000, audio.size()) / 160); + ids.insert(ids.end(), audio_tokens, asr.audio_id); if (!asr.native.contains("suffixes")) throw std::runtime_error("Re-export native prompt assets with --only mel for official language forcing"); Append(ids, asr.native["suffixes"][config.lang_id].get>()); + if (!prefix.empty()) + Append(ids, asr.tokenizer->Encode(prefix, false)); + if (ids.size() + config.max_new_tokens > asr.metadata["cache_capacity"].get()) + throw std::runtime_error("Audio, text prefix and generation budget exceed the exported KV capacity; " + "use shorter utterances or export a larger cache"); + auto features = runtime.Features(audio); + const int64_t frames = features.size() / 128; + auto encoded = asr.Encode(features, frames); PrepareCache(); - if (ids.size() + config.max_new_tokens > static_cast(capacity)) - throw std::runtime_error("Prompt and generation budget exceed the exported KV capacity"); // Clear on the shared CUDA stream. Unused NaN cache values can poison attention even when masked. for (auto& value : cache) ZeroTensor(*text, value); @@ -307,21 +323,48 @@ struct Qwen3Pipeline::Impl auto text_tokens = result.tokens; if (result.reached_eos) text_tokens.pop_back(); - const auto parsed = detail::ParseOutput(asr.tokenizer->Decode(text_tokens, false), - asr.native["languages"][config.lang_id].get()); + const auto raw = prefix + asr.tokenizer->Decode(text_tokens, false); + if (raw_output) + *raw_output = raw; + const auto parsed = detail::ParseOutput(raw, asr.native["languages"][config.lang_id].get()); result.language = parsed.first; result.text = parsed.second; return result; } + StreamingResult DecodeStream(size_t samples, bool final) + { + auto& s = *stream; + std::string prefix; + if (s.result.updates >= static_cast(s.config.unfixed_chunks) && !s.raw.empty()) + { + auto ids = asr.tokenizer->Encode(s.raw, false); + // Roll back incomplete UTF-8 tokens as well as the editable suffix. + ids.resize(ids.size() > static_cast(s.config.unfixed_tokens) ? ids.size() - s.config.unfixed_tokens + : 0); + while (!ids.empty()) + { + prefix = asr.tokenizer->Decode(ids, false); + if (ValidUtf8(prefix) && prefix.find("\xef\xbf\xbd") == std::string::npos) + break; + ids.pop_back(); + prefix.clear(); + } + } + auto result = TranscribeChunk(std::span(s.audio).first(samples), prefix, &s.raw); + s.result = {std::move(result.text), std::move(result.language), samples, + s.result.updates + 1, result.reached_eos, final}; + return s.result; + } + TranscriptionResult Transcribe(const din::io::Audio& audio) { + if (stream && !stream->result.final) + throw std::logic_error("Finish the active stream before offline transcription"); din::common::nvtx_scoped_range range{"qwen3.transcribe"}; const auto start = std::chrono::steady_clock::now(); din::io::Audio normalized; const auto* source = &NormalizeAudio(audio, normalized); - if (config.progress) - config.progress({din::common::ProgressStage::Transcribing, "Transcribing", 0, audio.Duration()}); TranscriptionResult result; result.reached_eos = true; std::string previous_language; @@ -354,12 +397,6 @@ struct Qwen3Pipeline::Impl result.segments.push_back({std::move(part.text), std::move(part.language), chunk.begin, chunk.end}); ++result.chunks_processed; result.reached_eos = result.reached_eos && part.reached_eos; - if (config.progress) - config.progress( - {din::common::ProgressStage::Transcribing, - "Completed chunk " + std::to_string(result.chunks_processed), - chunk.end == source->samples.size() ? audio.Duration() : static_cast(chunk.end) / kRate, - audio.Duration()}); } result.audio_seconds = static_cast(audio.Duration()); result.transcribe_seconds = std::chrono::duration(std::chrono::steady_clock::now() - start).count(); @@ -372,14 +409,45 @@ Qwen3Pipeline::Qwen3Pipeline(Qwen3Config config) { } Qwen3Pipeline::~Qwen3Pipeline() = default; +void Qwen3Pipeline::StartStream(StreamingConfig config) +{ + if (!std::isfinite(config.chunk_seconds) || config.chunk_seconds < 0.5f || config.chunk_seconds > 1200.f || + config.unfixed_chunks < 0 || config.unfixed_tokens < 0) + throw std::invalid_argument("Streaming requires chunk-seconds 0.5..1200 and nonnegative rollback settings"); + impl_->stream = Impl::Stream{config, static_cast(std::lround(config.chunk_seconds * kRate)), {}, {}, {}}; +} +std::vector Qwen3Pipeline::PushAudio(std::span pcm16k) +{ + if (!impl_->stream || impl_->stream->result.final) + throw std::logic_error("Call StartStream before supplying audio"); + for (float sample : pcm16k) + if (!std::isfinite(sample) || std::abs(sample) > 1.f) + throw std::invalid_argument("Streaming expects finite mono 16 kHz float PCM in [-1, 1]"); + auto& s = *impl_->stream; + if (pcm16k.size() > kMaxSamples - s.audio.size()) + throw std::invalid_argument("Streaming utterance is too long; finish it and start a new stream"); + s.audio.insert(s.audio.end(), pcm16k.begin(), pcm16k.end()); + std::vector updates; + while (s.audio.size() - s.result.samples_processed >= s.chunk_samples) + updates.push_back(impl_->DecodeStream(s.result.samples_processed + s.chunk_samples, false)); + return updates; +} +StreamingResult Qwen3Pipeline::FinishStream() +{ + if (!impl_->stream) + throw std::logic_error("Call StartStream before finishing audio"); + auto& s = *impl_->stream; + if (!s.result.final && s.audio.size() > s.result.samples_processed) + impl_->DecodeStream(s.audio.size(), true); + s.result.final = true; + return s.result; +} TranscriptionResult Qwen3Pipeline::Transcribe(const din::io::Audio& audio) { return impl_->Transcribe(audio); } TranscriptionResult Qwen3Pipeline::TranscribeFile(const std::filesystem::path& path) { - if (impl_->config.progress) - impl_->config.progress({din::common::ProgressStage::DecodingAudio, path.filename().string()}); return Transcribe(din::io::LoadAudio(path, kRate)); } } // namespace din::asr::qwen3 diff --git a/asr/qwen3/qwen3.h b/asr/qwen3/qwen3.h index e4a8a1c..303265f 100644 --- a/asr/qwen3/qwen3.h +++ b/asr/qwen3/qwen3.h @@ -4,11 +4,11 @@ #include #include #include +#include #include #include #include "audio.h" -#include "progress.h" namespace din::asr::qwen3 { @@ -21,7 +21,6 @@ struct Qwen3Config std::string lang_id = "auto"; int max_new_tokens = 1024; int max_chunk_seconds = 0; // Auto: 1200 s, limited by KV capacity; boundaries add up to 5 s. - din::common::ProgressCallback progress; }; struct TranscriptionSegment @@ -45,6 +44,23 @@ struct TranscriptionResult float transcribe_seconds = 0; }; +struct StreamingConfig +{ + float chunk_seconds = 2.f; + int unfixed_chunks = 2; + int unfixed_tokens = 5; +}; + +struct StreamingResult +{ + std::string text; + std::string language; + size_t samples_processed = 0; + size_t updates = 0; + bool reached_eos = true; + bool final = false; +}; + // One synchronous stream per instance. Reuse the instance across files. class Qwen3Pipeline { @@ -55,6 +71,10 @@ class Qwen3Pipeline Qwen3Pipeline& operator=(const Qwen3Pipeline&) = delete; TranscriptionResult Transcribe(const din::io::Audio& audio); TranscriptionResult TranscribeFile(const std::filesystem::path& path); + // One utterance per stream. Each update replaces the previous hypothesis. + void StartStream(StreamingConfig config = {}); + std::vector PushAudio(std::span pcm16k); + StreamingResult FinishStream(); private: struct Impl; diff --git a/common/ort_session.cpp b/common/ort_session.cpp index fbcce7c..46558e9 100644 --- a/common/ort_session.cpp +++ b/common/ort_session.cpp @@ -588,8 +588,6 @@ std::string CompileEpContextModel(Ort::Env& env, const std::string& model_path, const size_t size_threshold_external_init = 1024; compile_options.SetOutputModelExternalInitializersFile(output_external_wide.c_str(), size_threshold_external_init); - if (ep_context.progress) - ep_context.progress({ProgressStage::CompilingModel, fs::path(model_path).filename().string()}); const Ort::Status status = Ort::CompileModel(env, compile_options); if (!status.IsOK()) { @@ -783,8 +781,6 @@ OrtRunner::OrtRunner(Ort::Env& env_in, const std::string& model_path, const std: { DIN_NVTX_FUNC_RANGE(); auto session_model_path = model_path; - if (ep_context.progress) - ep_context.progress({ProgressStage::LoadingModel, fs::path(model_path).filename().string()}); if (provider == "trt-rtx") { @@ -822,8 +818,6 @@ OrtRunner::OrtRunner(Ort::Env& env_in, const std::string& model_path, const std: throw std::runtime_error("unsupported provider: " + provider); } - if (ep_context.progress) - ep_context.progress({ProgressStage::LoadingModel, fs::path(session_model_path).filename().string()}); #ifdef _WIN32 const auto wide = ToOrtPathString(session_model_path); { diff --git a/common/ort_session.h b/common/ort_session.h index 2fca6c1..127a599 100644 --- a/common/ort_session.h +++ b/common/ort_session.h @@ -13,7 +13,6 @@ #include #include -#include "progress.h" #include #include #include @@ -69,7 +68,6 @@ QueryCudaGraphicsInteropSharedMemoryInfo(int cuda_device_ordinal, bool high_prio struct EpContextOptions { std::string output_dir; - ProgressCallback progress; }; struct ModelProfile diff --git a/common/progress.h b/common/progress.h deleted file mode 100644 index 7c65d66..0000000 --- a/common/progress.h +++ /dev/null @@ -1,26 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 -#pragma once -#include -#include - -namespace din::common -{ -enum class ProgressStage -{ - DecodingAudio, - LoadingModel, - CompilingModel, - Transcribing, - Aligning -}; -struct InferenceProgress -{ - ProgressStage stage; - std::string detail; - double completed_audio_seconds = 0; - double total_audio_seconds = 0; -}; -// Invoked synchronously on the pipeline's calling thread. Keep callbacks short. -// Optional: existing CLI callers need not install a callback. -using ProgressCallback = std::function; -} // namespace din::common From 26fd323f8e2491104fb9708c8ffaff22ae698d19 Mon Sep 17 00:00:00 2001 From: contentis Date: Wed, 23 Sep 2026 10:46:45 +0200 Subject: [PATCH 3/8] Skip draft CI and scope model tests and export reuse --- .github/scripts/find_successful_artifact.py | 23 ++-- .github/scripts/model_ci.py | 112 ++++++++++++++++++ .github/workflows/ci.yml | 5 + .github/workflows/export-nvidia-asr-model.yml | 49 +++----- .github/workflows/export-sam2-model.yml | 69 ++++------- .github/workflows/export-whisper-model.yml | 50 +++----- .github/workflows/model-changes.yml | 32 +++++ .github/workflows/nvidia-asr.yml | 19 ++- .github/workflows/sam2.yml | 31 +++-- .github/workflows/whisper-asr.yml | 7 ++ 10 files changed, 257 insertions(+), 140 deletions(-) create mode 100644 .github/scripts/model_ci.py create mode 100644 .github/workflows/model-changes.yml diff --git a/.github/scripts/find_successful_artifact.py b/.github/scripts/find_successful_artifact.py index 82338ee..4174f0e 100644 --- a/.github/scripts/find_successful_artifact.py +++ b/.github/scripts/find_successful_artifact.py @@ -25,7 +25,10 @@ def github_api(path: str) -> object: return json.load(response) -def write_outputs(*, found: bool, artifact_name: str, run_id: str = "") -> None: +def write_outputs(*, found: bool, artifact_name: str, run_id: str = "", as_json: bool = False) -> None: + if as_json: + print(json.dumps({"found": found, "name": artifact_name if found else "", "run_id": run_id})) + return output_path = Path(os.environ["GITHUB_OUTPUT"]) with output_path.open("a", encoding="utf-8") as output: output.write(f"found={'true' if found else 'false'}\n") @@ -40,16 +43,15 @@ def main() -> None: parser.add_argument("--default-branch", required=True) parser.add_argument("--event-name", required=True) parser.add_argument("--source-branch", required=True) + parser.add_argument("--pull-request", default="") + parser.add_argument("--json", action="store_true") args = parser.parse_args() artifact_name = args.artifact_name allowed_branches = {args.default_branch} - if args.event_name != "pull_request": + if args.event_name != "pull_request" or args.pull_request: allowed_branches.add(args.source_branch) - response = github_api( - f"repos/{args.repository}/actions/artifacts" - f"?name={quote(artifact_name)}&per_page=100" - ) + response = github_api(f"repos/{args.repository}/actions/artifacts?name={quote(artifact_name)}&per_page=100") artifacts = sorted( ( artifact @@ -67,18 +69,23 @@ def main() -> None: run_id = int(artifact["workflow_run"]["id"]) if run_id not in conclusions: run = github_api(f"repos/{args.repository}/actions/runs/{run_id}") + same_pr = args.pull_request and any( + str(pr["number"]) == args.pull_request for pr in run.get("pull_requests", []) + ) + trusted_branch = run["event"] != "pull_request" and run["head_branch"] in allowed_branches conclusions[run_id] = ( - run["status"] == "completed" and run["conclusion"] == "success" + (trusted_branch or same_pr) and run["status"] == "completed" and run["conclusion"] == "success" ) if conclusions[run_id]: write_outputs( found=True, artifact_name=artifact_name, run_id=str(run_id), + as_json=args.json, ) return - write_outputs(found=False, artifact_name=artifact_name) + write_outputs(found=False, artifact_name=artifact_name, as_json=args.json) if __name__ == "__main__": diff --git a/.github/scripts/model_ci.py b/.github/scripts/model_ci.py new file mode 100644 index 0000000..207a4eb --- /dev/null +++ b/.github/scripts/model_ci.py @@ -0,0 +1,112 @@ +#!/usr/bin/env python3 +"""Select affected model tests and fingerprint their ONNX export inputs.""" + +import argparse +import fnmatch +import hashlib +import json +import os +import subprocess +from pathlib import Path + +SHARED_RUNTIME = [ + "CMakeLists.txt", + "CMakePresets.json", + "cmake/*", + "common/*.cpp", + "common/*.h", + "common/*CMakeLists.txt", + ".github/workflows/ci.yml", + ".github/scripts/*", + ".github/workflows/model-changes.yml", +] +EXPORTS = { + "whisper": [ + "asr/whisper/model_export/export_whisper.py", + "common/model_export/*.py", + ".github/workflows/export-whisper-model.yml", + ], + "parakeet": [ + "asr/rnnt/python/tools/export_parakeet_tdt_onnx.py", + "asr/rnnt/python/rnnt/parakeet_tdt/nemo_backend.py", + "asr/rnnt/python/rnnt/parakeet_tdt/__init__.py", + ], + "nemotron": [ + "asr/rnnt/python/tools/export_nemotron_onnx.py", + "asr/rnnt/python/rnnt/nemotron_asr/nemo_backend.py", + "asr/rnnt/python/rnnt/nemotron_asr/__init__.py", + ], + "sam2": [ + "vision/sam2/python/export_sam2_onnx.py", + "vision/sam2/python/modeling.py", + ".github/workflows/export-sam2-model.yml", + ], +} +for model in ("parakeet", "nemotron"): + EXPORTS[model] += [ + "asr/rnnt/python/tools/nemo_preprocessor_export.py", + "asr/rnnt/python/rnnt/__init__.py", + "asr/rnnt/python/rnnt/audio.py", + ".github/workflows/export-nvidia-asr-model.yml", + ] +RUNTIME = { + "whisper": ["asr/whisper/*", "assets/sample.wav", ".github/workflows/whisper-asr.yml"], + "sam2": ["vision/sam2/*", "assets/sam2-*", ".github/workflows/sam2.yml"], +} +for model, package in (("parakeet", "parakeet_tdt"), ("nemotron", "nemotron_asr")): + RUNTIME[model] = [ + f"asr/rnnt/cpp/{model}*", + "asr/rnnt/cpp/asr*", + "asr/rnnt/CMakeLists.txt", + f"asr/rnnt/python/rnnt/{package}/*", + "asr/rnnt/python/rnnt/validation/*", + "asr/rnnt/python/tests/*", + "assets/sample.wav", + ".github/workflows/nvidia-asr.yml", + ] + + +def affected(model, files): + patterns = SHARED_RUNTIME + RUNTIME[model] + EXPORTS[model] + return any(fnmatch.fnmatchcase(path, pattern) for path in files if not path.endswith(".md") for pattern in patterns) + + +def export_key(model, model_id, attention): + # Index blob IDs are independent of checkout line endings and include deleted/renamed inputs. + sources = subprocess.check_output(["git", "ls-files", "--stage", "--", *EXPORTS[model]]) + options = json.dumps([model, model_id, attention]).encode() + return hashlib.sha256(sources + options).hexdigest()[:16] + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("mode", choices=["changes", "key"]) + parser.add_argument("--model", choices=EXPORTS) + parser.add_argument("--model-id", default="") + parser.add_argument("--attention", default="") + args = parser.parse_args() + if args.mode == "key": + if not args.model or not args.model_id: + parser.error("key requires --model and --model-id") + print(export_key(args.model, args.model_id, args.attention)) + return + event = json.loads(Path(os.environ["GITHUB_EVENT_PATH"]).read_text()) + pr = event.get("pull_request") + base = pr["base"]["sha"] if pr else event.get("before") + files = None + if base and subprocess.run(["git", "cat-file", "-e", f"{base}^{{commit}}"], capture_output=True).returncode == 0: + if pr: + base = subprocess.check_output(["git", "merge-base", base, "HEAD"], text=True).strip() + files = ( + subprocess.check_output(["git", "diff", "--no-renames", "--name-only", "-z", base, "HEAD"]) + .decode() + .split("\0") + ) + with Path(os.environ["GITHUB_OUTPUT"]).open("a", encoding="utf-8") as output: + for model in EXPORTS: + changed = files is None or affected(model, files) + output.write(f"{model}={str(changed).lower()}\n") + + +if __name__ == "__main__": + main() diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 5d0965e..aae8c6c 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -4,6 +4,7 @@ on: push: branches: [main] pull_request: + types: [opened, synchronize, reopened, ready_for_review, converted_to_draft] workflow_dispatch: permissions: @@ -19,6 +20,7 @@ env: jobs: clang-format: + if: github.event_name != 'pull_request' || !github.event.pull_request.draft name: clang-format runs-on: ubuntu-24.04 @@ -47,6 +49,7 @@ jobs: | xargs -0 -r clang-format-22 --dry-run --Werror -- linux-x64: + if: github.event_name != 'pull_request' || !github.event.pull_request.draft name: Ubuntu 24.04 / Clang / x64 runs-on: ubuntu-24.04 container: nvidia/cuda:13.2.1-devel-ubuntu24.04 @@ -90,6 +93,7 @@ jobs: path: out/build/linux-x64/bin/Release/ linux-arm64: + if: github.event_name != 'pull_request' || !github.event.pull_request.draft name: Ubuntu 24.04 / Clang / ARM64 runs-on: ubuntu-24.04-arm container: nvidia/cuda:13.2.1-devel-ubuntu24.04 @@ -131,6 +135,7 @@ jobs: path: out/build/linux-arm64/bin/Release/ windows-x64: + if: github.event_name != 'pull_request' || !github.event.pull_request.draft name: Windows Server 2025 / MSVC / x64 runs-on: windows-2025 timeout-minutes: 90 diff --git a/.github/workflows/export-nvidia-asr-model.yml b/.github/workflows/export-nvidia-asr-model.yml index dc098d1..7bc2a6c 100644 --- a/.github/workflows/export-nvidia-asr-model.yml +++ b/.github/workflows/export-nvidia-asr-model.yml @@ -51,34 +51,11 @@ jobs: fetch-depth: 0 ref: ${{ inputs.source_ref }} - - name: Determine export requirement - id: changes - shell: bash - env: - BASE_SHA: ${{ github.event.pull_request.base.sha || github.event.before }} - EVENT_NAME: ${{ github.event_name }} - SOURCE_REF: ${{ inputs.source_ref }} - run: | - set -euo pipefail - if [[ "$EVENT_NAME" == "workflow_dispatch" ]]; then - force=false - elif [[ -z "$BASE_SHA" ]] || ! git cat-file -e "${BASE_SHA}^{commit}"; then - force=true - elif git diff --quiet "$BASE_SHA" "$SOURCE_REF" -- \ - 'asr/rnnt/python/*.py' \ - 'asr/rnnt/python/**/*.py' \ - '.github/workflows/export-nvidia-asr-model.yml' \ - '.github/workflows/nvidia-asr.yml'; then - force=false - else - force=true - fi - echo "force=$force" >> "$GITHUB_OUTPUT" - - name: Compute artifact names id: artifact-names shell: bash env: + MODEL_ID: ${{ inputs.model_id }} FALLBACK_PRECISION: ${{ inputs.fallback_precision }} MODEL_SLUG: ${{ inputs.model_slug }} MODEL_TYPE: ${{ inputs.model_type }} @@ -99,13 +76,14 @@ jobs: *) echo "Unsupported fallback precision: $FALLBACK_PRECISION" >&2; exit 2 ;; esac fi - echo "primary=asr-${MODEL_SLUG}-onnx-${PRECISION}" >> "$GITHUB_OUTPUT" + source_hash="$(python .github/scripts/model_ci.py key --model "${MODEL_TYPE}" --model-id "$MODEL_ID")" + echo "source_hash=$source_hash" >> "$GITHUB_OUTPUT" + echo "primary=asr-${MODEL_SLUG}-onnx-${PRECISION}-${source_hash}" >> "$GITHUB_OUTPUT" if [[ -n "$FALLBACK_PRECISION" ]]; then - echo "fallback=asr-${MODEL_SLUG}-onnx-${FALLBACK_PRECISION}" >> "$GITHUB_OUTPUT" + echo "fallback=asr-${MODEL_SLUG}-onnx-${FALLBACK_PRECISION}-${source_hash}" >> "$GITHUB_OUTPUT" fi - name: Find successful primary artifact - if: steps.changes.outputs.force != 'true' id: primary-artifact shell: bash env: @@ -113,17 +91,18 @@ jobs: EVENT_NAME: ${{ github.event_name }} GH_TOKEN: ${{ github.token }} SOURCE_BRANCH: ${{ github.head_ref || github.ref_name }} + PULL_REQUEST: ${{ github.event.pull_request.number }} run: | python .github/scripts/find_successful_artifact.py \ --repository "$GITHUB_REPOSITORY" \ --artifact-name "${{ steps.artifact-names.outputs.primary }}" \ --default-branch "$DEFAULT_BRANCH" \ --event-name "$EVENT_NAME" \ - --source-branch "$SOURCE_BRANCH" + --source-branch "$SOURCE_BRANCH" \ + --pull-request "$PULL_REQUEST" - name: Find successful fallback artifact if: >- - steps.changes.outputs.force != 'true' && steps.primary-artifact.outputs.found != 'true' && inputs.fallback_precision != '' id: fallback-artifact @@ -133,13 +112,15 @@ jobs: EVENT_NAME: ${{ github.event_name }} GH_TOKEN: ${{ github.token }} SOURCE_BRANCH: ${{ github.head_ref || github.ref_name }} + PULL_REQUEST: ${{ github.event.pull_request.number }} run: | python .github/scripts/find_successful_artifact.py \ --repository "$GITHUB_REPOSITORY" \ --artifact-name "${{ steps.artifact-names.outputs.fallback }}" \ --default-branch "$DEFAULT_BRANCH" \ --event-name "$EVENT_NAME" \ - --source-branch "$SOURCE_BRANCH" + --source-branch "$SOURCE_BRANCH" \ + --pull-request "$PULL_REQUEST" - name: Select export action id: selection @@ -148,15 +129,12 @@ jobs: FALLBACK_FOUND: ${{ steps.fallback-artifact.outputs.found }} FALLBACK_NAME: ${{ steps.fallback-artifact.outputs.name }} FALLBACK_RUN_ID: ${{ steps.fallback-artifact.outputs.run_id }} - FORCE_EXPORT: ${{ steps.changes.outputs.force }} PRIMARY_FOUND: ${{ steps.primary-artifact.outputs.found }} PRIMARY_NAME: ${{ steps.primary-artifact.outputs.name }} PRIMARY_RUN_ID: ${{ steps.primary-artifact.outputs.run_id }} run: | set -euo pipefail - if [[ "$FORCE_EXPORT" == "true" ]]; then - echo "export=true" >> "$GITHUB_OUTPUT" - elif [[ "$PRIMARY_FOUND" == "true" ]]; then + if [[ "$PRIMARY_FOUND" == "true" ]]; then echo "export=false" >> "$GITHUB_OUTPUT" echo "name=$PRIMARY_NAME" >> "$GITHUB_OUTPUT" echo "run_id=$PRIMARY_RUN_ID" >> "$GITHUB_OUTPUT" @@ -213,6 +191,7 @@ jobs: PRECISION: ${{ inputs.precision }} PYTHONPATH: asr/rnnt/python SOURCE_REF: ${{ inputs.source_ref }} + SOURCE_HASH: ${{ steps.artifact-names.outputs.source_hash }} run: | set -uo pipefail export_dtype() { @@ -288,7 +267,7 @@ jobs: json.dumps(manifest, indent=2) + "\n", encoding="utf-8" ) PY - echo "name=asr-${MODEL_SLUG}-onnx-${selected_precision}" >> "$GITHUB_OUTPUT" + echo "name=asr-${MODEL_SLUG}-onnx-${selected_precision}-${SOURCE_HASH}" >> "$GITHUB_OUTPUT" - name: Upload NVIDIA ASR ONNX if: steps.selection.outputs.export == 'true' diff --git a/.github/workflows/export-sam2-model.yml b/.github/workflows/export-sam2-model.yml index 227015b..c2df020 100644 --- a/.github/workflows/export-sam2-model.yml +++ b/.github/workflows/export-sam2-model.yml @@ -41,43 +41,15 @@ jobs: fetch-depth: 0 ref: ${{ inputs.source_ref }} - - name: Determine export requirement - id: changes - shell: bash - env: - BASE_SHA: ${{ github.event.pull_request.base.sha || github.event.before }} - EVENT_NAME: ${{ github.event_name }} - SOURCE_REF: ${{ inputs.source_ref }} - run: | - set -euo pipefail - if [[ "$EVENT_NAME" == "workflow_dispatch" ]]; then - force=false - elif [[ -z "$BASE_SHA" ]] || ! git cat-file -e "${BASE_SHA}^{commit}"; then - force=true - elif git diff --quiet "$BASE_SHA" "$SOURCE_REF" -- \ - 'vision/sam2/python/*.py' \ - 'vision/sam2/python/**/*.py' \ - '.github/workflows/export-sam2-model.yml' \ - '.github/workflows/sam2.yml'; then - force=false - else - force=true - fi - echo "force=$force" >> "$GITHUB_OUTPUT" - - name: Compute model artifact key id: model-key shell: bash env: MODEL_SLUG: ${{ inputs.model_slug }} + MODEL_ID: ${{ inputs.model_id }} run: | set -euo pipefail - source_hash="$( - git ls-files -z -- 'vision/sam2/python/*.py' 'vision/sam2/python/**/*.py' \ - | xargs -0 sha256sum \ - | sha256sum \ - | cut -d' ' -f1 - )" + source_hash="$(python .github/scripts/model_ci.py key --model sam2 --model-id "$MODEL_ID")" echo "source_hash=$source_hash" >> "$GITHUB_OUTPUT" echo "artifact_name=vision-${MODEL_SLUG}-onnx-v1-${source_hash:0:16}" >> "$GITHUB_OUTPUT" @@ -85,24 +57,23 @@ jobs: id: existing-artifact shell: bash env: - FORCE_EXPORT: ${{ steps.changes.outputs.force }} + DEFAULT_BRANCH: ${{ github.event.repository.default_branch }} + EVENT_NAME: ${{ github.event_name }} GH_TOKEN: ${{ github.token }} + SOURCE_BRANCH: ${{ github.head_ref || github.ref_name }} + PULL_REQUEST: ${{ github.event.pull_request.number }} MODEL_ARTIFACT: ${{ steps.model-key.outputs.artifact_name }} run: | - set -euo pipefail - artifact_run_id="$( - gh api "repos/${GITHUB_REPOSITORY}/actions/artifacts?name=${MODEL_ARTIFACT}&per_page=100" \ - --jq '.artifacts | map(select(.expired | not)) | sort_by(.created_at) | last | .workflow_run.id // empty' - )" - if [[ "$FORCE_EXPORT" == "true" || -z "$artifact_run_id" ]]; then - echo "export=true" >> "$GITHUB_OUTPUT" - else - echo "export=false" >> "$GITHUB_OUTPUT" - fi - echo "run_id=$artifact_run_id" >> "$GITHUB_OUTPUT" + python .github/scripts/find_successful_artifact.py \ + --repository "$GITHUB_REPOSITORY" \ + --artifact-name "$MODEL_ARTIFACT" \ + --default-branch "$DEFAULT_BRANCH" \ + --event-name "$EVENT_NAME" \ + --source-branch "$SOURCE_BRANCH" \ + --pull-request "$PULL_REQUEST" - name: Reclaim runner disk space - if: steps.existing-artifact.outputs.export == 'true' + if: steps.existing-artifact.outputs.found != 'true' shell: bash run: | set -euxo pipefail @@ -115,13 +86,13 @@ jobs: df -h / - name: Set up Python - if: steps.existing-artifact.outputs.export == 'true' + if: steps.existing-artifact.outputs.found != 'true' uses: actions/setup-python@v7 with: python-version: '3.12' - name: Install SAM2 export dependencies - if: steps.existing-artifact.outputs.export == 'true' + if: steps.existing-artifact.outputs.found != 'true' shell: bash run: | set -euxo pipefail @@ -138,7 +109,7 @@ jobs: "sam-2 @ git+https://github.com/facebookresearch/sam2.git" - name: Export SAM2 ONNX - if: steps.existing-artifact.outputs.export == 'true' + if: steps.existing-artifact.outputs.found != 'true' shell: bash env: MODEL_ID: ${{ inputs.model_id }} @@ -187,7 +158,7 @@ jobs: PY - name: Upload SAM2 ONNX - if: steps.existing-artifact.outputs.export == 'true' + if: steps.existing-artifact.outputs.found != 'true' uses: actions/upload-artifact@v7 with: name: ${{ steps.model-key.outputs.artifact_name }} @@ -199,11 +170,11 @@ jobs: id: final-artifact shell: bash env: - EXPORTED: ${{ steps.existing-artifact.outputs.export }} + FOUND: ${{ steps.existing-artifact.outputs.found }} EXISTING_RUN_ID: ${{ steps.existing-artifact.outputs.run_id }} run: | set -euo pipefail - if [[ "$EXPORTED" == "true" ]]; then + if [[ "$FOUND" != "true" ]]; then echo "run_id=$GITHUB_RUN_ID" >> "$GITHUB_OUTPUT" else echo "run_id=$EXISTING_RUN_ID" >> "$GITHUB_OUTPUT" diff --git a/.github/workflows/export-whisper-model.yml b/.github/workflows/export-whisper-model.yml index 0f895e8..8f5ddf4 100644 --- a/.github/workflows/export-whisper-model.yml +++ b/.github/workflows/export-whisper-model.yml @@ -52,34 +52,12 @@ jobs: fetch-depth: 0 ref: ${{ inputs.source_ref }} - - name: Determine export requirement - id: changes - shell: bash - env: - BASE_SHA: ${{ github.event.pull_request.base.sha || github.event.before }} - EVENT_NAME: ${{ github.event_name }} - SOURCE_REF: ${{ inputs.source_ref }} - run: | - set -euo pipefail - if [[ "$EVENT_NAME" == "workflow_dispatch" ]]; then - force=false - elif [[ -z "$BASE_SHA" ]] || ! git cat-file -e "${BASE_SHA}^{commit}"; then - force=true - elif git diff --quiet "$BASE_SHA" "$SOURCE_REF" -- \ - 'asr/whisper/*.py' \ - 'asr/whisper/**/*.py' \ - '.github/workflows/export-whisper-model.yml' \ - '.github/workflows/whisper-asr.yml'; then - force=false - else - force=true - fi - echo "force=$force" >> "$GITHUB_OUTPUT" - - name: Compute artifact names id: artifact-names shell: bash env: + MODEL_ID: ${{ inputs.model_id }} + ATTENTION: ${{ inputs.attention }} FALLBACK_PRECISION: ${{ inputs.fallback_precision }} MODEL_SLUG: ${{ inputs.model_slug }} PRECISION: ${{ inputs.precision }} @@ -95,13 +73,14 @@ jobs: *) echo "Unsupported fallback precision: $FALLBACK_PRECISION" >&2; exit 2 ;; esac fi - echo "primary=asr-${MODEL_SLUG}-onnx-${PRECISION}" >> "$GITHUB_OUTPUT" + source_hash="$(python .github/scripts/model_ci.py key --model "whisper" --model-id "$MODEL_ID" --attention "$ATTENTION")" + echo "source_hash=$source_hash" >> "$GITHUB_OUTPUT" + echo "primary=asr-${MODEL_SLUG}-onnx-${PRECISION}-${source_hash}" >> "$GITHUB_OUTPUT" if [[ -n "$FALLBACK_PRECISION" ]]; then - echo "fallback=asr-${MODEL_SLUG}-onnx-${FALLBACK_PRECISION}" >> "$GITHUB_OUTPUT" + echo "fallback=asr-${MODEL_SLUG}-onnx-${FALLBACK_PRECISION}-${source_hash}" >> "$GITHUB_OUTPUT" fi - name: Find successful primary artifact - if: steps.changes.outputs.force != 'true' id: primary-artifact shell: bash env: @@ -109,17 +88,18 @@ jobs: EVENT_NAME: ${{ github.event_name }} GH_TOKEN: ${{ github.token }} SOURCE_BRANCH: ${{ github.head_ref || github.ref_name }} + PULL_REQUEST: ${{ github.event.pull_request.number }} run: | python .github/scripts/find_successful_artifact.py \ --repository "$GITHUB_REPOSITORY" \ --artifact-name "${{ steps.artifact-names.outputs.primary }}" \ --default-branch "$DEFAULT_BRANCH" \ --event-name "$EVENT_NAME" \ - --source-branch "$SOURCE_BRANCH" + --source-branch "$SOURCE_BRANCH" \ + --pull-request "$PULL_REQUEST" - name: Find successful fallback artifact if: >- - steps.changes.outputs.force != 'true' && steps.primary-artifact.outputs.found != 'true' && inputs.fallback_precision != '' id: fallback-artifact @@ -129,13 +109,15 @@ jobs: EVENT_NAME: ${{ github.event_name }} GH_TOKEN: ${{ github.token }} SOURCE_BRANCH: ${{ github.head_ref || github.ref_name }} + PULL_REQUEST: ${{ github.event.pull_request.number }} run: | python .github/scripts/find_successful_artifact.py \ --repository "$GITHUB_REPOSITORY" \ --artifact-name "${{ steps.artifact-names.outputs.fallback }}" \ --default-branch "$DEFAULT_BRANCH" \ --event-name "$EVENT_NAME" \ - --source-branch "$SOURCE_BRANCH" + --source-branch "$SOURCE_BRANCH" \ + --pull-request "$PULL_REQUEST" - name: Select export action id: selection @@ -144,15 +126,12 @@ jobs: FALLBACK_FOUND: ${{ steps.fallback-artifact.outputs.found }} FALLBACK_NAME: ${{ steps.fallback-artifact.outputs.name }} FALLBACK_RUN_ID: ${{ steps.fallback-artifact.outputs.run_id }} - FORCE_EXPORT: ${{ steps.changes.outputs.force }} PRIMARY_FOUND: ${{ steps.primary-artifact.outputs.found }} PRIMARY_NAME: ${{ steps.primary-artifact.outputs.name }} PRIMARY_RUN_ID: ${{ steps.primary-artifact.outputs.run_id }} run: | set -euo pipefail - if [[ "$FORCE_EXPORT" == "true" ]]; then - echo "export=true" >> "$GITHUB_OUTPUT" - elif [[ "$PRIMARY_FOUND" == "true" ]]; then + if [[ "$PRIMARY_FOUND" == "true" ]]; then echo "export=false" >> "$GITHUB_OUTPUT" echo "name=$PRIMARY_NAME" >> "$GITHUB_OUTPUT" echo "run_id=$PRIMARY_RUN_ID" >> "$GITHUB_OUTPUT" @@ -204,6 +183,7 @@ jobs: MODEL_SLUG: ${{ inputs.model_slug }} PRECISION: ${{ inputs.precision }} SOURCE_REF: ${{ inputs.source_ref }} + SOURCE_HASH: ${{ steps.artifact-names.outputs.source_hash }} run: | set -uo pipefail run_export() { @@ -250,7 +230,7 @@ jobs: json.dumps(manifest, indent=2) + "\n", encoding="utf-8" ) PY - echo "name=asr-${MODEL_SLUG}-onnx-${selected_precision}" >> "$GITHUB_OUTPUT" + echo "name=asr-${MODEL_SLUG}-onnx-${selected_precision}-${SOURCE_HASH}" >> "$GITHUB_OUTPUT" - name: Upload Whisper ONNX if: steps.selection.outputs.export == 'true' diff --git a/.github/workflows/model-changes.yml b/.github/workflows/model-changes.yml new file mode 100644 index 0000000..d945b3e --- /dev/null +++ b/.github/workflows/model-changes.yml @@ -0,0 +1,32 @@ +name: Detect model changes + +on: + workflow_call: + outputs: + whisper: + value: ${{ jobs.changes.outputs.whisper }} + parakeet: + value: ${{ jobs.changes.outputs.parakeet }} + nemotron: + value: ${{ jobs.changes.outputs.nemotron }} + sam2: + value: ${{ jobs.changes.outputs.sam2 }} + +permissions: + contents: read + +jobs: + changes: + runs-on: ubuntu-24.04 + outputs: + whisper: ${{ steps.changes.outputs.whisper }} + parakeet: ${{ steps.changes.outputs.parakeet }} + nemotron: ${{ steps.changes.outputs.nemotron }} + sam2: ${{ steps.changes.outputs.sam2 }} + steps: + - uses: actions/checkout@v7 + with: + fetch-depth: 0 + ref: ${{ github.event.pull_request.head.sha || github.sha }} + - id: changes + run: python .github/scripts/model_ci.py changes diff --git a/.github/workflows/nvidia-asr.yml b/.github/workflows/nvidia-asr.yml index 780708c..0eb6940 100644 --- a/.github/workflows/nvidia-asr.yml +++ b/.github/workflows/nvidia-asr.yml @@ -4,6 +4,7 @@ on: push: branches: [main] pull_request: + types: [opened, synchronize, reopened, ready_for_review, converted_to_draft] workflow_dispatch: inputs: build_run_id: @@ -20,7 +21,13 @@ concurrency: cancel-in-progress: true jobs: + changes: + if: github.event_name != 'pull_request' || !github.event.pull_request.draft + uses: ./.github/workflows/model-changes.yml + requirements: + needs: changes + if: needs.changes.outputs.parakeet == 'true' || needs.changes.outputs.nemotron == 'true' name: Wait for Build runs-on: ubuntu-24.04 timeout-minutes: 100 @@ -62,8 +69,9 @@ jobs: echo "source_ref=$source_ref" >> "$GITHUB_OUTPUT" export-parakeet: + if: needs.changes.outputs.parakeet == 'true' name: Export Parakeet TDT ONNX - needs: requirements + needs: [requirements, changes] uses: ./.github/workflows/export-nvidia-asr-model.yml with: model_slug: parakeet-tdt @@ -75,8 +83,9 @@ jobs: source_ref: ${{ needs.requirements.outputs.source_ref }} export-nemotron: + if: needs.changes.outputs.nemotron == 'true' name: Export Nemotron 3.5 ASR Streaming ONNX - needs: requirements + needs: [requirements, changes] uses: ./.github/workflows/export-nvidia-asr-model.yml with: model_slug: nemotron-3.5-asr-streaming-0.6b @@ -89,7 +98,10 @@ jobs: artifacts: name: Collect NVIDIA ASR artifacts - needs: [export-parakeet, export-nemotron] + needs: [requirements, export-parakeet, export-nemotron] + if: >- + !cancelled() && !failure() && + (needs.export-parakeet.result == 'success' || needs.export-nemotron.result == 'success') runs-on: ubuntu-24.04 outputs: matrix: ${{ steps.matrix.outputs.matrix }} @@ -127,6 +139,7 @@ jobs: }, ] } + matrix["include"] = [model for model in matrix["include"] if model["artifact_name"]] with Path(os.environ["GITHUB_OUTPUT"]).open("a", encoding="utf-8") as output: output.write(f"matrix={json.dumps(matrix, separators=(',', ':'))}\n") PY diff --git a/.github/workflows/sam2.yml b/.github/workflows/sam2.yml index cb12fa8..d6b6cee 100644 --- a/.github/workflows/sam2.yml +++ b/.github/workflows/sam2.yml @@ -4,6 +4,7 @@ on: push: branches: [main] pull_request: + types: [opened, synchronize, reopened, ready_for_review, converted_to_draft] workflow_dispatch: inputs: build_run_id: @@ -20,7 +21,13 @@ concurrency: cancel-in-progress: true jobs: + changes: + if: github.event_name != 'pull_request' || !github.event.pull_request.draft + uses: ./.github/workflows/model-changes.yml + requirements: + needs: changes + if: needs.changes.outputs.sam2 == 'true' name: Wait for Build runs-on: ubuntu-24.04 timeout-minutes: 100 @@ -105,25 +112,29 @@ jobs: shell: bash env: GH_TOKEN: ${{ github.token }} + DEFAULT_BRANCH: ${{ github.event.repository.default_branch }} + EVENT_NAME: ${{ github.event_name }} + SOURCE_BRANCH: ${{ github.head_ref || github.ref_name }} + PULL_REQUEST: ${{ github.event.pull_request.number }} run: | set -euo pipefail - source_hash="$( - git ls-files -z -- 'vision/sam2/python/*.py' 'vision/sam2/python/**/*.py' \ - | xargs -0 sha256sum \ - | sha256sum \ - | cut -d' ' -f1 - )" models='{}' for slug in \ sam2.1-hiera-tiny \ sam2.1-hiera-small \ sam2.1-hiera-base-plus \ sam2.1-hiera-large; do + source_hash="$(python .github/scripts/model_ci.py key --model sam2 --model-id "facebook/${slug}")" artifact_name="vision-${slug}-onnx-v1-${source_hash:0:16}" - artifact_run_id="$( - gh api "repos/${GITHUB_REPOSITORY}/actions/artifacts?name=${artifact_name}&per_page=100" \ - --jq '.artifacts | map(select(.expired | not)) | sort_by(.created_at) | last | .workflow_run.id // empty' - )" + artifact_run_id="$(gh api "repos/${GITHUB_REPOSITORY}/actions/runs/${GITHUB_RUN_ID}/artifacts?per_page=100" \ + --jq ".artifacts[] | select(.name == \"$artifact_name\" and (.expired | not)) | \"$GITHUB_RUN_ID\"")" + if [[ -z "$artifact_run_id" ]]; then + artifact="$(python .github/scripts/find_successful_artifact.py \ + --repository "$GITHUB_REPOSITORY" --artifact-name "$artifact_name" \ + --default-branch "$DEFAULT_BRANCH" --event-name "$EVENT_NAME" \ + --source-branch "$SOURCE_BRANCH" --pull-request "$PULL_REQUEST" --json)" + artifact_run_id="$(jq -r '.run_id' <<< "$artifact")" + fi test -n "$artifact_run_id" models="$(jq \ --arg slug "$slug" \ diff --git a/.github/workflows/whisper-asr.yml b/.github/workflows/whisper-asr.yml index 6d7d64f..0c1e59f 100644 --- a/.github/workflows/whisper-asr.yml +++ b/.github/workflows/whisper-asr.yml @@ -4,6 +4,7 @@ on: push: branches: [main] pull_request: + types: [opened, synchronize, reopened, ready_for_review, converted_to_draft] workflow_dispatch: inputs: build_run_id: @@ -20,7 +21,13 @@ concurrency: cancel-in-progress: true jobs: + changes: + if: github.event_name != 'pull_request' || !github.event.pull_request.draft + uses: ./.github/workflows/model-changes.yml + requirements: + needs: changes + if: needs.changes.outputs.whisper == 'true' name: Wait for Build runs-on: ubuntu-24.04 timeout-minutes: 100 From f7bbe1719d163257a1038f071f2912e369af9bb9 Mon Sep 17 00:00:00 2001 From: contentis Date: Wed, 23 Sep 2026 10:48:30 +0200 Subject: [PATCH 4/8] Ignore skipped draft builds when resolving model test inputs --- .github/workflows/nvidia-asr.yml | 2 +- .github/workflows/sam2.yml | 2 +- .github/workflows/whisper-asr.yml | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/.github/workflows/nvidia-asr.yml b/.github/workflows/nvidia-asr.yml index 0eb6940..501f3d0 100644 --- a/.github/workflows/nvidia-asr.yml +++ b/.github/workflows/nvidia-asr.yml @@ -50,7 +50,7 @@ jobs: while [[ -z "$run_id" ]]; do run_id="$( gh api "repos/${GITHUB_REPOSITORY}/actions/workflows/ci.yml/runs?per_page=100" \ - --jq ".workflow_runs | map(select(.head_sha == \"$SOURCE_SHA\" and .event == \"$EVENT_NAME\")) | sort_by(.created_at) | last | .id // empty" + --jq ".workflow_runs | map(select(.head_sha == \"$SOURCE_SHA\" and .event == \"$EVENT_NAME\" and .conclusion != \"skipped\")) | sort_by(.created_at) | last | .id // empty" )" [[ -n "$run_id" ]] || sleep 10 done diff --git a/.github/workflows/sam2.yml b/.github/workflows/sam2.yml index d6b6cee..3149e44 100644 --- a/.github/workflows/sam2.yml +++ b/.github/workflows/sam2.yml @@ -50,7 +50,7 @@ jobs: while [[ -z "$run_id" ]]; do run_id="$( gh api "repos/${GITHUB_REPOSITORY}/actions/workflows/ci.yml/runs?per_page=100" \ - --jq ".workflow_runs | map(select(.head_sha == \"$SOURCE_SHA\" and .event == \"$EVENT_NAME\")) | sort_by(.created_at) | last | .id // empty" + --jq ".workflow_runs | map(select(.head_sha == \"$SOURCE_SHA\" and .event == \"$EVENT_NAME\" and .conclusion != \"skipped\")) | sort_by(.created_at) | last | .id // empty" )" [[ -n "$run_id" ]] || sleep 10 done diff --git a/.github/workflows/whisper-asr.yml b/.github/workflows/whisper-asr.yml index 0c1e59f..3fdd8de 100644 --- a/.github/workflows/whisper-asr.yml +++ b/.github/workflows/whisper-asr.yml @@ -50,7 +50,7 @@ jobs: while [[ -z "$run_id" ]]; do run_id="$( gh api "repos/${GITHUB_REPOSITORY}/actions/workflows/ci.yml/runs?per_page=100" \ - --jq ".workflow_runs | map(select(.head_sha == \"$SOURCE_SHA\" and .event == \"$EVENT_NAME\")) | sort_by(.created_at) | last | .id // empty" + --jq ".workflow_runs | map(select(.head_sha == \"$SOURCE_SHA\" and .event == \"$EVENT_NAME\" and .conclusion != \"skipped\")) | sort_by(.created_at) | last | .id // empty" )" [[ -n "$run_id" ]] || sleep 10 done From a310a436126b46d99842282f7cf97274adab2b2f Mon Sep 17 00:00:00 2001 From: contentis Date: Wed, 23 Sep 2026 11:00:27 +0200 Subject: [PATCH 5/8] Move CI changes to a separate PR [skip ci] --- .github/scripts/find_successful_artifact.py | 23 ++-- .github/scripts/model_ci.py | 112 ------------------ .github/workflows/ci.yml | 5 - .github/workflows/export-nvidia-asr-model.yml | 49 +++++--- .github/workflows/export-sam2-model.yml | 69 +++++++---- .github/workflows/export-whisper-model.yml | 50 +++++--- .github/workflows/model-changes.yml | 32 ----- .github/workflows/nvidia-asr.yml | 21 +--- .github/workflows/sam2.yml | 33 ++---- .github/workflows/whisper-asr.yml | 9 +- 10 files changed, 143 insertions(+), 260 deletions(-) delete mode 100644 .github/scripts/model_ci.py delete mode 100644 .github/workflows/model-changes.yml diff --git a/.github/scripts/find_successful_artifact.py b/.github/scripts/find_successful_artifact.py index 4174f0e..82338ee 100644 --- a/.github/scripts/find_successful_artifact.py +++ b/.github/scripts/find_successful_artifact.py @@ -25,10 +25,7 @@ def github_api(path: str) -> object: return json.load(response) -def write_outputs(*, found: bool, artifact_name: str, run_id: str = "", as_json: bool = False) -> None: - if as_json: - print(json.dumps({"found": found, "name": artifact_name if found else "", "run_id": run_id})) - return +def write_outputs(*, found: bool, artifact_name: str, run_id: str = "") -> None: output_path = Path(os.environ["GITHUB_OUTPUT"]) with output_path.open("a", encoding="utf-8") as output: output.write(f"found={'true' if found else 'false'}\n") @@ -43,15 +40,16 @@ def main() -> None: parser.add_argument("--default-branch", required=True) parser.add_argument("--event-name", required=True) parser.add_argument("--source-branch", required=True) - parser.add_argument("--pull-request", default="") - parser.add_argument("--json", action="store_true") args = parser.parse_args() artifact_name = args.artifact_name allowed_branches = {args.default_branch} - if args.event_name != "pull_request" or args.pull_request: + if args.event_name != "pull_request": allowed_branches.add(args.source_branch) - response = github_api(f"repos/{args.repository}/actions/artifacts?name={quote(artifact_name)}&per_page=100") + response = github_api( + f"repos/{args.repository}/actions/artifacts" + f"?name={quote(artifact_name)}&per_page=100" + ) artifacts = sorted( ( artifact @@ -69,23 +67,18 @@ def main() -> None: run_id = int(artifact["workflow_run"]["id"]) if run_id not in conclusions: run = github_api(f"repos/{args.repository}/actions/runs/{run_id}") - same_pr = args.pull_request and any( - str(pr["number"]) == args.pull_request for pr in run.get("pull_requests", []) - ) - trusted_branch = run["event"] != "pull_request" and run["head_branch"] in allowed_branches conclusions[run_id] = ( - (trusted_branch or same_pr) and run["status"] == "completed" and run["conclusion"] == "success" + run["status"] == "completed" and run["conclusion"] == "success" ) if conclusions[run_id]: write_outputs( found=True, artifact_name=artifact_name, run_id=str(run_id), - as_json=args.json, ) return - write_outputs(found=False, artifact_name=artifact_name, as_json=args.json) + write_outputs(found=False, artifact_name=artifact_name) if __name__ == "__main__": diff --git a/.github/scripts/model_ci.py b/.github/scripts/model_ci.py deleted file mode 100644 index 207a4eb..0000000 --- a/.github/scripts/model_ci.py +++ /dev/null @@ -1,112 +0,0 @@ -#!/usr/bin/env python3 -"""Select affected model tests and fingerprint their ONNX export inputs.""" - -import argparse -import fnmatch -import hashlib -import json -import os -import subprocess -from pathlib import Path - -SHARED_RUNTIME = [ - "CMakeLists.txt", - "CMakePresets.json", - "cmake/*", - "common/*.cpp", - "common/*.h", - "common/*CMakeLists.txt", - ".github/workflows/ci.yml", - ".github/scripts/*", - ".github/workflows/model-changes.yml", -] -EXPORTS = { - "whisper": [ - "asr/whisper/model_export/export_whisper.py", - "common/model_export/*.py", - ".github/workflows/export-whisper-model.yml", - ], - "parakeet": [ - "asr/rnnt/python/tools/export_parakeet_tdt_onnx.py", - "asr/rnnt/python/rnnt/parakeet_tdt/nemo_backend.py", - "asr/rnnt/python/rnnt/parakeet_tdt/__init__.py", - ], - "nemotron": [ - "asr/rnnt/python/tools/export_nemotron_onnx.py", - "asr/rnnt/python/rnnt/nemotron_asr/nemo_backend.py", - "asr/rnnt/python/rnnt/nemotron_asr/__init__.py", - ], - "sam2": [ - "vision/sam2/python/export_sam2_onnx.py", - "vision/sam2/python/modeling.py", - ".github/workflows/export-sam2-model.yml", - ], -} -for model in ("parakeet", "nemotron"): - EXPORTS[model] += [ - "asr/rnnt/python/tools/nemo_preprocessor_export.py", - "asr/rnnt/python/rnnt/__init__.py", - "asr/rnnt/python/rnnt/audio.py", - ".github/workflows/export-nvidia-asr-model.yml", - ] -RUNTIME = { - "whisper": ["asr/whisper/*", "assets/sample.wav", ".github/workflows/whisper-asr.yml"], - "sam2": ["vision/sam2/*", "assets/sam2-*", ".github/workflows/sam2.yml"], -} -for model, package in (("parakeet", "parakeet_tdt"), ("nemotron", "nemotron_asr")): - RUNTIME[model] = [ - f"asr/rnnt/cpp/{model}*", - "asr/rnnt/cpp/asr*", - "asr/rnnt/CMakeLists.txt", - f"asr/rnnt/python/rnnt/{package}/*", - "asr/rnnt/python/rnnt/validation/*", - "asr/rnnt/python/tests/*", - "assets/sample.wav", - ".github/workflows/nvidia-asr.yml", - ] - - -def affected(model, files): - patterns = SHARED_RUNTIME + RUNTIME[model] + EXPORTS[model] - return any(fnmatch.fnmatchcase(path, pattern) for path in files if not path.endswith(".md") for pattern in patterns) - - -def export_key(model, model_id, attention): - # Index blob IDs are independent of checkout line endings and include deleted/renamed inputs. - sources = subprocess.check_output(["git", "ls-files", "--stage", "--", *EXPORTS[model]]) - options = json.dumps([model, model_id, attention]).encode() - return hashlib.sha256(sources + options).hexdigest()[:16] - - -def main(): - parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument("mode", choices=["changes", "key"]) - parser.add_argument("--model", choices=EXPORTS) - parser.add_argument("--model-id", default="") - parser.add_argument("--attention", default="") - args = parser.parse_args() - if args.mode == "key": - if not args.model or not args.model_id: - parser.error("key requires --model and --model-id") - print(export_key(args.model, args.model_id, args.attention)) - return - event = json.loads(Path(os.environ["GITHUB_EVENT_PATH"]).read_text()) - pr = event.get("pull_request") - base = pr["base"]["sha"] if pr else event.get("before") - files = None - if base and subprocess.run(["git", "cat-file", "-e", f"{base}^{{commit}}"], capture_output=True).returncode == 0: - if pr: - base = subprocess.check_output(["git", "merge-base", base, "HEAD"], text=True).strip() - files = ( - subprocess.check_output(["git", "diff", "--no-renames", "--name-only", "-z", base, "HEAD"]) - .decode() - .split("\0") - ) - with Path(os.environ["GITHUB_OUTPUT"]).open("a", encoding="utf-8") as output: - for model in EXPORTS: - changed = files is None or affected(model, files) - output.write(f"{model}={str(changed).lower()}\n") - - -if __name__ == "__main__": - main() diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index aae8c6c..5d0965e 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -4,7 +4,6 @@ on: push: branches: [main] pull_request: - types: [opened, synchronize, reopened, ready_for_review, converted_to_draft] workflow_dispatch: permissions: @@ -20,7 +19,6 @@ env: jobs: clang-format: - if: github.event_name != 'pull_request' || !github.event.pull_request.draft name: clang-format runs-on: ubuntu-24.04 @@ -49,7 +47,6 @@ jobs: | xargs -0 -r clang-format-22 --dry-run --Werror -- linux-x64: - if: github.event_name != 'pull_request' || !github.event.pull_request.draft name: Ubuntu 24.04 / Clang / x64 runs-on: ubuntu-24.04 container: nvidia/cuda:13.2.1-devel-ubuntu24.04 @@ -93,7 +90,6 @@ jobs: path: out/build/linux-x64/bin/Release/ linux-arm64: - if: github.event_name != 'pull_request' || !github.event.pull_request.draft name: Ubuntu 24.04 / Clang / ARM64 runs-on: ubuntu-24.04-arm container: nvidia/cuda:13.2.1-devel-ubuntu24.04 @@ -135,7 +131,6 @@ jobs: path: out/build/linux-arm64/bin/Release/ windows-x64: - if: github.event_name != 'pull_request' || !github.event.pull_request.draft name: Windows Server 2025 / MSVC / x64 runs-on: windows-2025 timeout-minutes: 90 diff --git a/.github/workflows/export-nvidia-asr-model.yml b/.github/workflows/export-nvidia-asr-model.yml index 7bc2a6c..dc098d1 100644 --- a/.github/workflows/export-nvidia-asr-model.yml +++ b/.github/workflows/export-nvidia-asr-model.yml @@ -51,11 +51,34 @@ jobs: fetch-depth: 0 ref: ${{ inputs.source_ref }} + - name: Determine export requirement + id: changes + shell: bash + env: + BASE_SHA: ${{ github.event.pull_request.base.sha || github.event.before }} + EVENT_NAME: ${{ github.event_name }} + SOURCE_REF: ${{ inputs.source_ref }} + run: | + set -euo pipefail + if [[ "$EVENT_NAME" == "workflow_dispatch" ]]; then + force=false + elif [[ -z "$BASE_SHA" ]] || ! git cat-file -e "${BASE_SHA}^{commit}"; then + force=true + elif git diff --quiet "$BASE_SHA" "$SOURCE_REF" -- \ + 'asr/rnnt/python/*.py' \ + 'asr/rnnt/python/**/*.py' \ + '.github/workflows/export-nvidia-asr-model.yml' \ + '.github/workflows/nvidia-asr.yml'; then + force=false + else + force=true + fi + echo "force=$force" >> "$GITHUB_OUTPUT" + - name: Compute artifact names id: artifact-names shell: bash env: - MODEL_ID: ${{ inputs.model_id }} FALLBACK_PRECISION: ${{ inputs.fallback_precision }} MODEL_SLUG: ${{ inputs.model_slug }} MODEL_TYPE: ${{ inputs.model_type }} @@ -76,14 +99,13 @@ jobs: *) echo "Unsupported fallback precision: $FALLBACK_PRECISION" >&2; exit 2 ;; esac fi - source_hash="$(python .github/scripts/model_ci.py key --model "${MODEL_TYPE}" --model-id "$MODEL_ID")" - echo "source_hash=$source_hash" >> "$GITHUB_OUTPUT" - echo "primary=asr-${MODEL_SLUG}-onnx-${PRECISION}-${source_hash}" >> "$GITHUB_OUTPUT" + echo "primary=asr-${MODEL_SLUG}-onnx-${PRECISION}" >> "$GITHUB_OUTPUT" if [[ -n "$FALLBACK_PRECISION" ]]; then - echo "fallback=asr-${MODEL_SLUG}-onnx-${FALLBACK_PRECISION}-${source_hash}" >> "$GITHUB_OUTPUT" + echo "fallback=asr-${MODEL_SLUG}-onnx-${FALLBACK_PRECISION}" >> "$GITHUB_OUTPUT" fi - name: Find successful primary artifact + if: steps.changes.outputs.force != 'true' id: primary-artifact shell: bash env: @@ -91,18 +113,17 @@ jobs: EVENT_NAME: ${{ github.event_name }} GH_TOKEN: ${{ github.token }} SOURCE_BRANCH: ${{ github.head_ref || github.ref_name }} - PULL_REQUEST: ${{ github.event.pull_request.number }} run: | python .github/scripts/find_successful_artifact.py \ --repository "$GITHUB_REPOSITORY" \ --artifact-name "${{ steps.artifact-names.outputs.primary }}" \ --default-branch "$DEFAULT_BRANCH" \ --event-name "$EVENT_NAME" \ - --source-branch "$SOURCE_BRANCH" \ - --pull-request "$PULL_REQUEST" + --source-branch "$SOURCE_BRANCH" - name: Find successful fallback artifact if: >- + steps.changes.outputs.force != 'true' && steps.primary-artifact.outputs.found != 'true' && inputs.fallback_precision != '' id: fallback-artifact @@ -112,15 +133,13 @@ jobs: EVENT_NAME: ${{ github.event_name }} GH_TOKEN: ${{ github.token }} SOURCE_BRANCH: ${{ github.head_ref || github.ref_name }} - PULL_REQUEST: ${{ github.event.pull_request.number }} run: | python .github/scripts/find_successful_artifact.py \ --repository "$GITHUB_REPOSITORY" \ --artifact-name "${{ steps.artifact-names.outputs.fallback }}" \ --default-branch "$DEFAULT_BRANCH" \ --event-name "$EVENT_NAME" \ - --source-branch "$SOURCE_BRANCH" \ - --pull-request "$PULL_REQUEST" + --source-branch "$SOURCE_BRANCH" - name: Select export action id: selection @@ -129,12 +148,15 @@ jobs: FALLBACK_FOUND: ${{ steps.fallback-artifact.outputs.found }} FALLBACK_NAME: ${{ steps.fallback-artifact.outputs.name }} FALLBACK_RUN_ID: ${{ steps.fallback-artifact.outputs.run_id }} + FORCE_EXPORT: ${{ steps.changes.outputs.force }} PRIMARY_FOUND: ${{ steps.primary-artifact.outputs.found }} PRIMARY_NAME: ${{ steps.primary-artifact.outputs.name }} PRIMARY_RUN_ID: ${{ steps.primary-artifact.outputs.run_id }} run: | set -euo pipefail - if [[ "$PRIMARY_FOUND" == "true" ]]; then + if [[ "$FORCE_EXPORT" == "true" ]]; then + echo "export=true" >> "$GITHUB_OUTPUT" + elif [[ "$PRIMARY_FOUND" == "true" ]]; then echo "export=false" >> "$GITHUB_OUTPUT" echo "name=$PRIMARY_NAME" >> "$GITHUB_OUTPUT" echo "run_id=$PRIMARY_RUN_ID" >> "$GITHUB_OUTPUT" @@ -191,7 +213,6 @@ jobs: PRECISION: ${{ inputs.precision }} PYTHONPATH: asr/rnnt/python SOURCE_REF: ${{ inputs.source_ref }} - SOURCE_HASH: ${{ steps.artifact-names.outputs.source_hash }} run: | set -uo pipefail export_dtype() { @@ -267,7 +288,7 @@ jobs: json.dumps(manifest, indent=2) + "\n", encoding="utf-8" ) PY - echo "name=asr-${MODEL_SLUG}-onnx-${selected_precision}-${SOURCE_HASH}" >> "$GITHUB_OUTPUT" + echo "name=asr-${MODEL_SLUG}-onnx-${selected_precision}" >> "$GITHUB_OUTPUT" - name: Upload NVIDIA ASR ONNX if: steps.selection.outputs.export == 'true' diff --git a/.github/workflows/export-sam2-model.yml b/.github/workflows/export-sam2-model.yml index c2df020..227015b 100644 --- a/.github/workflows/export-sam2-model.yml +++ b/.github/workflows/export-sam2-model.yml @@ -41,15 +41,43 @@ jobs: fetch-depth: 0 ref: ${{ inputs.source_ref }} + - name: Determine export requirement + id: changes + shell: bash + env: + BASE_SHA: ${{ github.event.pull_request.base.sha || github.event.before }} + EVENT_NAME: ${{ github.event_name }} + SOURCE_REF: ${{ inputs.source_ref }} + run: | + set -euo pipefail + if [[ "$EVENT_NAME" == "workflow_dispatch" ]]; then + force=false + elif [[ -z "$BASE_SHA" ]] || ! git cat-file -e "${BASE_SHA}^{commit}"; then + force=true + elif git diff --quiet "$BASE_SHA" "$SOURCE_REF" -- \ + 'vision/sam2/python/*.py' \ + 'vision/sam2/python/**/*.py' \ + '.github/workflows/export-sam2-model.yml' \ + '.github/workflows/sam2.yml'; then + force=false + else + force=true + fi + echo "force=$force" >> "$GITHUB_OUTPUT" + - name: Compute model artifact key id: model-key shell: bash env: MODEL_SLUG: ${{ inputs.model_slug }} - MODEL_ID: ${{ inputs.model_id }} run: | set -euo pipefail - source_hash="$(python .github/scripts/model_ci.py key --model sam2 --model-id "$MODEL_ID")" + source_hash="$( + git ls-files -z -- 'vision/sam2/python/*.py' 'vision/sam2/python/**/*.py' \ + | xargs -0 sha256sum \ + | sha256sum \ + | cut -d' ' -f1 + )" echo "source_hash=$source_hash" >> "$GITHUB_OUTPUT" echo "artifact_name=vision-${MODEL_SLUG}-onnx-v1-${source_hash:0:16}" >> "$GITHUB_OUTPUT" @@ -57,23 +85,24 @@ jobs: id: existing-artifact shell: bash env: - DEFAULT_BRANCH: ${{ github.event.repository.default_branch }} - EVENT_NAME: ${{ github.event_name }} + FORCE_EXPORT: ${{ steps.changes.outputs.force }} GH_TOKEN: ${{ github.token }} - SOURCE_BRANCH: ${{ github.head_ref || github.ref_name }} - PULL_REQUEST: ${{ github.event.pull_request.number }} MODEL_ARTIFACT: ${{ steps.model-key.outputs.artifact_name }} run: | - python .github/scripts/find_successful_artifact.py \ - --repository "$GITHUB_REPOSITORY" \ - --artifact-name "$MODEL_ARTIFACT" \ - --default-branch "$DEFAULT_BRANCH" \ - --event-name "$EVENT_NAME" \ - --source-branch "$SOURCE_BRANCH" \ - --pull-request "$PULL_REQUEST" + set -euo pipefail + artifact_run_id="$( + gh api "repos/${GITHUB_REPOSITORY}/actions/artifacts?name=${MODEL_ARTIFACT}&per_page=100" \ + --jq '.artifacts | map(select(.expired | not)) | sort_by(.created_at) | last | .workflow_run.id // empty' + )" + if [[ "$FORCE_EXPORT" == "true" || -z "$artifact_run_id" ]]; then + echo "export=true" >> "$GITHUB_OUTPUT" + else + echo "export=false" >> "$GITHUB_OUTPUT" + fi + echo "run_id=$artifact_run_id" >> "$GITHUB_OUTPUT" - name: Reclaim runner disk space - if: steps.existing-artifact.outputs.found != 'true' + if: steps.existing-artifact.outputs.export == 'true' shell: bash run: | set -euxo pipefail @@ -86,13 +115,13 @@ jobs: df -h / - name: Set up Python - if: steps.existing-artifact.outputs.found != 'true' + if: steps.existing-artifact.outputs.export == 'true' uses: actions/setup-python@v7 with: python-version: '3.12' - name: Install SAM2 export dependencies - if: steps.existing-artifact.outputs.found != 'true' + if: steps.existing-artifact.outputs.export == 'true' shell: bash run: | set -euxo pipefail @@ -109,7 +138,7 @@ jobs: "sam-2 @ git+https://github.com/facebookresearch/sam2.git" - name: Export SAM2 ONNX - if: steps.existing-artifact.outputs.found != 'true' + if: steps.existing-artifact.outputs.export == 'true' shell: bash env: MODEL_ID: ${{ inputs.model_id }} @@ -158,7 +187,7 @@ jobs: PY - name: Upload SAM2 ONNX - if: steps.existing-artifact.outputs.found != 'true' + if: steps.existing-artifact.outputs.export == 'true' uses: actions/upload-artifact@v7 with: name: ${{ steps.model-key.outputs.artifact_name }} @@ -170,11 +199,11 @@ jobs: id: final-artifact shell: bash env: - FOUND: ${{ steps.existing-artifact.outputs.found }} + EXPORTED: ${{ steps.existing-artifact.outputs.export }} EXISTING_RUN_ID: ${{ steps.existing-artifact.outputs.run_id }} run: | set -euo pipefail - if [[ "$FOUND" != "true" ]]; then + if [[ "$EXPORTED" == "true" ]]; then echo "run_id=$GITHUB_RUN_ID" >> "$GITHUB_OUTPUT" else echo "run_id=$EXISTING_RUN_ID" >> "$GITHUB_OUTPUT" diff --git a/.github/workflows/export-whisper-model.yml b/.github/workflows/export-whisper-model.yml index 8f5ddf4..0f895e8 100644 --- a/.github/workflows/export-whisper-model.yml +++ b/.github/workflows/export-whisper-model.yml @@ -52,12 +52,34 @@ jobs: fetch-depth: 0 ref: ${{ inputs.source_ref }} + - name: Determine export requirement + id: changes + shell: bash + env: + BASE_SHA: ${{ github.event.pull_request.base.sha || github.event.before }} + EVENT_NAME: ${{ github.event_name }} + SOURCE_REF: ${{ inputs.source_ref }} + run: | + set -euo pipefail + if [[ "$EVENT_NAME" == "workflow_dispatch" ]]; then + force=false + elif [[ -z "$BASE_SHA" ]] || ! git cat-file -e "${BASE_SHA}^{commit}"; then + force=true + elif git diff --quiet "$BASE_SHA" "$SOURCE_REF" -- \ + 'asr/whisper/*.py' \ + 'asr/whisper/**/*.py' \ + '.github/workflows/export-whisper-model.yml' \ + '.github/workflows/whisper-asr.yml'; then + force=false + else + force=true + fi + echo "force=$force" >> "$GITHUB_OUTPUT" + - name: Compute artifact names id: artifact-names shell: bash env: - MODEL_ID: ${{ inputs.model_id }} - ATTENTION: ${{ inputs.attention }} FALLBACK_PRECISION: ${{ inputs.fallback_precision }} MODEL_SLUG: ${{ inputs.model_slug }} PRECISION: ${{ inputs.precision }} @@ -73,14 +95,13 @@ jobs: *) echo "Unsupported fallback precision: $FALLBACK_PRECISION" >&2; exit 2 ;; esac fi - source_hash="$(python .github/scripts/model_ci.py key --model "whisper" --model-id "$MODEL_ID" --attention "$ATTENTION")" - echo "source_hash=$source_hash" >> "$GITHUB_OUTPUT" - echo "primary=asr-${MODEL_SLUG}-onnx-${PRECISION}-${source_hash}" >> "$GITHUB_OUTPUT" + echo "primary=asr-${MODEL_SLUG}-onnx-${PRECISION}" >> "$GITHUB_OUTPUT" if [[ -n "$FALLBACK_PRECISION" ]]; then - echo "fallback=asr-${MODEL_SLUG}-onnx-${FALLBACK_PRECISION}-${source_hash}" >> "$GITHUB_OUTPUT" + echo "fallback=asr-${MODEL_SLUG}-onnx-${FALLBACK_PRECISION}" >> "$GITHUB_OUTPUT" fi - name: Find successful primary artifact + if: steps.changes.outputs.force != 'true' id: primary-artifact shell: bash env: @@ -88,18 +109,17 @@ jobs: EVENT_NAME: ${{ github.event_name }} GH_TOKEN: ${{ github.token }} SOURCE_BRANCH: ${{ github.head_ref || github.ref_name }} - PULL_REQUEST: ${{ github.event.pull_request.number }} run: | python .github/scripts/find_successful_artifact.py \ --repository "$GITHUB_REPOSITORY" \ --artifact-name "${{ steps.artifact-names.outputs.primary }}" \ --default-branch "$DEFAULT_BRANCH" \ --event-name "$EVENT_NAME" \ - --source-branch "$SOURCE_BRANCH" \ - --pull-request "$PULL_REQUEST" + --source-branch "$SOURCE_BRANCH" - name: Find successful fallback artifact if: >- + steps.changes.outputs.force != 'true' && steps.primary-artifact.outputs.found != 'true' && inputs.fallback_precision != '' id: fallback-artifact @@ -109,15 +129,13 @@ jobs: EVENT_NAME: ${{ github.event_name }} GH_TOKEN: ${{ github.token }} SOURCE_BRANCH: ${{ github.head_ref || github.ref_name }} - PULL_REQUEST: ${{ github.event.pull_request.number }} run: | python .github/scripts/find_successful_artifact.py \ --repository "$GITHUB_REPOSITORY" \ --artifact-name "${{ steps.artifact-names.outputs.fallback }}" \ --default-branch "$DEFAULT_BRANCH" \ --event-name "$EVENT_NAME" \ - --source-branch "$SOURCE_BRANCH" \ - --pull-request "$PULL_REQUEST" + --source-branch "$SOURCE_BRANCH" - name: Select export action id: selection @@ -126,12 +144,15 @@ jobs: FALLBACK_FOUND: ${{ steps.fallback-artifact.outputs.found }} FALLBACK_NAME: ${{ steps.fallback-artifact.outputs.name }} FALLBACK_RUN_ID: ${{ steps.fallback-artifact.outputs.run_id }} + FORCE_EXPORT: ${{ steps.changes.outputs.force }} PRIMARY_FOUND: ${{ steps.primary-artifact.outputs.found }} PRIMARY_NAME: ${{ steps.primary-artifact.outputs.name }} PRIMARY_RUN_ID: ${{ steps.primary-artifact.outputs.run_id }} run: | set -euo pipefail - if [[ "$PRIMARY_FOUND" == "true" ]]; then + if [[ "$FORCE_EXPORT" == "true" ]]; then + echo "export=true" >> "$GITHUB_OUTPUT" + elif [[ "$PRIMARY_FOUND" == "true" ]]; then echo "export=false" >> "$GITHUB_OUTPUT" echo "name=$PRIMARY_NAME" >> "$GITHUB_OUTPUT" echo "run_id=$PRIMARY_RUN_ID" >> "$GITHUB_OUTPUT" @@ -183,7 +204,6 @@ jobs: MODEL_SLUG: ${{ inputs.model_slug }} PRECISION: ${{ inputs.precision }} SOURCE_REF: ${{ inputs.source_ref }} - SOURCE_HASH: ${{ steps.artifact-names.outputs.source_hash }} run: | set -uo pipefail run_export() { @@ -230,7 +250,7 @@ jobs: json.dumps(manifest, indent=2) + "\n", encoding="utf-8" ) PY - echo "name=asr-${MODEL_SLUG}-onnx-${selected_precision}-${SOURCE_HASH}" >> "$GITHUB_OUTPUT" + echo "name=asr-${MODEL_SLUG}-onnx-${selected_precision}" >> "$GITHUB_OUTPUT" - name: Upload Whisper ONNX if: steps.selection.outputs.export == 'true' diff --git a/.github/workflows/model-changes.yml b/.github/workflows/model-changes.yml deleted file mode 100644 index d945b3e..0000000 --- a/.github/workflows/model-changes.yml +++ /dev/null @@ -1,32 +0,0 @@ -name: Detect model changes - -on: - workflow_call: - outputs: - whisper: - value: ${{ jobs.changes.outputs.whisper }} - parakeet: - value: ${{ jobs.changes.outputs.parakeet }} - nemotron: - value: ${{ jobs.changes.outputs.nemotron }} - sam2: - value: ${{ jobs.changes.outputs.sam2 }} - -permissions: - contents: read - -jobs: - changes: - runs-on: ubuntu-24.04 - outputs: - whisper: ${{ steps.changes.outputs.whisper }} - parakeet: ${{ steps.changes.outputs.parakeet }} - nemotron: ${{ steps.changes.outputs.nemotron }} - sam2: ${{ steps.changes.outputs.sam2 }} - steps: - - uses: actions/checkout@v7 - with: - fetch-depth: 0 - ref: ${{ github.event.pull_request.head.sha || github.sha }} - - id: changes - run: python .github/scripts/model_ci.py changes diff --git a/.github/workflows/nvidia-asr.yml b/.github/workflows/nvidia-asr.yml index 501f3d0..780708c 100644 --- a/.github/workflows/nvidia-asr.yml +++ b/.github/workflows/nvidia-asr.yml @@ -4,7 +4,6 @@ on: push: branches: [main] pull_request: - types: [opened, synchronize, reopened, ready_for_review, converted_to_draft] workflow_dispatch: inputs: build_run_id: @@ -21,13 +20,7 @@ concurrency: cancel-in-progress: true jobs: - changes: - if: github.event_name != 'pull_request' || !github.event.pull_request.draft - uses: ./.github/workflows/model-changes.yml - requirements: - needs: changes - if: needs.changes.outputs.parakeet == 'true' || needs.changes.outputs.nemotron == 'true' name: Wait for Build runs-on: ubuntu-24.04 timeout-minutes: 100 @@ -50,7 +43,7 @@ jobs: while [[ -z "$run_id" ]]; do run_id="$( gh api "repos/${GITHUB_REPOSITORY}/actions/workflows/ci.yml/runs?per_page=100" \ - --jq ".workflow_runs | map(select(.head_sha == \"$SOURCE_SHA\" and .event == \"$EVENT_NAME\" and .conclusion != \"skipped\")) | sort_by(.created_at) | last | .id // empty" + --jq ".workflow_runs | map(select(.head_sha == \"$SOURCE_SHA\" and .event == \"$EVENT_NAME\")) | sort_by(.created_at) | last | .id // empty" )" [[ -n "$run_id" ]] || sleep 10 done @@ -69,9 +62,8 @@ jobs: echo "source_ref=$source_ref" >> "$GITHUB_OUTPUT" export-parakeet: - if: needs.changes.outputs.parakeet == 'true' name: Export Parakeet TDT ONNX - needs: [requirements, changes] + needs: requirements uses: ./.github/workflows/export-nvidia-asr-model.yml with: model_slug: parakeet-tdt @@ -83,9 +75,8 @@ jobs: source_ref: ${{ needs.requirements.outputs.source_ref }} export-nemotron: - if: needs.changes.outputs.nemotron == 'true' name: Export Nemotron 3.5 ASR Streaming ONNX - needs: [requirements, changes] + needs: requirements uses: ./.github/workflows/export-nvidia-asr-model.yml with: model_slug: nemotron-3.5-asr-streaming-0.6b @@ -98,10 +89,7 @@ jobs: artifacts: name: Collect NVIDIA ASR artifacts - needs: [requirements, export-parakeet, export-nemotron] - if: >- - !cancelled() && !failure() && - (needs.export-parakeet.result == 'success' || needs.export-nemotron.result == 'success') + needs: [export-parakeet, export-nemotron] runs-on: ubuntu-24.04 outputs: matrix: ${{ steps.matrix.outputs.matrix }} @@ -139,7 +127,6 @@ jobs: }, ] } - matrix["include"] = [model for model in matrix["include"] if model["artifact_name"]] with Path(os.environ["GITHUB_OUTPUT"]).open("a", encoding="utf-8") as output: output.write(f"matrix={json.dumps(matrix, separators=(',', ':'))}\n") PY diff --git a/.github/workflows/sam2.yml b/.github/workflows/sam2.yml index 3149e44..cb12fa8 100644 --- a/.github/workflows/sam2.yml +++ b/.github/workflows/sam2.yml @@ -4,7 +4,6 @@ on: push: branches: [main] pull_request: - types: [opened, synchronize, reopened, ready_for_review, converted_to_draft] workflow_dispatch: inputs: build_run_id: @@ -21,13 +20,7 @@ concurrency: cancel-in-progress: true jobs: - changes: - if: github.event_name != 'pull_request' || !github.event.pull_request.draft - uses: ./.github/workflows/model-changes.yml - requirements: - needs: changes - if: needs.changes.outputs.sam2 == 'true' name: Wait for Build runs-on: ubuntu-24.04 timeout-minutes: 100 @@ -50,7 +43,7 @@ jobs: while [[ -z "$run_id" ]]; do run_id="$( gh api "repos/${GITHUB_REPOSITORY}/actions/workflows/ci.yml/runs?per_page=100" \ - --jq ".workflow_runs | map(select(.head_sha == \"$SOURCE_SHA\" and .event == \"$EVENT_NAME\" and .conclusion != \"skipped\")) | sort_by(.created_at) | last | .id // empty" + --jq ".workflow_runs | map(select(.head_sha == \"$SOURCE_SHA\" and .event == \"$EVENT_NAME\")) | sort_by(.created_at) | last | .id // empty" )" [[ -n "$run_id" ]] || sleep 10 done @@ -112,29 +105,25 @@ jobs: shell: bash env: GH_TOKEN: ${{ github.token }} - DEFAULT_BRANCH: ${{ github.event.repository.default_branch }} - EVENT_NAME: ${{ github.event_name }} - SOURCE_BRANCH: ${{ github.head_ref || github.ref_name }} - PULL_REQUEST: ${{ github.event.pull_request.number }} run: | set -euo pipefail + source_hash="$( + git ls-files -z -- 'vision/sam2/python/*.py' 'vision/sam2/python/**/*.py' \ + | xargs -0 sha256sum \ + | sha256sum \ + | cut -d' ' -f1 + )" models='{}' for slug in \ sam2.1-hiera-tiny \ sam2.1-hiera-small \ sam2.1-hiera-base-plus \ sam2.1-hiera-large; do - source_hash="$(python .github/scripts/model_ci.py key --model sam2 --model-id "facebook/${slug}")" artifact_name="vision-${slug}-onnx-v1-${source_hash:0:16}" - artifact_run_id="$(gh api "repos/${GITHUB_REPOSITORY}/actions/runs/${GITHUB_RUN_ID}/artifacts?per_page=100" \ - --jq ".artifacts[] | select(.name == \"$artifact_name\" and (.expired | not)) | \"$GITHUB_RUN_ID\"")" - if [[ -z "$artifact_run_id" ]]; then - artifact="$(python .github/scripts/find_successful_artifact.py \ - --repository "$GITHUB_REPOSITORY" --artifact-name "$artifact_name" \ - --default-branch "$DEFAULT_BRANCH" --event-name "$EVENT_NAME" \ - --source-branch "$SOURCE_BRANCH" --pull-request "$PULL_REQUEST" --json)" - artifact_run_id="$(jq -r '.run_id' <<< "$artifact")" - fi + artifact_run_id="$( + gh api "repos/${GITHUB_REPOSITORY}/actions/artifacts?name=${artifact_name}&per_page=100" \ + --jq '.artifacts | map(select(.expired | not)) | sort_by(.created_at) | last | .workflow_run.id // empty' + )" test -n "$artifact_run_id" models="$(jq \ --arg slug "$slug" \ diff --git a/.github/workflows/whisper-asr.yml b/.github/workflows/whisper-asr.yml index 3fdd8de..6d7d64f 100644 --- a/.github/workflows/whisper-asr.yml +++ b/.github/workflows/whisper-asr.yml @@ -4,7 +4,6 @@ on: push: branches: [main] pull_request: - types: [opened, synchronize, reopened, ready_for_review, converted_to_draft] workflow_dispatch: inputs: build_run_id: @@ -21,13 +20,7 @@ concurrency: cancel-in-progress: true jobs: - changes: - if: github.event_name != 'pull_request' || !github.event.pull_request.draft - uses: ./.github/workflows/model-changes.yml - requirements: - needs: changes - if: needs.changes.outputs.whisper == 'true' name: Wait for Build runs-on: ubuntu-24.04 timeout-minutes: 100 @@ -50,7 +43,7 @@ jobs: while [[ -z "$run_id" ]]; do run_id="$( gh api "repos/${GITHUB_REPOSITORY}/actions/workflows/ci.yml/runs?per_page=100" \ - --jq ".workflow_runs | map(select(.head_sha == \"$SOURCE_SHA\" and .event == \"$EVENT_NAME\" and .conclusion != \"skipped\")) | sort_by(.created_at) | last | .id // empty" + --jq ".workflow_runs | map(select(.head_sha == \"$SOURCE_SHA\" and .event == \"$EVENT_NAME\")) | sort_by(.created_at) | last | .id // empty" )" [[ -n "$run_id" ]] || sleep 10 done From 1754e9d09a6d109a0f76d465e60984b2a21fa5b5 Mon Sep 17 00:00:00 2001 From: contentis Date: Wed, 23 Sep 2026 12:00:26 +0200 Subject: [PATCH 6/8] Simplify Qwen3 ASR and remove bundled Japanese segmentation --- THIRD_PARTY_NOTICES.md | 213 +++++++++++++++++- asr/qwen3/CMakeLists.txt | 1 - asr/qwen3/README.md | 181 +++++---------- asr/qwen3/detail/japanese.h | 192 ---------------- asr/qwen3/detail/text.h | 22 -- asr/qwen3/forced_aligner.cpp | 19 +- asr/qwen3/model_export/detail/japanese.py | 149 ------------ .../model_export/detail/nagisa.LICENSE.txt | 21 -- asr/qwen3/model_export/export_qwen3_asr.py | 10 +- asr/qwen3/qwen3.cpp | 6 +- asr/qwen3/{detail => }/runtime.h | 0 common/io/unicode_regex.cpp | 15 ++ common/io/unicode_regex.h | 2 + 13 files changed, 288 insertions(+), 543 deletions(-) delete mode 100644 asr/qwen3/detail/japanese.h delete mode 100644 asr/qwen3/detail/text.h delete mode 100644 asr/qwen3/model_export/detail/japanese.py delete mode 100644 asr/qwen3/model_export/detail/nagisa.LICENSE.txt rename asr/qwen3/{detail => }/runtime.h (100%) diff --git a/THIRD_PARTY_NOTICES.md b/THIRD_PARTY_NOTICES.md index 0a48691..a721e4e 100644 --- a/THIRD_PARTY_NOTICES.md +++ b/THIRD_PARTY_NOTICES.md @@ -6,11 +6,10 @@ DIN Deploy is distributed under the Apache License, Version 2.0. The project sou |---|---| | argparse | MIT: | | miniaudio | MIT: | -| Nagisa word segmentation | MIT: | | lodepng | zlib: | | nlohmann/json | MIT: | -| PCRE2 | BSD-3-Clause WITH PCRE2-exception: | -| utf8proc | MIT and Unicode data license: | +| PCRE2 | BSD-3-Clause WITH PCRE2-exception ([notice](#pcre2-notice)): | +| utf8proc | MIT and Unicode data license ([notices](#utf8proc-notices)): | | NVIDIA NVTX | Apache-2.0 with LLVM exception: | | ONNX Runtime | MIT: | | ONNX Runtime TensorRT RTX EP ABI | Apache-2.0: | @@ -81,8 +80,6 @@ For components supplied through an SDK or binary package, the corresponding vend ## Model/Artifact -- Nagisa `nagisa_v001` word segmenter: MIT: - - `black-forest-labs/FLUX.2-klein-4b`: Apache-2.0 - `black-forest-labs/FLUX.2-klein-4b-fp8`: Apache-2.0 - `black-forest-labs/FLUX.2-klein-4b-nvfp4`: Apache-2.0 @@ -103,3 +100,209 @@ For components supplied through an SDK or binary package, the corresponding vend - `Qwen/Qwen3-ASR-0.6B-hf`: Apache-2.0 - `Qwen/Qwen3-ASR-1.7B-hf`: Apache-2.0 - `Qwen/Qwen3-ForcedAligner-0.6B-hf`: Apache-2.0 + +## PCRE2 notice + +```text +PCRE2 License +============= + +| SPDX-License-Identifier: | BSD-3-Clause WITH PCRE2-exception | +|---------|-------| + +PCRE2 is a library of functions to support regular expressions whose syntax +and semantics are as close as possible to those of the Perl 5 language. + +Releases 10.00 and above of PCRE2 are distributed under the terms of the "BSD" +licence, as specified below, with one exemption for certain binary +redistributions. The documentation for PCRE2, supplied in the "doc" directory, +is distributed under the same terms as the software itself. The data in the +testdata directory is not copyrighted and is in the public domain. + +The basic library functions are written in C and are freestanding. Also +included in the distribution is a just-in-time compiler that can be used to +optimize pattern matching. This is an optional feature that can be omitted when +the library is built. + + +COPYRIGHT +--------- + +### The basic library functions + + Written by: Philip Hazel + Email local part: Philip.Hazel + Email domain: gmail.com + + Retired from University of Cambridge Computing Service, + Cambridge, England. + + Copyright (c) 1997-2007 University of Cambridge + Copyright (c) 2007-2024 Philip Hazel + All rights reserved. + +### PCRE2 Just-In-Time compilation support + + Written by: Zoltan Herczeg + Email local part: hzmester + Email domain: freemail.hu + + Copyright (c) 2010-2024 Zoltan Herczeg + All rights reserved. + +### Stack-less Just-In-Time compiler + + Written by: Zoltan Herczeg + Email local part: hzmester + Email domain: freemail.hu + + Copyright (c) 2009-2024 Zoltan Herczeg + All rights reserved. + +### All other contributions + +Many other contributors have participated in the authorship of PCRE2. As PCRE2 +has never required a Contributor Licensing Agreement, or other copyright +assignment agreement, all contributions have copyright retained by each +original contributor or their employer. + + +THE "BSD" LICENCE +----------------- + +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that the following conditions are met: + +* Redistributions of source code must retain the above copyright notices, + this list of conditions and the following disclaimer. + +* Redistributions in binary form must reproduce the above copyright + notices, this list of conditions and the following disclaimer in the + documentation and/or other materials provided with the distribution. + +* Neither the name of the University of Cambridge nor the names of any + contributors may be used to endorse or promote products derived from this + software without specific prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT OWNER OR CONTRIBUTORS BE +LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR +CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF +SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN +CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) +ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE +POSSIBILITY OF SUCH DAMAGE. + + +EXEMPTION FOR BINARY LIBRARY-LIKE PACKAGES +------------------------------------------ + +The second condition in the BSD licence (covering binary redistributions) does +not apply all the way down a chain of software. If binary package A includes +PCRE2, it must respect the condition, but if package B is software that +includes package A, the condition is not imposed on package B unless it uses +PCRE2 independently. + +End +``` + +## utf8proc notices + +```text +## utf8proc license ## + +**utf8proc** is a software package originally developed +by Jan Behrens and the rest of the Public Software Group, who +deserve nearly all of the credit for this library, that is now maintained by the Julia-language developers. Like the original utf8proc, +whose copyright and license statements are reproduced below, all new +work on the utf8proc library is licensed under the [MIT "expat" +license](http://opensource.org/licenses/MIT): + +*Copyright © 2014-2021 by Steven G. Johnson, Jiahao Chen, Tony Kelman, Jonas Fonseca, and other contributors listed in the git history.* + +Permission is hereby granted, free of charge, to any person obtaining a +copy of this software and associated documentation files (the "Software"), +to deal in the Software without restriction, including without limitation +the rights to use, copy, modify, merge, publish, distribute, sublicense, +and/or sell copies of the Software, and to permit persons to whom the +Software is furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in +all copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING +FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER +DEALINGS IN THE SOFTWARE. + +## Original utf8proc license ## + +*Copyright (c) 2009, 2013 Public Software Group e. V., Berlin, Germany* + +Permission is hereby granted, free of charge, to any person obtaining a +copy of this software and associated documentation files (the "Software"), +to deal in the Software without restriction, including without limitation +the rights to use, copy, modify, merge, publish, distribute, sublicense, +and/or sell copies of the Software, and to permit persons to whom the +Software is furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in +all copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING +FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER +DEALINGS IN THE SOFTWARE. + +## Unicode data license ## + +This software contains data (`utf8proc_data.c`) derived from processing +the Unicode data files. The following license applies to that data: + +**COPYRIGHT AND PERMISSION NOTICE** + +*Copyright (c) 1991-2007 Unicode, Inc. All rights reserved. Distributed +under the Terms of Use in http://www.unicode.org/copyright.html.* + +Permission is hereby granted, free of charge, to any person obtaining a +copy of the Unicode data files and any associated documentation (the "Data +Files") or Unicode software and any associated documentation (the +"Software") to deal in the Data Files or Software without restriction, +including without limitation the rights to use, copy, modify, merge, +publish, distribute, and/or sell copies of the Data Files or Software, and +to permit persons to whom the Data Files or Software are furnished to do +so, provided that (a) the above copyright notice(s) and this permission +notice appear with all copies of the Data Files or Software, (b) both the +above copyright notice(s) and this permission notice appear in associated +documentation, and (c) there is clear notice in each modified Data File or +in the Software as well as in the documentation associated with the Data +File(s) or Software that the data or software has been modified. + +THE DATA FILES AND SOFTWARE ARE PROVIDED "AS IS", WITHOUT WARRANTY OF ANY +KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF +MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT OF +THIRD PARTY RIGHTS. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR HOLDERS +INCLUDED IN THIS NOTICE BE LIABLE FOR ANY CLAIM, OR ANY SPECIAL INDIRECT OR +CONSEQUENTIAL DAMAGES, OR ANY DAMAGES WHATSOEVER RESULTING FROM LOSS OF +USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR OTHER +TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +PERFORMANCE OF THE DATA FILES OR SOFTWARE. + +Except as contained in this notice, the name of a copyright holder shall +not be used in advertising or otherwise to promote the sale, use or other +dealings in these Data Files or Software without prior written +authorization of the copyright holder. + +Unicode and the Unicode logo are trademarks of Unicode, Inc., and may be +registered in some jurisdictions. All other trademarks and registered +trademarks mentioned herein are the property of their respective owners. +``` diff --git a/asr/qwen3/CMakeLists.txt b/asr/qwen3/CMakeLists.txt index c084c7b..3426310 100644 --- a/asr/qwen3/CMakeLists.txt +++ b/asr/qwen3/CMakeLists.txt @@ -2,7 +2,6 @@ add_din_shared_library(din_asr_qwen3 STATIC qwen3.cpp forced_aligner.cpp) target_include_directories(din_asr_qwen3 PUBLIC "${CMAKE_CURRENT_SOURCE_DIR}") target_link_libraries(din_asr_qwen3 PUBLIC din_common_ort din_common_io nlohmann_json::nlohmann_json) -target_link_libraries(din_asr_qwen3 PRIVATE utf8proc) find_package(CUDAToolkit QUIET) if(CUDAToolkit_FOUND) target_link_libraries(din_asr_qwen3 PRIVATE CUDA::cudart) diff --git a/asr/qwen3/README.md b/asr/qwen3/README.md index 6229c14..033ab19 100644 --- a/asr/qwen3/README.md +++ b/asr/qwen3/README.md @@ -1,14 +1,17 @@ # Qwen3 ASR and forced alignment -C++ inference with CPU or TensorRT RTX. Use FP32 exports for CPU; TensorRT RTX supports BF16 (default), FP16 and FP32. +ONNX Runtime inference on CPU or TensorRT RTX, with independent ASR and alignment APIs. ## Supported models -| Model | Hugging Face ID | Checkpoint | BF16 export | FP16 export | FP32 export | Recommended export | -| --- | --- | --- | --- | --- | --- | --- | -| ASR 0.6B | `Qwen/Qwen3-ASR-0.6B-hf` | BF16 | ✓ | ✓ | ✓ | BF16 | -| ASR 1.7B | `Qwen/Qwen3-ASR-1.7B-hf` | BF16 | ✓ | ✓ | ✓ | BF16 | -| Forced Aligner 0.6B | `Qwen/Qwen3-ForcedAligner-0.6B-hf` | BF16 | ✓ | ✓ | ✓ | BF16 | +| Model | Hugging Face ID | +|---|---| +| ASR 0.6B | `Qwen/Qwen3-ASR-0.6B-hf` | +| ASR 1.7B | `Qwen/Qwen3-ASR-1.7B-hf` | +| Forced Aligner 0.6B | `Qwen/Qwen3-ForcedAligner-0.6B-hf` | + +All models support BF16 (original precision, default), FP16 and FP32 exports. +Use FP32 for CPU inference. ## Supported capabilities @@ -17,151 +20,71 @@ C++ inference with CPU or TensorRT RTX. Use FP32 exports for CPU; TensorRT RTX s | Offline, single stream | ✓ | | Online / streaming, single utterance | ✓ | | Batched inference | — | -| Long-form audio | ✓ | -| ASR with / without forced alignment | ✓ | -| Standalone alignment of supplied text | ✓ | +| Long-form ASR | ✓ | +| Automatic language detection / language hint | ✓ | +| ASR: 30 languages and 22 Chinese dialects | ✓ | +| Standalone alignment / alignment of any ASR output | ✓ | | Long-form alignment with timed transcript segments (C++ API) | ✓ | | Long-form alignment of unsegmented text | — | -| Automatic language identification / language hint | ✓ | -| Multilingual ASR: 30 languages and 22 Chinese dialects | ✓ | | Word timestamps: en, de, es, fr, it, pt, ru, ko | ✓ | | Chinese / Cantonese character timestamps | ✓ | -| Japanese word timestamps (Nagisa) | ✓ | -| Character alignment units: all 11 languages (sample extension) | ✓ | +| Automatic Japanese word timestamps | — | +| Character timestamps: all 11 alignment languages (sample extension) | ✓ | | Caller-supplied alignment units | ✓ | -| All 11 upstream alignment languages | ✓ | -ASR uses the [upstream model's language support](https://github.com/QwenLM/Qwen3-ASR). -ASR accepts all 30 upstream language codes/names and `auto`; dialects use automatic -recognition or the corresponding language hint, not separate dialect switches. -Alignment supports Chinese, Cantonese, English, German, Spanish, French, Italian, -Portuguese, Russian, Korean and Japanese. Japanese word boundaries use upstream -Nagisa; Chinese/Cantonese default to characters, keeping Latin words together. -Language names/codes are case-insensitive. ASR's other languages require alignment -to be disabled. Language coverage is not an accuracy guarantee for every dialect. +See [upstream language support](https://github.com/QwenLM/Qwen3-ASR). +Japanese alignment requires character mode or supplied units. -## Export +## Export and validate -Run Python commands from `asr/qwen3/model_export`, with the dependencies in -`../requirements.txt` and a CUDA-enabled PyTorch installation. +Run commands from the repository root; install the export dependencies first. ```bash -python -X utf8 export_qwen3_asr.py --size 0.6B --output D:/models/qwen3-asr-0.6b-onnx-bf16 -python -X utf8 export_qwen3_asr.py --size 1.7B --output D:/models/qwen3-asr-1.7b-onnx-bf16 -python -X utf8 export_qwen3_asr.py --task aligner --output D:/models/qwen3-aligner-onnx-bf16 +pip install -r asr/qwen3/requirements.txt +python asr/qwen3/model_export/export_qwen3_asr.py --output models/qwen3-asr +python asr/qwen3/model_export/export_qwen3_asr.py --task aligner --output models/qwen3-aligner ``` -Use `--dtype fp16` or `--dtype fp32` for converted exports; `--dtype original` -keeps BF16. This applies to both ASR sizes and the aligner. +HF downloads checkpoints automatically. Use `--size 1.7B` for the larger ASR model, +or `--dtype fp16` / `--dtype fp32` to change precision; use a separate output directory. +Each ASR export contains one encoder, one decoder and the shared log-mel graph. ```bash -python -X utf8 export_qwen3_asr.py --dtype fp16 --output D:/models/qwen3-asr-0.6b-onnx-fp16 -python -X utf8 export_qwen3_asr.py --dtype fp32 --output D:/models/qwen3-asr-0.6b-onnx-fp32 +python asr/qwen3/model_export/validate_qwen3_asr.py --onnx-dir models/qwen3-asr --audio audio.wav +python asr/qwen3/model_export/validate_qwen3_asr.py --task aligner --onnx-dir models/qwen3-aligner --audio audio.wav --transcript transcript.txt --language English ``` -The C++ pipeline reads precision from each export; ASR and aligner can use different -precisions. FP32 uses decomposed attention for TensorRT RTX compatibility; BF16/FP16 -use fused attention. Log-mel stays FP32. Use separate output directories per precision. - -HF downloads checkpoints automatically. Use `--model` for a local checkpoint or -`--revision` to pin the source. Keep each export directory intact. Log-mel -processing reuses the shared Whisper frontend. - -Aligner exports also include Nagisa's small FP32 word segmenter and vocabulary. -It runs on CPU without Python. Add it to an existing aligner export with -`--task aligner --only japanese --output `; no ASR/aligner -weights or GPU engines need rebuilding. - -ASR exports contain one encoder and one decoder, each with one weight file, plus the -shared log-mel graph. Prefill and token generation update one KV bank in place. -The decoder uses two fixed TensorRT profiles (512-token prefill and one-token -steps), compiled once and cached; audio length does not create more encoder/decoder -profiles. The two engines may each retain weights in GPU memory. Changed exports -get new cache keys; clear compiled caches when changing the GPU or runtime. - -`--cache-capacity` sets the token ceiling (default 8192; multiples of 512 up to -16384). Long audio uses upstream quiet-boundary splitting with enough room reserved -for `--max-new-tokens`. Reaching that generation limit returns `reached_eos=false`. -Attention still scans the allocated cache. BF16/FP16 KV uses 896 MiB at 8192 slots. -Re-export older ASR artifacts for format 3. +## Build and run -## Verify +Follow the [repository build setup](../../README.md). Replace `` and +`` with your build and executable directories; append `.exe` on Windows. -```bash -python -X utf8 validate_qwen3_asr.py --onnx-dir D:/models/qwen3-asr-0.6b-onnx-bf16 --audio audio.mp3 -python -X utf8 validate_qwen3_asr.py --onnx-dir D:/models/qwen3-asr-1.7b-onnx-bf16 --audio audio.mp3 -python -X utf8 validate_qwen3_asr.py --task aligner --onnx-dir D:/models/qwen3-aligner-onnx-bf16 --audio audio.mp3 --transcript transcript.txt --language Chinese +```text +cmake --build --config Release --target din_asr_qwen3_cli din_asr_qwen3_aligner_cli +/din_asr_qwen3_cli audio.wav --model-dir models/qwen3-asr +/din_asr_qwen3_aligner_cli audio.wav --model-dir models/qwen3-aligner --transcript transcript.txt --lang-id en ``` -The validator compares encoder, prefill and cached-token outputs against HF, -then uses HF `generate()` for a short end-to-end token/EOS check. Alignment uses -HF transcript preparation and span decoding. Strict BF16 numerical comparisons -can fail despite matching tokens/spans. Long-form specialized decoding can change -words; long-form alignment has small endpoint differences from HF. - -## Build +Both CLIs accept `--provider cpu|trt-rtx` (default: trt-rtx). +Use `--granularity characters` for character alignment or `--units` for one supplied +alignment unit per transcript line. Character timestamps have 80 ms resolution. +Japanese word segmentation can be an optional external preprocessing pass supplied through `AlignUnits` / `--units`. -Build from the repository root with the same TensorRT RTX setup as Whisper. -An NVIDIA GPU supporting the selected TensorRT RTX precision is required. +## Long-form and streaming -```powershell -cmake --build out\build\windows-x64 --target din_asr_qwen3_cli din_asr_qwen3_aligner_cli -``` +Long recordings are split automatically. `--max-new-tokens` defaults to 1024 per +chunk; `reached_eos=false` means the transcript is incomplete. Increase the budget +or reduce `--max-chunk-seconds`. Export-time `--cache-capacity` defaults to 8192 tokens. -## Run +Add `--stream` for streaming transcription; `- --stream` reads mono 16 kHz float32 +PCM from stdin. Outputs are replacement hypotheses, not incremental text. +Streaming reprocesses the current utterance, so latency grows with its length; +start a new stream before exceeding the exported context capacity. -```powershell -out\build\windows-x64\bin\din_asr_qwen3_cli.exe audio.mp3 --model-dir D:\models\qwen3-asr-1.7b-onnx-bf16 -out\build\windows-x64\bin\din_asr_qwen3_aligner_cli.exe audio.mp3 --model-dir D:\models\qwen3-aligner-onnx-bf16 --transcript transcript.txt --lang-id zh -out\build\windows-x64\bin\din_asr_qwen3_aligner_cli.exe audio.mp3 --model-dir D:\models\qwen3-aligner-onnx-bf16 --transcript transcript.txt --lang-id ja --granularity characters -out\build\windows-x64\bin\din_asr_qwen3_cli.exe audio.mp3 --model-dir D:\models\qwen3-asr-0.6b-onnx-bf16 --stream -``` +## C++ integration -Multi-configuration builds add the configuration (for example, `Release`) under `bin`. -ASR and alignment are independent APIs in the same library: - -| API / CLI | Input | Output | -|---|---|---| -| `Qwen3Pipeline` / `din_asr_qwen3_cli` | Audio | Text, tokens, language and audio chunk boundaries | -| `Qwen3ForcedAligner` / `din_asr_qwen3_aligner_cli` | Audio, supplied text and language | Word/character timestamps | - -Include `qwen3.h` for ASR or `forced_aligner.h` for alignment. The aligner loads -no ASR model; use text from Whisper, Parakeet, Nemotron, Qwen ASR or a text file. -Both CLIs accept `--provider cpu|trt-rtx`, `--model-dir` and cache options independently. -Reuse instances; each processes one synchronous call at a time. - -`Align` takes mono 16 kHz audio; `AlignFile` decodes it automatically. Alignment -accepts at most 180 seconds per call. For long recordings, use -`AlignSegments(audio, segments)` with text from any ASR. Each `AlignmentSegment` -contains `text`, half-open `start_sample` / `end_sample` offsets at 16 kHz, and -`language` (default English). The method reuses the aligner and returns timestamps -relative to the full recording, in segment order. Bounds and the 180-second limit -are checked before inference; overlapping intervals are preserved without -deduplication. Each interval must contain all speech for its text. - -`--granularity characters` / `AlignmentGranularity::Characters` aligns Unicode -graphemes, retaining combining marks and omitting spaces/punctuation except apostrophes. -`AlignUnits(audio, units)` bypasses text splitting; CLI `--units` reads one unit -per transcript line. Limits: 2048 units and 8192 context tokens per call. Character -alignment is a sample extension: the model's 80 ms timestamp bins can give adjacent -characters identical times; sub-word accuracy is not guaranteed. - -Streaming uses `StartStream()`, `PushAudio(mono16k)` and `FinishStream()`. -The CLI emits replacement hypotheses as JSON lines; use `- --stream` to read -little-endian float32 mono 16 kHz PCM from stdin. `--chunk-seconds` defaults to 2 -(minimum 0.5); `--unfixed-chunks 2 --unfixed-tokens 5` matches upstream rollback. -Each update reprocesses accumulated audio using the existing engines, so latency -grows with utterance length. The exported KV ceiling still applies and overflow -raises an error; start a new stream for the next utterance. Streaming has no live -timestamps: align the final text separately. `FinishStream()` flushes the tail; -UTF-8-safe rollback also applies to that tail. - -ASR returns `segments` with text, detected language and half-open `start_sample` / -`end_sample` offsets at 16 kHz. Its upstream quiet-boundary splitter uses a -1200-second target and a ±5-second search, further limited by KV capacity. -`--max-chunk-seconds` overrides the target (6–1200; 0 uses the default). -For subsequent alignment, use a target of 175 seconds or less to reserve the -search margin within the 180-second limit. Chunk boundaries are not word timestamps. - -`--max-new-tokens` defaults to 1024 per chunk. If `reached_eos` is false, increase -the budget or reduce `--max-chunk-seconds`; the transcript is incomplete. +Use `Qwen3Pipeline` ([qwen3.h](qwen3.h)) for transcription and +`Qwen3ForcedAligner` ([forced_aligner.h](forced_aligner.h)) for text from any ASR. +`AlignSegments` accepts timed transcript segments and returns recording-relative +timestamps. Alignment is limited to 180 seconds, 2048 units and 8192 context tokens +per segment; use `--max-chunk-seconds 175` when transcribing for subsequent alignment. diff --git a/asr/qwen3/detail/japanese.h b/asr/qwen3/detail/japanese.h deleted file mode 100644 index 5345d73..0000000 --- a/asr/qwen3/detail/japanese.h +++ /dev/null @@ -1,192 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 -#pragma once -#include -#include -#include - -#include "runtime.h" -#include "text.h" -#include "unicode_regex.h" - -namespace din::asr::qwen3::detail -{ -// Nagisa preprocessing, dictionary features and BMES decoding; neural inference stays in ORT. -class JapaneseTokenizer -{ - using Vocabulary = std::unordered_map; - Vocabulary unigrams_, bigrams_, words_; - int window_; - int64_t padding_word_; - std::array, 6> transitions_; - Ort::Session session_{nullptr}; - - static std::string Utf8(int32_t code) - { - utf8proc_uint8_t data[4]; - const auto size = utf8proc_encode_char(code, data); - return {reinterpret_cast(data), static_cast(size)}; - } - static int64_t Lookup(const Vocabulary& vocabulary, const std::string& text) - { - const auto found = vocabulary.find(text); - return found == vocabulary.end() ? vocabulary.at("oov") : found->second; - } - -public: - JapaneseTokenizer(Ort::Env& env, const std::filesystem::path& dir) - { - if (!std::filesystem::exists(dir / "japanese.json")) - throw std::runtime_error("Japanese word alignment requires --task aligner --only japanese export"); - const auto meta = ReadJson(dir / "japanese.json"); - unigrams_ = meta.at("unigrams").get(); - bigrams_ = meta.at("bigrams").get(); - words_ = meta.at("words").get(); - window_ = meta.at("window"); - padding_word_ = meta.at("padding_word"); - transitions_ = meta.at("transitions").get(); - Ort::SessionOptions options; - options.SetIntraOpNumThreads(1); - session_ = Ort::Session(env, (dir / "japanese.onnx").c_str(), options); - } - - std::vector Words(std::string text) - { - if (!ValidUtf8(text)) - throw std::invalid_argument("Japanese alignment requires valid UTF-8 text"); - static const din::io::UnicodeRegex leading(R"(^[\s\x{1c}-\x{1f}]+)"), trailing(R"([\s\x{1c}-\x{1f}]+$)"); - const auto head = leading.FindAll(text); - if (!head.empty()) - text.erase(0, head[0].size()); - const auto tail = trailing.FindAll(text); - if (!tail.empty()) - text.resize(text.size() - tail[0].size()); - utf8proc_uint8_t* normalized = nullptr; - const auto length = - utf8proc_map(reinterpret_cast(text.data()), text.size(), &normalized, - static_cast(UTF8PROC_STABLE | UTF8PROC_COMPAT | UTF8PROC_COMPOSE)); - const std::unique_ptr owner(normalized, std::free); - if (length < 0) - throw std::invalid_argument(utf8proc_errmsg(length)); - std::vector characters, lower; - std::vector types; - for (size_t i = 0; i < static_cast(length);) - { - int32_t code; - i += utf8proc_iterate(normalized + i, length - i, &code); - if (code == 0x130) - code = 'I'; - if (code == ' ') - code = 0x3000; - characters.push_back(Utf8(code)); - code = utf8proc_tolower(code); - lower.push_back(Utf8(code)); - types.push_back(code >= 0x3040 && code <= 0x309f ? 0 - : code >= 0x30a1 && code <= 0x30fa ? 1 - : code >= 0x4e00 && code <= 0x9fa5 ? 2 - : code >= 'a' && code <= 'z' ? 3 - : code >= '0' && code <= '9' ? 4 - : 5); - } - const int64_t count = characters.size(); - if (!count) - return {}; - std::array, 5> features; - for (size_t i = 0; i < 3; ++i) - features[i].resize(count * window_, i == 2 ? 6 : 1); - for (size_t i = 3; i < 5; ++i) - features[i].resize(count * 8, padding_word_); - for (int64_t i = 0; i < count; ++i) - { - for (int j = 0; j < window_; ++j) - { - const auto at = i + j - window_ / 2; - if (at < 0 || at >= count) - continue; - features[0][i * window_ + j] = Lookup(unigrams_, lower[at]); - features[1][i * window_ + j] = Lookup(bigrams_, lower[at] + (at + 1 < count ? lower[at + 1] : "")); - features[2][i * window_ + j] = types[at]; - } - for (int direction = 0; direction < 2; ++direction) - { - std::string word; - int matches = 0; - for (int j = 0; j < 8; ++j) - { - const auto at = direction ? i - j : i + j; - if (at < 0 || at >= count) - break; - word = direction ? lower[at] + word : word + lower[at]; - if (const auto found = words_.find(word); found != words_.end()) - features[3 + direction][i * 8 + matches++] = found->second; - } - if (!matches) - features[3 + direction][i * 8] = words_.at("oov"); - } - } - const auto memory = Ort::MemoryInfo::CreateCpu(OrtArenaAllocator, OrtMemTypeDefault); - std::vector inputs; - for (size_t i = 0; i < features.size(); ++i) - { - const int64_t shape[]{count, i < 3 ? window_ : 8}; - inputs.push_back(Ort::Value::CreateTensor(memory, features[i].data(), features[i].size(), shape, 2)); - } - const char* names[]{"unigrams", "bigrams", "types", "word_starts", "word_ends"}; - const char* output_name = "emissions"; - auto outputs = session_.Run(Ort::RunOptions{}, names, inputs.data(), inputs.size(), &output_name, 1); - const auto* emissions = outputs[0].GetTensorData(); - std::vector> parents(count); - std::array scores; - scores.fill(-1e10f); - scores[4] = 0; - for (int64_t i = 0; i < count; ++i) - { - std::array next; - for (int to = 0; to < 6; ++to) - { - int best = 0; - for (int from = 1; from < 6; ++from) - if (scores[from] + transitions_[to][from] > scores[best] + transitions_[to][best]) - best = from; - parents[i][to] = best; - next[to] = scores[best] + transitions_[to][best] + emissions[i * 6 + to]; - } - // A common offset preserves the best path while keeping long sequences in FP32 range. - const float maximum = *std::max_element(next.begin(), next.end()); - for (int j = 0; j < 6; ++j) - scores[j] = next[j] - maximum; - } - int tag = 0; - for (int i = 1; i < 6; ++i) - if (scores[i] + transitions_[5][i] > scores[tag] + transitions_[5][tag]) - tag = i; - std::vector tags(count); - for (int64_t i = count; i-- > 0;) - { - tags[i] = tag; - tag = parents[i][tag]; - } - static const din::io::UnicodeRegex kept(R"([\p{L}\p{N}']+)"); - std::vector result; - std::string word; - auto emit = [&] - { - std::string cleaned; - for (const auto& part : kept.FindAll(word)) - cleaned += part; - if (!cleaned.empty()) - result.push_back(std::move(cleaned)); - word.clear(); - }; - for (int64_t i = 0; i < count; ++i) - { - if (tags[i] == 3) - emit(); - word += characters[i]; - if (tags[i] == 2 || tags[i] == 3) - emit(); - } - emit(); - return result; - } -}; -} // namespace din::asr::qwen3::detail diff --git a/asr/qwen3/detail/text.h b/asr/qwen3/detail/text.h deleted file mode 100644 index 63c0e51..0000000 --- a/asr/qwen3/detail/text.h +++ /dev/null @@ -1,22 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 -#pragma once -#include - -#include - -namespace din::asr::qwen3::detail -{ -inline bool ValidUtf8(const std::string& text) -{ - for (size_t i = 0; i < text.size();) - { - utf8proc_int32_t code; - const auto n = - utf8proc_iterate(reinterpret_cast(text.data() + i), text.size() - i, &code); - if (n < 0) - return false; - i += n; - } - return true; -} -} // namespace din::asr::qwen3::detail diff --git a/asr/qwen3/forced_aligner.cpp b/asr/qwen3/forced_aligner.cpp index 958287c..6b444ef 100644 --- a/asr/qwen3/forced_aligner.cpp +++ b/asr/qwen3/forced_aligner.cpp @@ -1,15 +1,13 @@ // SPDX-License-Identifier: Apache-2.0 #include "forced_aligner.h" -#include "detail/japanese.h" -#include "detail/runtime.h" -#include "detail/text.h" +#include "runtime.h" #include "unicode_regex.h" namespace din::asr::qwen3::detail { std::vector AlignmentUnits(const std::string& text, const std::string& language, - AlignmentGranularity granularity, JapaneseTokenizer* japanese_tokenizer) + AlignmentGranularity granularity) { static constexpr std::pair languages[] = { {"zh", "chinese"}, {"yue", "cantonese"}, {"en", "english"}, {"de", "german"}, @@ -24,7 +22,7 @@ std::vector AlignmentUnits(const std::string& text, const std::stri if (found == std::end(languages)) throw std::invalid_argument("Forced alignment supports zh, yue, en, de, es, fr, it, pt, ru, ko and ja"); - if (!ValidUtf8(text)) + if (!din::io::ValidUtf8(text)) throw std::invalid_argument("Alignment requires valid UTF-8 text"); if (granularity == AlignmentGranularity::Characters) { @@ -36,7 +34,8 @@ std::vector AlignmentUnits(const std::string& text, const std::stri return result; } if (found->first == "ja") - return japanese_tokenizer->Words(text); + throw std::invalid_argument( + "Japanese word alignment requires supplied units; use AlignUnits/--units or --granularity characters"); // HF keeps Unicode letters/numbers and ASCII apostrophes, dropping punctuation and marks. static const din::io::UnicodeRegex kept(R"([\p{L}\p{N}'\s\x{1c}-\x{1f}]+)"); @@ -199,7 +198,6 @@ struct Qwen3ForcedAligner::Impl ForcedAlignerConfig config; Runtime runtime; AlignmentEngine engine; - std::unique_ptr japanese; explicit Impl(ForcedAlignerConfig cfg) : config(std::move(cfg)) , runtime(config.provider, config.ep_cache_dir, config.ep_context_dir) @@ -209,15 +207,12 @@ struct Qwen3ForcedAligner::Impl } std::vector Units(const std::string& text, const std::string& language) { - const auto lang = LowerLanguage(language); - if ((lang == "ja" || lang == "japanese") && config.granularity == AlignmentGranularity::Words && !japanese) - japanese = std::make_unique(runtime.env, config.model_dir); - return AlignmentUnits(text, language, config.granularity, japanese.get()); + return AlignmentUnits(text, language, config.granularity); } std::vector AlignAudio(const din::io::Audio& audio, const std::vector& words) { for (const auto& word : words) - if (word.empty() || !ValidUtf8(word)) + if (word.empty() || !din::io::ValidUtf8(word)) throw std::invalid_argument("Alignment units must be nonempty UTF-8 text"); if (words.empty()) return {}; diff --git a/asr/qwen3/model_export/detail/japanese.py b/asr/qwen3/model_export/detail/japanese.py deleted file mode 100644 index c48afe8..0000000 --- a/asr/qwen3/model_export/detail/japanese.py +++ /dev/null @@ -1,149 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -"""Export Nagisa's word segmenter as an FP32 ONNX LSTM; no POS model is needed.""" - -import gzip -import hashlib -import json -import pickle -import re -from array import array -from pathlib import Path - -import onnx -import torch -from onnx import TensorProto, helper - -REVISION = "3c4bb48d3ba7451e3314b35337c79f1256ade0cf" -HASHES = { - "dict": "968ac9e6c7a53051ef24d8561673dd31de81b2feb9b5bff01b1d3b6b2473113c", - "hp": "6737f76b588315fe2fe05d05c99939d3142a0f1c6468f49f4c23e7003192f204", - "model": "9db9abc06a927c56e18af8d485e20a14908138752be459a83f0dc7ac85368c1b", -} - - -def export_japanese(output, source=None): - source = Path(source or Path(torch.hub.get_dir()) / "nagisa" / REVISION) - source.mkdir(parents=True, exist_ok=True) - for suffix, digest in HASHES.items(): - path = source / f"nagisa_v001.{suffix}" - if not path.exists(): - torch.hub.download_url_to_file( - f"https://raw.githubusercontent.com/taishi-i/nagisa/{REVISION}/nagisa/data/{path.name}", - str(path), - hash_prefix=digest, - ) - if hashlib.sha256(path.read_bytes()).hexdigest() != digest: - raise ValueError(f"Unexpected Nagisa asset: {path}") - # Only load the authenticated, pinned upstream dictionaries above. - vocabs = pickle.loads(gzip.decompress((source / "nagisa_v001.dict").read_bytes())) - hp = pickle.loads(gzip.decompress((source / "nagisa_v001.hp").read_bytes())) - parameters, lookups = [], [] - with (source / "nagisa_v001.model").open("rb") as file: - while header := file.readline(): - match = re.fullmatch(rb"#(Parameter|LookupParameter)# (\S+) \{([\d,]+)\} (\d+) ZERO_GRAD\s*", header) - if not match: - raise ValueError("Unsupported Nagisa parameter layout") - kind, name, dimensions, count = match.groups() - shape = [int(n) for n in dimensions.split(b",")] - data = torch.tensor(array("f", (float(n) for n in file.read(int(count)).split())), dtype=torch.float32) - if len(shape) == 2: - data = data.reshape(shape[::-1]) - if kind == b"Parameter": - data = data.T - (parameters if kind == b"Parameter" else lookups).append((name.decode(), data.contiguous())) - - nodes, constants = [], [] - - def const(name, data): - data = data.contiguous().clone() - tensor = TensorProto(name=name, data_type=TensorProto.INT64 if data.dtype == torch.int64 else TensorProto.FLOAT) - tensor.dims.extend(data.shape) - tensor.raw_data = bytes(data.untyped_storage()) - constants.append(tensor) - return name - - def node(op, inputs, name, **attributes): - nodes.append(helper.make_node(op, inputs, [name], **attributes)) - return name - - width = hp["WINDOW_SIZE"] - inputs, pieces = [], [] - word_table = torch.cat([lookups[2][1], torch.zeros(1, hp["DIM_WORD"])]) - const("word_table", word_table) - const("axis1", torch.tensor([1])) - for name, table in [("unigrams", lookups[0][1]), ("bigrams", lookups[1][1]), ("types", lookups[3][1])]: - inputs.append(helper.make_tensor_value_info(name, TensorProto.INT64, ["characters", width])) - gathered = node("Gather", [const(name + "_table", table), name], name + "_vectors", axis=0) - pieces.append(node("Flatten", [gathered], name + "_flat", axis=1)) - for name in ("word_starts", "word_ends"): - inputs.append(helper.make_tensor_value_info(name, TensorProto.INT64, ["characters", 8])) - gathered = node("Gather", ["word_table", name], name + "_vectors", axis=0) - pieces.append(node("ReduceSum", [gathered, "axis1"], name + "_sum", keepdims=0)) - x = node("Concat", pieces, "features", axis=1) - x = node("Unsqueeze", [x, "axis1"], "sequence") - hidden = hp["DIM_HIDDEN"] // 2 - # DyNet gates i,f,o,g -> ONNX gates i,o,f,g; DyNet adds +1 to the forget bias. - order = torch.cat([torch.arange(i * hidden, (i + 1) * hidden) for i in (0, 2, 1, 3)]) - for layer in range(hp["LAYERS"]): - weights, recurrent, biases = [], [], [] - for direction in range(2): - offset = (2 * layer + direction) * 3 - wx, wh, bias = [p[1] for p in parameters[offset : offset + 3]] - bias = bias.clone() - bias[hidden : 2 * hidden] += 1 - weights.append(wx[order]) - recurrent.append(wh[order]) - biases.append(torch.cat([bias[order], torch.zeros(4 * hidden)])) - name = f"lstm_{layer}" - x = node( - "LSTM", - [ - x, - const(name + "_w", torch.stack(weights)), - const(name + "_r", torch.stack(recurrent)), - const(name + "_b", torch.stack(biases)), - ], - name, - hidden_size=hidden, - direction="bidirectional", - ) - x = node("Transpose", [x], name + "_ordered", perm=[0, 2, 1, 3]) - x = node("Reshape", [x, const(name + "_shape", torch.tensor([-1, 1, 2 * hidden]))], name + "_flat") - plain = [p for p in parameters if p[0].count("/") == 1] - x = node("Squeeze", [x, "axis1"], "hidden") - x = node("MatMul", [x, const("projection", plain[0][1].T)], "projected") - node("Add", [x, const("bias", plain[1][1])], "emissions") - graph = helper.make_graph( - nodes, - "nagisa_word_segmentation", - inputs, - [helper.make_tensor_value_info("emissions", TensorProto.FLOAT, ["characters", 6])], - constants, - ) - model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 17)], ir_version=8) - output = Path(output) - output.mkdir(parents=True, exist_ok=True) - path = output / "japanese.onnx" - path.with_suffix(".onnx.data").unlink(missing_ok=True) - onnx.save_model( - model, - path, - save_as_external_data=True, - all_tensors_to_one_file=True, - location="japanese.onnx.data", - size_threshold=1024, - ) - onnx.checker.check_model(str(path)) - metadata = { - "revision": REVISION, - "window": width, - "unigrams": vocabs[0], - "bigrams": vocabs[1], - "words": vocabs[2], - "padding_word": len(word_table) - 1, - "transitions": lookups[5][1].tolist(), - } - (output / "japanese.json").write_text(json.dumps(metadata, ensure_ascii=False), encoding="utf-8") - (output / "japanese.LICENSE.txt").write_text( - Path(__file__).with_name("nagisa.LICENSE.txt").read_text(encoding="utf-8"), encoding="utf-8" - ) diff --git a/asr/qwen3/model_export/detail/nagisa.LICENSE.txt b/asr/qwen3/model_export/detail/nagisa.LICENSE.txt deleted file mode 100644 index 52360b0..0000000 --- a/asr/qwen3/model_export/detail/nagisa.LICENSE.txt +++ /dev/null @@ -1,21 +0,0 @@ -MIT License - -Copyright (c) 2018 taishi-i - -Permission is hereby granted, free of charge, to any person obtaining a copy -of this software and associated documentation files (the "Software"), to deal -in the Software without restriction, including without limitation the rights -to use, copy, modify, merge, publish, distribute, sublicense, and/or sell -copies of the Software, and to permit persons to whom the Software is -furnished to do so, subject to the following conditions: - -The above copyright notice and this permission notice shall be included in all -copies or substantial portions of the Software. - -THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR -IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, -FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE -AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER -LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, -OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE -SOFTWARE. diff --git a/asr/qwen3/model_export/export_qwen3_asr.py b/asr/qwen3/model_export/export_qwen3_asr.py index dee041d..a46fefc 100644 --- a/asr/qwen3/model_export/export_qwen3_asr.py +++ b/asr/qwen3/model_export/export_qwen3_asr.py @@ -326,7 +326,7 @@ def main(): parser.add_argument("--output", type=Path, help="Defaults to the ONNX artifact directory for --task") parser.add_argument("--task", choices=["asr", "aligner"], default="asr") parser.add_argument("--dtype", choices=["original", "fp16", "fp32"], default="original") - parser.add_argument("--only", choices=["mel", "encoder", "decoder", "aligner", "japanese"]) + parser.add_argument("--only", choices=["mel", "encoder", "decoder", "aligner"]) parser.add_argument("--threads", type=int, default=4) parser.add_argument("--cache-capacity", type=int, choices=range(512, 16385, 512), default=8192, metavar="TOKENS") args = parser.parse_args() @@ -337,14 +337,6 @@ def main(): precision = "bf16" if args.dtype == "original" else args.dtype size_suffix = "-1.7b" if args.task == "asr" and "1.7b" in args.model.lower() else "" args.output = args.output or Path(f"artifacts/qwen3/{prefix}onnx-{precision}{size_suffix}") - if args.only == "japanese" and args.task != "aligner": - parser.error("--only japanese requires --task aligner") - if args.task == "aligner" and args.only in (None, "japanese"): - from asr.qwen3.model_export.detail.japanese import export_japanese - - export_japanese(args.output) - if args.only == "japanese": - return if args.only in ("decoder", "aligner") and args.only != ("decoder" if args.task == "asr" else "aligner"): parser.error("--only must match --task") metadata_path = args.output / "metadata.json" diff --git a/asr/qwen3/qwen3.cpp b/asr/qwen3/qwen3.cpp index 85e488d..ec9eb84 100644 --- a/asr/qwen3/qwen3.cpp +++ b/asr/qwen3/qwen3.cpp @@ -8,8 +8,8 @@ #include #include -#include "detail/runtime.h" -#include "detail/text.h" +#include "runtime.h" +#include "unicode_regex.h" namespace din::asr::qwen3::detail { @@ -345,7 +345,7 @@ struct Qwen3Pipeline::Impl while (!ids.empty()) { prefix = asr.tokenizer->Decode(ids, false); - if (ValidUtf8(prefix) && prefix.find("\xef\xbf\xbd") == std::string::npos) + if (din::io::ValidUtf8(prefix) && prefix.find("\xef\xbf\xbd") == std::string::npos) break; ids.pop_back(); prefix.clear(); diff --git a/asr/qwen3/detail/runtime.h b/asr/qwen3/runtime.h similarity index 100% rename from asr/qwen3/detail/runtime.h rename to asr/qwen3/runtime.h diff --git a/common/io/unicode_regex.cpp b/common/io/unicode_regex.cpp index f0adb66..1397b18 100644 --- a/common/io/unicode_regex.cpp +++ b/common/io/unicode_regex.cpp @@ -6,9 +6,24 @@ #define PCRE2_CODE_UNIT_WIDTH 8 #include +#include namespace din::io { +bool ValidUtf8(std::string_view text) +{ + for (size_t i = 0; i < text.size();) + { + utf8proc_int32_t code; + const auto n = + utf8proc_iterate(reinterpret_cast(text.data() + i), text.size() - i, &code); + if (n < 0) + return false; + i += n; + } + return true; +} + UnicodeRegex::UnicodeRegex(std::string_view pattern) { int error; diff --git a/common/io/unicode_regex.h b/common/io/unicode_regex.h index 7bbc463..59f6e66 100644 --- a/common/io/unicode_regex.h +++ b/common/io/unicode_regex.h @@ -9,6 +9,8 @@ struct pcre2_real_code_8; namespace din::io { +bool ValidUtf8(std::string_view text); + // Compiled once; matching uses per-call state and validates UTF-8. class UnicodeRegex { From bd067b1b458872a2d511ae8617a2172b9cee06fb Mon Sep 17 00:00:00 2001 From: contentis Date: Wed, 23 Sep 2026 12:19:46 +0200 Subject: [PATCH 7/8] Simplify Qwen3 runtime and validate partial export checkpoints --- asr/qwen3/forced_aligner.cpp | 114 ++++++++----------- asr/qwen3/model_export/export_qwen3_asr.py | 22 ++-- asr/qwen3/model_export/validate_qwen3_asr.py | 38 ++----- asr/qwen3/qwen3.cpp | 3 - asr/qwen3/runtime.h | 7 +- common/io/tokenizer.cpp | 2 +- 6 files changed, 74 insertions(+), 112 deletions(-) diff --git a/asr/qwen3/forced_aligner.cpp b/asr/qwen3/forced_aligner.cpp index 6b444ef..4aada7c 100644 --- a/asr/qwen3/forced_aligner.cpp +++ b/asr/qwen3/forced_aligner.cpp @@ -97,42 +97,56 @@ std::vector FixTimestamps(const std::vector& data) return result; } -struct AlignmentEngine +} // namespace din::asr::qwen3::detail + +namespace din::asr::qwen3 { +using namespace detail; +struct Qwen3ForcedAligner::Impl +{ + ForcedAlignerConfig config; + Runtime runtime; AudioModel model; std::unique_ptr text; - AlignmentEngine(Runtime& runtime, const std::filesystem::path& dir); - std::vector Align(const std::vector& features, const std::vector& words); -}; - -AlignmentEngine::AlignmentEngine(Runtime& runtime, const std::filesystem::path& dir) - : model(runtime, dir, "aligner") -{ - if (!model.metadata.value("timestamp_bins", false)) - throw std::runtime_error("Re-export the aligner with --only aligner for GPU timestamp selection"); - const auto shape = [&](int64_t seq, int slots) + explicit Impl(ForcedAlignerConfig cfg) + : config(std::move(cfg)) + , runtime(config.provider, config.ep_cache_dir, config.ep_context_dir) + , model(runtime, config.model_dir, "aligner") { - return TextShape(seq, model.hidden, seq) + ",timestamp_indices:" + std::to_string(slots); - }; - text = runtime.Runner(dir, "aligner", shape(4, 2), shape(128, 32), shape(8192, 4096)); -} -std::vector AlignmentEngine::Align(const std::vector& features, - const std::vector& words) + if (!model.metadata.value("timestamp_bins", false)) + throw std::runtime_error("Re-export the aligner with --only aligner for GPU timestamp selection"); + const auto shape = [&](int64_t seq, int slots) + { + return TextShape(seq, model.hidden, seq) + ",timestamp_indices:" + std::to_string(slots); + }; + text = runtime.Runner(config.model_dir, "aligner", shape(4, 2), shape(128, 32), shape(8192, 4096)); + runtime.LoadMel(config.model_dir); + } + std::vector AlignFeatures(const std::vector& features, const std::vector& words); + std::vector AlignAudio(const din::io::Audio& audio, const std::vector& words) + { + if (words.size() > 2048) + throw std::invalid_argument("Native alignment is limited to 2048 alignment units per chunk"); + for (const auto& word : words) + if (word.empty() || !din::io::ValidUtf8(word)) + throw std::invalid_argument("Alignment units must be nonempty UTF-8 text"); + if (words.empty()) + return {}; + din::io::Audio normalized; + const auto& source = NormalizeAudio(audio, normalized); + if (source.samples.size() > 180 * kRate) + throw std::invalid_argument("Standalone alignment accepts up to 180 seconds; supply audio/text segments"); + return AlignFeatures(runtime.Features(source.samples), words); + } +}; +std::vector Qwen3ForcedAligner::Impl::AlignFeatures(const std::vector& features, + const std::vector& words) { din::common::nvtx_scoped_range range{"qwen3.align"}; const int64_t frames = features.size() / 128; std::vector result; - if (words.empty()) - return {}; - if (words.size() > 2048) - throw std::runtime_error("Native alignment is limited to 2048 alignment units per chunk"); - auto audio = [&] - { - din::common::nvtx_scoped_range range{"qwen3.align_encoder"}; - return model.Encode(features, frames); - }(); std::vector ids = model.native["audio_start"]; - ids.insert(ids.end(), audio.tokens, model.audio_id); + ids.insert(ids.end(), AudioTokens(frames), model.audio_id); Append(ids, model.native["audio_end"].get>()); std::vector slots; const int64_t timestamp_id = model.metadata["timestamp_token_id"]; @@ -152,6 +166,11 @@ std::vector AlignmentEngine::Align(const std::vector& feat const auto seq = static_cast(ids.size()); if (seq > 8192) throw std::runtime_error("Alignment exceeds the native context limit"); + auto audio = [&] + { + din::common::nvtx_scoped_range range{"qwen3.align_encoder"}; + return model.Encode(features, frames); + }(); TextInputs input(*text, seq, model.hidden, seq, model.dtype); input.Fill(ids, 0, model.audio_id, &audio); const int64_t labels = model.metadata["num_labels"]; @@ -188,42 +207,6 @@ std::vector AlignmentEngine::Align(const std::vector& feat result.push_back({words[i], times[2 * i] / 1000.f, times[2 * i + 1] / 1000.f}); return result; } -} // namespace din::asr::qwen3::detail - -namespace din::asr::qwen3 -{ -using namespace detail; -struct Qwen3ForcedAligner::Impl -{ - ForcedAlignerConfig config; - Runtime runtime; - AlignmentEngine engine; - explicit Impl(ForcedAlignerConfig cfg) - : config(std::move(cfg)) - , runtime(config.provider, config.ep_cache_dir, config.ep_context_dir) - , engine(runtime, config.model_dir) - { - runtime.LoadMel(config.model_dir); - } - std::vector Units(const std::string& text, const std::string& language) - { - return AlignmentUnits(text, language, config.granularity); - } - std::vector AlignAudio(const din::io::Audio& audio, const std::vector& words) - { - for (const auto& word : words) - if (word.empty() || !din::io::ValidUtf8(word)) - throw std::invalid_argument("Alignment units must be nonempty UTF-8 text"); - if (words.empty()) - return {}; - din::io::Audio normalized; - const auto& source = NormalizeAudio(audio, normalized); - if (source.samples.size() > 180 * kRate) - throw std::invalid_argument("Standalone alignment accepts up to 180 seconds; supply audio/text segments"); - const auto features = runtime.Features(source.samples); - return engine.Align(features, words); - } -}; Qwen3ForcedAligner::Qwen3ForcedAligner(ForcedAlignerConfig config) { impl_ = std::make_unique(std::move(config)); @@ -232,7 +215,7 @@ Qwen3ForcedAligner::~Qwen3ForcedAligner() = default; std::vector Qwen3ForcedAligner::Align(const din::io::Audio& audio, const std::string& text, const std::string& language) { - return impl_->AlignAudio(audio, impl_->Units(text, language)); + return impl_->AlignAudio(audio, AlignmentUnits(text, language, impl_->config.granularity)); } std::vector Qwen3ForcedAligner::AlignUnits(const din::io::Audio& audio, const std::vector& units) @@ -262,7 +245,8 @@ std::vector Qwen3ForcedAligner::AlignSegments(const din::io::Audi for (const auto& segment : segments) { clip.samples.assign(audio.samples.begin() + segment.start_sample, audio.samples.begin() + segment.end_sample); - auto aligned = impl_->AlignAudio(clip, impl_->Units(segment.text, segment.language)); + auto aligned = + impl_->AlignAudio(clip, AlignmentUnits(segment.text, segment.language, impl_->config.granularity)); const float offset = static_cast(segment.start_sample) / kRate; for (auto& word : aligned) { diff --git a/asr/qwen3/model_export/export_qwen3_asr.py b/asr/qwen3/model_export/export_qwen3_asr.py index a46fefc..feeace2 100644 --- a/asr/qwen3/model_export/export_qwen3_asr.py +++ b/asr/qwen3/model_export/export_qwen3_asr.py @@ -64,7 +64,6 @@ def encode(text): data["prefixes"] = prefixes data["suffixes"] = suffixes data["languages"] = languages - data["suffix"] = suffixes["auto"] (output / "native.json").write_text(json.dumps(data, indent=2), encoding="utf-8") @@ -321,7 +320,7 @@ def export_aligner(model, output): def main(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--model", "--checkpoint", dest="model", help="HF model ID or local checkpoint directory") - parser.add_argument("--size", choices=["0.6B", "1.7B"], default="0.6B", help="ASR model size") + parser.add_argument("--size", choices=["0.6B", "1.7B"], help="ASR model size (default: 0.6B)") parser.add_argument("--revision", help="Optional HF revision or commit") parser.add_argument("--output", type=Path, help="Defaults to the ONNX artifact directory for --task") parser.add_argument("--task", choices=["asr", "aligner"], default="asr") @@ -331,22 +330,25 @@ def main(): parser.add_argument("--cache-capacity", type=int, choices=range(512, 16385, 512), default=8192, metavar="TOKENS") args = parser.parse_args() prefix = "aligner-" if args.task == "aligner" else "" - args.model = args.model or ( - "Qwen/Qwen3-ForcedAligner-0.6B-hf" if args.task == "aligner" else f"Qwen/Qwen3-ASR-{args.size}-hf" + default_model = ( + "Qwen/Qwen3-ForcedAligner-0.6B-hf" if args.task == "aligner" else f"Qwen/Qwen3-ASR-{args.size or '0.6B'}-hf" ) precision = "bf16" if args.dtype == "original" else args.dtype - size_suffix = "-1.7b" if args.task == "asr" and "1.7b" in args.model.lower() else "" + size_suffix = "-1.7b" if args.task == "asr" and "1.7b" in (args.model or default_model).lower() else "" args.output = args.output or Path(f"artifacts/qwen3/{prefix}onnx-{precision}{size_suffix}") if args.only in ("decoder", "aligner") and args.only != ("decoder" if args.task == "asr" else "aligner"): parser.error("--only must match --task") metadata_path = args.output / "metadata.json" existing = json.loads(metadata_path.read_text(encoding="utf-8")) if args.only and metadata_path.exists() else {} - requested = {"original": "bfloat16", "fp16": "float16", "fp32": "float32"}[args.dtype] - if ( - existing - and args.only != "mel" - and (existing["dtype"] != requested or existing["task"] != args.task or existing.get("quantization")) + source = existing.get("source", {}) + args.model = args.model or (source.get("model") if args.size is None else None) or default_model + args.revision = args.revision or source.get("revision") + if existing and ( + existing["task"] != args.task or source["model"] != args.model or source.get("revision") != args.revision ): + parser.error("Partial export requires the original task and checkpoint; use a new directory") + requested = {"original": "bfloat16", "fp16": "float16", "fp32": "float32"}[args.dtype] + if existing and args.only != "mel" and (existing["dtype"] != requested or existing.get("quantization")): parser.error("Partial export requires the same task and precision, without quantization; use a new directory") torch.set_num_threads(args.threads) args.output.mkdir(parents=True, exist_ok=True) diff --git a/asr/qwen3/model_export/validate_qwen3_asr.py b/asr/qwen3/model_export/validate_qwen3_asr.py index 5f3f7ec..9ea1f41 100644 --- a/asr/qwen3/model_export/validate_qwen3_asr.py +++ b/asr/qwen3/model_export/validate_qwen3_asr.py @@ -34,23 +34,6 @@ def session_options(provider, threads, extra_options=None): return options, None -def pack_audio(features, mask, config): - chunk_size = config["n_window"] * 2 - if features.shape[0] != 1 or features.shape[-1] % chunk_size: - raise ValueError("Expected batch-one features padded to the encoder chunk size") - chunks = features.reshape(1, features.shape[1], -1, chunk_size)[0].permute(1, 0, 2) - lengths = mask.reshape(-1, chunk_size).sum(1) - for _ in range(3): - lengths = (lengths + 1) // 2 - indices = (torch.arange((chunk_size + 7) // 8)[None] < lengths[:, None]).flatten().nonzero().flatten() - if not len(indices): - raise ValueError("Audio has no valid feature frames") - window = int(lengths.max()) * (config["n_window_infer"] // chunk_size) - groups = torch.arange(len(indices)) // window - bias = torch.zeros(len(indices), len(indices)).masked_fill(groups[:, None] != groups[None], -1e4) - return chunks, indices, bias[None, None] - - def decoder_inputs(ids, embeddings, audio_token_id, hidden_size, past_length, capacity=None): ids = ids.cpu().long().reshape(1, -1) seq = ids.shape[1] @@ -75,10 +58,8 @@ def decoder_inputs(ids, embeddings, audio_token_id, hidden_size, past_length, ca class OnnxAudioModel: def __init__(self, directory, task, threads=4, provider=None): - self.text_graph = "decoder.onnx" if task == "asr" else "aligner.onnx" + text_graph = "decoder.onnx" if task == "asr" else "aligner.onnx" directory = Path(directory) - self.directory = directory - self.threads = threads self.metadata = json.loads((directory / "metadata.json").read_text(encoding="utf-8")) if self.metadata["task"] != task or self.metadata["format_version"] != (3 if task == "asr" else 2): raise ValueError("Incompatible export; re-export the model with the current exporter") @@ -90,7 +71,7 @@ def __init__(self, directory, task, threads=4, provider=None): self.provider, threads, {f"nv_profile_{bound}_shapes": encoder_shape for bound in ("min", "opt", "max")} ) self.encoder = ort.InferenceSession(str(directory / "encoder.onnx"), options, providers=providers) - if self.text_graph == "decoder.onnx": + if text_graph == "decoder.onnx": c = self.metadata["text_config"] capacity = self.metadata["cache_capacity"] @@ -108,11 +89,11 @@ def shapes(sequence): threads, {f"nv_profile_{bound}_shapes": shapes(sequence) for bound in ("min", "opt", "max")}, ) - sessions.append(ort.InferenceSession(str(directory / self.text_graph), options, providers=providers)) + sessions.append(ort.InferenceSession(str(directory / text_graph), options, providers=providers)) self.decoder, self.prefill = sessions else: options, providers = session_options(self.provider, threads) - self.decoder = ort.InferenceSession(str(directory / self.text_graph), options, providers=providers) + self.decoder = ort.InferenceSession(str(directory / text_graph), options, providers=providers) def run(self, session, feed, inplace=False): values = {} @@ -150,18 +131,19 @@ def encode(self, inputs): features, mask = inputs["input_features"], inputs["input_features_mask"] outputs = [] for start in range(0, int(mask.sum()), window): - chunks, indices, _ = pack_audio( - features[..., start : start + window], mask[..., start : start + window], config - ) + chunk_size = config["n_window"] * 2 + chunks = features[..., start : start + window].reshape(128, -1, chunk_size).permute(1, 0, 2) + lengths = mask[..., start : start + window].reshape(-1, chunk_size).sum(1) + valid_tokens = int(((lengths + 7) // 8).sum()) tokens = window // 100 * 13 bias = torch.zeros(1, 1, tokens, tokens) - bias[..., len(indices) :] = -1e4 + bias[..., valid_tokens:] = -1e4 packed = torch.zeros(window // 100, 128, 100) packed[: len(chunks)] = chunks outputs.append( self.run( self.encoder, {"mel_chunks": packed, "valid_indices": torch.arange(tokens), "attention_bias": bias} - )[0][: len(indices)] + )[0][:valid_tokens] ) return torch.cat(outputs) diff --git a/asr/qwen3/qwen3.cpp b/asr/qwen3/qwen3.cpp index ec9eb84..ca6c992 100644 --- a/asr/qwen3/qwen3.cpp +++ b/asr/qwen3/qwen3.cpp @@ -82,10 +82,7 @@ inline std::vector SplitAudio(std::span samples, size_t chunks.push_back({start, samples.size()}); return chunks; } -} // namespace din::asr::qwen3::detail -namespace din::asr::qwen3::detail -{ std::string Trim(const std::string& text) { const auto first = text.find_first_not_of(" \r\n\t"); diff --git a/asr/qwen3/runtime.h b/asr/qwen3/runtime.h index e539dee..6ecc177 100644 --- a/asr/qwen3/runtime.h +++ b/asr/qwen3/runtime.h @@ -433,15 +433,14 @@ struct AudioModel struct EncoderBuffers { - int64_t frames, tokens; + int64_t tokens; FloatBuffer mel, bias; Buffer indices; Ort::Value output; Ort::IoBinding binding; EncoderBuffers(OrtRunner& runner, int64_t count, int64_t hidden, ONNXTensorElementDataType dtype) - : frames(count) - , tokens(AudioTokens(count)) + : tokens(AudioTokens(count)) , mel(runner, {(count + 99) / 100, 128, 100}, dtype, true) , bias(runner, {1, 1, tokens, tokens}, dtype, true) , indices(runner, {tokens}, runner.HasDeviceIo()) @@ -449,10 +448,8 @@ struct AudioModel , binding(runner.session) { // A call contains exactly one independent HF encoder window. - bias.Fill(0.f); for (int64_t i = 0; i < tokens; ++i) indices.HostData()[i] = i; - bias.CopyAsyncToDevice(); indices.CopyAsyncToDevice(); binding.BindInput("mel_chunks", mel.BindingValue()); binding.BindInput("valid_indices", indices.BindingValue()); diff --git a/common/io/tokenizer.cpp b/common/io/tokenizer.cpp index 46c7671..99eb7d7 100644 --- a/common/io/tokenizer.cpp +++ b/common/io/tokenizer.cpp @@ -705,7 +705,7 @@ std::string Tokenizer::Decode(const std::vector& ids, bool skip_special { if (decode_mode_ == DecodeMode::ByteBpe) { - const auto byte_decoder = BuildByteDecoder(); + static const auto byte_decoder = BuildByteDecoder(); std::string text; for (auto id : ids) { From 1a9a191c01a7eaf0cf8439b7bdaa7771ea8ec738 Mon Sep 17 00:00:00 2001 From: lspindler Date: Thu, 24 Sep 2026 09:56:46 +0200 Subject: [PATCH 8/8] Fix Qwen3 JSON diagnostics and truncated UTF-8 output Signed-off-by: lspindler --- asr/qwen3/main.cpp | 24 ++++++------- asr/qwen3/qwen3.cpp | 6 +++- asr/qwen3/tests/CMakeLists.txt | 3 ++ asr/qwen3/tests/text_output_test.cpp | 50 ++++++++++++++++++++++++++++ asr/qwen3/text_output.h | 35 +++++++++++++++++++ common/ort_session.cpp | 18 +++++----- 6 files changed, 112 insertions(+), 24 deletions(-) create mode 100644 asr/qwen3/tests/CMakeLists.txt create mode 100644 asr/qwen3/tests/text_output_test.cpp create mode 100644 asr/qwen3/text_output.h diff --git a/asr/qwen3/main.cpp b/asr/qwen3/main.cpp index 676384a..0bf2c0d 100644 --- a/asr/qwen3/main.cpp +++ b/asr/qwen3/main.cpp @@ -46,26 +46,22 @@ int main(int argc, char** argv) config.ep_context_dir = parser.get("--ep-context-dir"); config.max_new_tokens = parser.get("--max-new-tokens"); config.max_chunk_seconds = parser.get("--max-chunk-seconds"); - std::ostream results(std::cout.rdbuf()); - // Keep runtime diagnostics off the streaming JSON output. - if (parser.get("--stream")) - std::cout.rdbuf(std::cerr.rdbuf()); Qwen3Pipeline pipeline(std::move(config)); if (parser.get("--stream")) { pipeline.StartStream({parser.get("--chunk-seconds"), parser.get("--unfixed-chunks"), parser.get("--unfixed-tokens")}); - auto print = [&results](const StreamingResult& result) + auto print = [](const StreamingResult& result) { - results << nlohmann::json{{"transcription", result.text}, - {"language", result.language}, - {"samples_processed", result.samples_processed}, - {"sample_rate", 16000}, - {"updates", result.updates}, - {"final", result.final}, - {"reached_eos", result.reached_eos}} - .dump() - << std::endl; + std::cout << nlohmann::json{{"transcription", result.text}, + {"language", result.language}, + {"samples_processed", result.samples_processed}, + {"sample_rate", 16000}, + {"updates", result.updates}, + {"final", result.final}, + {"reached_eos", result.reached_eos}} + .dump() + << std::endl; }; auto feed = [&](std::span samples) { diff --git a/asr/qwen3/qwen3.cpp b/asr/qwen3/qwen3.cpp index ca6c992..f73bf24 100644 --- a/asr/qwen3/qwen3.cpp +++ b/asr/qwen3/qwen3.cpp @@ -9,6 +9,7 @@ #include #include "runtime.h" +#include "text_output.h" #include "unicode_regex.h" namespace din::asr::qwen3::detail @@ -320,7 +321,10 @@ struct Qwen3Pipeline::Impl auto text_tokens = result.tokens; if (result.reached_eos) text_tokens.pop_back(); - const auto raw = prefix + asr.tokenizer->Decode(text_tokens, false); + auto decoded = asr.tokenizer->Decode(text_tokens, false); + if (!result.reached_eos) + detail::TrimIncompleteUtf8Suffix(decoded); + const auto raw = prefix + decoded; if (raw_output) *raw_output = raw; const auto parsed = detail::ParseOutput(raw, asr.native["languages"][config.lang_id].get()); diff --git a/asr/qwen3/tests/CMakeLists.txt b/asr/qwen3/tests/CMakeLists.txt new file mode 100644 index 0000000..a706074 --- /dev/null +++ b/asr/qwen3/tests/CMakeLists.txt @@ -0,0 +1,3 @@ +add_executable(din_qwen3_text_output_test text_output_test.cpp) +target_link_libraries(din_qwen3_text_output_test PRIVATE nlohmann_json::nlohmann_json) +add_test(NAME din_qwen3_text_output COMMAND din_qwen3_text_output_test) diff --git a/asr/qwen3/tests/text_output_test.cpp b/asr/qwen3/tests/text_output_test.cpp new file mode 100644 index 0000000..b72ac12 --- /dev/null +++ b/asr/qwen3/tests/text_output_test.cpp @@ -0,0 +1,50 @@ +// SPDX-License-Identifier: Apache-2.0 +#include +#include +#include +#include + +#include "../text_output.h" +#include + +int main() +{ + using din::asr::qwen3::detail::TrimIncompleteUtf8Suffix; + try + { + // Cover every truncation of two-, three-, and four-byte characters. + const std::vector characters = {"\xC2\xA2", "\xE4\xB8\xAD", "\xF0\x9F\x98\x80"}; + for (const auto& character : characters) + { + for (size_t length = 0; length <= character.size(); ++length) + { + const auto prefix = std::string("earlier chunk ") + characters[1] + " "; + auto text = prefix + character.substr(0, length); + TrimIncompleteUtf8Suffix(text); + const auto expected = prefix + (length == character.size() ? character : ""); + if (text != expected) + throw std::runtime_error("Trimming lost complete transcript text"); + const nlohmann::json output{{"transcription", text}, {"reached_eos", false}}; + const auto parsed = nlohmann::json::parse(output.dump()); + if (parsed.at("transcription") != expected || parsed.at("reached_eos") != false) + throw std::runtime_error("Partial transcript did not survive JSON serialization"); + } + } + // Do not hide other encoding errors by removing arbitrary trailing bytes. + for (const std::string original : + {"", "ASCII", "\x80", "\xFFtext", "\xC0", "\xE0\x80", "\xED\xA0", "\xF0\x80", "\xF4\x90", "\xF5\x80"}) + { + auto text = original; + TrimIncompleteUtf8Suffix(text); + if (text != original) + throw std::runtime_error("Modified complete text or an unrelated encoding error"); + } + std::cout << "Qwen3 text output checks passed\n"; + return 0; + } + catch (const std::exception& error) + { + std::cerr << error.what() << '\n'; + return 1; + } +} diff --git a/asr/qwen3/text_output.h b/asr/qwen3/text_output.h new file mode 100644 index 0000000..6a141c6 --- /dev/null +++ b/asr/qwen3/text_output.h @@ -0,0 +1,35 @@ +// SPDX-License-Identifier: Apache-2.0 +#pragma once + +#include + +namespace din::asr::qwen3::detail +{ +// A generation limit may split a byte-BPE character. Preserve complete characters +// and malformed sequences; only discard a valid prefix of an unfinished character. +inline void TrimIncompleteUtf8Suffix(std::string& text) +{ + if (text.empty()) + return; + size_t start = text.size() - 1; + while (start > 0 && (static_cast(text[start]) & 0xC0) == 0x80) + --start; + const auto lead = static_cast(text[start]); + const size_t expected = lead >= 0xC2 && lead <= 0xDF ? 2 + : lead >= 0xE0 && lead <= 0xEF ? 3 + : lead >= 0xF0 && lead <= 0xF4 ? 4 + : 0; + const auto available = text.size() - start; + if (!expected || available >= expected) + return; + if (available > 1) + { + const auto second = static_cast(text[start + 1]); + // Reject overlong encodings, surrogates, and code points beyond U+10FFFF. + if ((lead == 0xE0 && second < 0xA0) || (lead == 0xED && second > 0x9F) || (lead == 0xF0 && second < 0x90) || + (lead == 0xF4 && second > 0x8F)) + return; + } + text.resize(start); +} +} // namespace din::asr::qwen3::detail diff --git a/common/ort_session.cpp b/common/ort_session.cpp index 46558e9..17b73a5 100644 --- a/common/ort_session.cpp +++ b/common/ort_session.cpp @@ -431,10 +431,10 @@ bool RegisterTensorRTRTXExecutionProvider(Ort::Env& env) env.RegisterExecutionProviderLibrary(kDinNvTensorRTRTXExecutionProvider, provider_library_path.c_str()); const auto ep_devices = env.GetEpDevices(); - std::cout << "Execution provider devices after TRT RTX registration:\n"; + std::cerr << "Execution provider devices after TRT RTX registration:\n"; for (const auto& device : ep_devices) { - std::cout << " " << device.EpName() << " vendor=" << device.EpVendor() + std::cerr << " " << device.EpName() << " vendor=" << device.EpVendor() << " device_id=" << device.Device().DeviceId() << '\n'; } return true; @@ -633,7 +633,7 @@ int ChooseCudaDeviceOrdinal(Ort::ConstEpDevice ep_device, const char* override_e const auto ort_luid = luid_text != nullptr ? ParseUint64(*luid_text) : std::optional{}; if (luid_text != nullptr && !ort_luid.has_value()) { - std::cout << "Ignoring unparsable ORT LUID metadata value: " << *luid_text << std::endl; + std::cerr << "Ignoring unparsable ORT LUID metadata value: " << *luid_text << std::endl; } #else const std::string* pci_bus_id_text = FindMetadataValue(metadata, "pci_bus_id"); @@ -641,7 +641,7 @@ int ChooseCudaDeviceOrdinal(Ort::ConstEpDevice ep_device, const char* override_e pci_bus_id_text != nullptr ? ParsePciBusId(*pci_bus_id_text) : std::optional{}; if (pci_bus_id_text != nullptr && !ort_pci_bus_id.has_value()) { - std::cout << "Ignoring unparsable ORT pci_bus_id metadata value: " << *pci_bus_id_text << std::endl; + std::cerr << "Ignoring unparsable ORT pci_bus_id metadata value: " << *pci_bus_id_text << std::endl; } #endif @@ -693,24 +693,24 @@ int ChooseCudaDeviceOrdinal(Ort::ConstEpDevice ep_device, const char* override_e if (metadata_matches.size() > 1) { - std::cout << "Multiple CUDA devices match ORT device metadata"; + std::cerr << "Multiple CUDA devices match ORT device metadata"; } else { - std::cout << "No CUDA device matches ORT device metadata"; + std::cerr << "No CUDA device matches ORT device metadata"; } #ifdef _WIN32 if (luid_text != nullptr) { - std::cout << " LUID=" << *luid_text; + std::cerr << " LUID=" << *luid_text; } #else if (pci_bus_id_text != nullptr) { - std::cout << " pci_bus_id=" << *pci_bus_id_text; + std::cerr << " pci_bus_id=" << *pci_bus_id_text; } #endif - std::cout << " for ORT hardware device_id=" << ort_hardware_device_id << "; defaulting to CUDA ordinal 0. Set " + std::cerr << " for ORT hardware device_id=" << ort_hardware_device_id << "; defaulting to CUDA ordinal 0. Set " << (override_env_var != nullptr ? override_env_var : "FLUX_CUDA_DEVICE_ID") << " to override." << std::endl;