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..a721e4e 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 ([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: |
| 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,213 @@ 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
+
+## 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
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..033ab19
--- /dev/null
+++ b/asr/qwen3/README.md
@@ -0,0 +1,90 @@
+# Qwen3 ASR and forced alignment
+
+ONNX Runtime inference on CPU or TensorRT RTX, with independent ASR and alignment APIs.
+
+## Supported models
+
+| 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
+
+| Model / upstream toolkit capability | C++ sample |
+|---|:---:|
+| Offline, single stream | ✓ |
+| Online / streaming, single utterance | ✓ |
+| Batched inference | — |
+| 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 | — |
+| Word timestamps: en, de, es, fr, it, pt, ru, ko | ✓ |
+| Chinese / Cantonese character timestamps | ✓ |
+| Automatic Japanese word timestamps | — |
+| Character timestamps: all 11 alignment languages (sample extension) | ✓ |
+| Caller-supplied alignment units | ✓ |
+
+See [upstream language support](https://github.com/QwenLM/Qwen3-ASR).
+Japanese alignment requires character mode or supplied units.
+
+## Export and validate
+
+Run commands from the repository root; install the export dependencies first.
+
+```bash
+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
+```
+
+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 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
+```
+
+## Build and run
+
+Follow the [repository build setup](../../README.md). Replace `` and
+`` with your build and executable directories; append `.exe` on Windows.
+
+```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
+```
+
+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`.
+
+## Long-form and streaming
+
+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.
+
+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.
+
+## C++ integration
+
+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/aligner_main.cpp b/asr/qwen3/aligner_main.cpp
new file mode 100644
index 0000000..f3e74e3
--- /dev/null
+++ b/asr/qwen3/aligner_main.cpp
@@ -0,0 +1,81 @@
+// SPDX-License-Identifier: Apache-2.0
+#include
+#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("--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());
+ 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");
+ 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)
+ 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();
+ 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},
+ {"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/forced_aligner.cpp b/asr/qwen3/forced_aligner.cpp
new file mode 100644
index 0000000..4aada7c
--- /dev/null
+++ b/asr/qwen3/forced_aligner.cpp
@@ -0,0 +1,265 @@
+// SPDX-License-Identifier: Apache-2.0
+#include "forced_aligner.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)
+{
+ 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");
+
+ if (!din::io::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")
+ 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}]+)");
+ 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}]+)");
+ 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;
+}
+
+} // 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;
+ 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")
+ {
+ 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;
+ std::vector ids = model.native["audio_start"];
+ 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"];
+ for (const auto& word : words)
+ {
+ 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());
+ 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");
+ 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"];
+ 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;
+}
+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, AlignmentUnits(text, language, impl_->config.granularity));
+}
+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, AlignmentUnits(segment.text, segment.language, impl_->config.granularity));
+ 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)
+{
+ 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..586979b
--- /dev/null
+++ b/asr/qwen3/forced_aligner.h
@@ -0,0 +1,61 @@
+// SPDX-License-Identifier: Apache-2.0
+#pragma once
+#include
+#include
+#include
+#include
+#include
+
+#include "audio.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";
+ AlignmentGranularity granularity = AlignmentGranularity::Words;
+};
+
+struct WordTimestamp
+{
+ std::string text;
+ float start_time = 0;
+ 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
+{
+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");
+ 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;
+ 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..0bf2c0d
--- /dev/null
+++ b/asr/qwen3/main.cpp
@@ -0,0 +1,118 @@
+// SPDX-License-Identifier: Apache-2.0
+#include
+#include
+#include
+#ifdef _WIN32
+#include
+#include
+#endif
+
+#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 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());
+ 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("--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>()
+ .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));
+ if (parser.get("--stream"))
+ {
+ pipeline.StartStream({parser.get("--chunk-seconds"), parser.get("--unfixed-chunks"),
+ parser.get("--unfixed-tokens")});
+ auto print = [](const StreamingResult& result)
+ {
+ 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)
+ {
+ 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},
+ {"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..feeace2
--- /dev/null
+++ b/asr/qwen3/model_export/export_qwen3_asr.py
@@ -0,0 +1,414 @@
+# 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
+ (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"], 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")
+ 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 ""
+ 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 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 {}
+ 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)
+ 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..9ea1f41
--- /dev/null
+++ b/asr/qwen3/model_export/validate_qwen3_asr.py
@@ -0,0 +1,373 @@
+# 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 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):
+ text_graph = "decoder.onnx" if task == "asr" else "aligner.onnx"
+ directory = Path(directory)
+ 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 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 / text_graph), options, providers=providers))
+ self.decoder, self.prefill = sessions
+ else:
+ options, providers = session_options(self.provider, threads)
+ self.decoder = ort.InferenceSession(str(directory / 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):
+ 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[..., 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][:valid_tokens]
+ )
+ 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..f73bf24
--- /dev/null
+++ b/asr/qwen3/qwen3.cpp
@@ -0,0 +1,454 @@
+// SPDX-License-Identifier: Apache-2.0
+#include "qwen3.h"
+
+#include
+#include
+#include
+#include
+#include
+#include
+
+#include "runtime.h"
+#include "text_output.h"
+#include "unicode_regex.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;
+}
+
+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;
+ 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)
+ , 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, 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");
+ std::vector ids = asr.native["prefixes"][config.lang_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();
+ // 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();
+ 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());
+ 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 (din::io::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);
+ 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;
+ }
+ 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;
+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)
+{
+ 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..303265f
--- /dev/null
+++ b/asr/qwen3/qwen3.h
@@ -0,0 +1,84 @@
+// SPDX-License-Identifier: Apache-2.0
+#pragma once
+
+#include
+#include
+#include
+#include
+#include
+#include
+
+#include "audio.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.
+};
+
+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;
+};
+
+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
+{
+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);
+ // 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;
+ 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/asr/qwen3/runtime.h b/asr/qwen3/runtime.h
new file mode 100644
index 0000000..6ecc177
--- /dev/null
+++ b/asr/qwen3/runtime.h
@@ -0,0 +1,549 @@
+// 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;
+ 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)
+ : provider(execution_provider)
+ , ep_cache_dir(cache)
+ , ep_context_dir(context)
+ {
+ 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()}, 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 tokens;
+ FloatBuffer mel, bias;
+ Buffer indices;
+ Ort::Value output;
+ Ort::IoBinding binding;
+
+ EncoderBuffers(OrtRunner& runner, int64_t count, int64_t hidden, ONNXTensorElementDataType dtype)
+ : 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.
+ for (int64_t i = 0; i < tokens; ++i)
+ indices.HostData()[i] = i;
+ 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/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