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 + +#include "../text_output.h" +#include + +int main() +{ + using din::asr::qwen3::detail::TrimIncompleteUtf8Suffix; + try + { + // Cover every truncation of two-, three-, and four-byte characters. + const std::vector characters = {"\xC2\xA2", "\xE4\xB8\xAD", "\xF0\x9F\x98\x80"}; + for (const auto& character : characters) + { + for (size_t length = 0; length <= character.size(); ++length) + { + const auto prefix = std::string("earlier chunk ") + characters[1] + " "; + auto text = prefix + character.substr(0, length); + TrimIncompleteUtf8Suffix(text); + const auto expected = prefix + (length == character.size() ? character : ""); + if (text != expected) + throw std::runtime_error("Trimming lost complete transcript text"); + const nlohmann::json output{{"transcription", text}, {"reached_eos", false}}; + const auto parsed = nlohmann::json::parse(output.dump()); + if (parsed.at("transcription") != expected || parsed.at("reached_eos") != false) + throw std::runtime_error("Partial transcript did not survive JSON serialization"); + } + } + // Do not hide other encoding errors by removing arbitrary trailing bytes. + for (const std::string original : + {"", "ASCII", "\x80", "\xFFtext", "\xC0", "\xE0\x80", "\xED\xA0", "\xF0\x80", "\xF4\x90", "\xF5\x80"}) + { + auto text = original; + TrimIncompleteUtf8Suffix(text); + if (text != original) + throw std::runtime_error("Modified complete text or an unrelated encoding error"); + } + std::cout << "Qwen3 text output checks passed\n"; + return 0; + } + catch (const std::exception& error) + { + std::cerr << error.what() << '\n'; + return 1; + } +} diff --git a/asr/qwen3/text_output.h b/asr/qwen3/text_output.h new file mode 100644 index 0000000..6a141c6 --- /dev/null +++ b/asr/qwen3/text_output.h @@ -0,0 +1,35 @@ +// SPDX-License-Identifier: Apache-2.0 +#pragma once + +#include + +namespace din::asr::qwen3::detail +{ +// A generation limit may split a byte-BPE character. Preserve complete characters +// and malformed sequences; only discard a valid prefix of an unfinished character. +inline void TrimIncompleteUtf8Suffix(std::string& text) +{ + if (text.empty()) + return; + size_t start = text.size() - 1; + while (start > 0 && (static_cast(text[start]) & 0xC0) == 0x80) + --start; + const auto lead = static_cast(text[start]); + const size_t expected = lead >= 0xC2 && lead <= 0xDF ? 2 + : lead >= 0xE0 && lead <= 0xEF ? 3 + : lead >= 0xF0 && lead <= 0xF4 ? 4 + : 0; + const auto available = text.size() - start; + if (!expected || available >= expected) + return; + if (available > 1) + { + const auto second = static_cast(text[start + 1]); + // Reject overlong encodings, surrogates, and code points beyond U+10FFFF. + if ((lead == 0xE0 && second < 0xA0) || (lead == 0xED && second > 0x9F) || (lead == 0xF0 && second < 0x90) || + (lead == 0xF4 && second > 0x8F)) + return; + } + text.resize(start); +} +} // namespace din::asr::qwen3::detail diff --git a/cmake/pcre2.cmake b/cmake/pcre2.cmake new file mode 100644 index 0000000..215d138 --- /dev/null +++ b/cmake/pcre2.cmake @@ -0,0 +1,16 @@ +# SPDX-License-Identifier: Apache-2.0 +include_guard(GLOBAL) +include(FetchContent) + +set(PCRE2_BUILD_PCRE2_8 ON CACHE BOOL "" FORCE) +set(PCRE2_BUILD_PCRE2_16 OFF CACHE BOOL "" FORCE) +set(PCRE2_BUILD_PCRE2_32 OFF CACHE BOOL "" FORCE) +set(PCRE2_BUILD_PCRE2GREP OFF CACHE BOOL "" FORCE) +set(PCRE2_BUILD_TESTS OFF CACHE BOOL "" FORCE) +set(PCRE2_STATIC_PIC ON CACHE BOOL "" FORCE) +FetchContent_Declare(pcre2 + GIT_REPOSITORY https://github.com/PCRE2Project/pcre2.git + GIT_TAG pcre2-10.46 + GIT_SHALLOW TRUE +) +FetchContent_MakeAvailable(pcre2) diff --git a/cmake/utf8proc.cmake b/cmake/utf8proc.cmake new file mode 100644 index 0000000..c4addbd --- /dev/null +++ b/cmake/utf8proc.cmake @@ -0,0 +1,12 @@ +# SPDX-License-Identifier: Apache-2.0 +include_guard(GLOBAL) +include(FetchContent) + +set(UTF8PROC_ENABLE_TESTING OFF CACHE BOOL "" FORCE) +FetchContent_Declare(utf8proc + GIT_REPOSITORY https://github.com/JuliaStrings/utf8proc.git + GIT_TAG v2.11.0 + GIT_SHALLOW TRUE +) +FetchContent_MakeAvailable(utf8proc) +set_target_properties(utf8proc PROPERTIES POSITION_INDEPENDENT_CODE ON) diff --git a/common/io/CMakeLists.txt b/common/io/CMakeLists.txt index e84b9a9..e613f5a 100644 --- a/common/io/CMakeLists.txt +++ b/common/io/CMakeLists.txt @@ -1,9 +1,12 @@ include(${PROJECT_SOURCE_DIR}/cmake/lodepng.cmake) include(${PROJECT_SOURCE_DIR}/cmake/miniaudio.cmake) +include(${PROJECT_SOURCE_DIR}/cmake/pcre2.cmake) +include(${PROJECT_SOURCE_DIR}/cmake/utf8proc.cmake) add_library(din_common_io STATIC image.cpp tokenizer.cpp + unicode_regex.cpp ) target_include_directories(din_common_io @@ -16,6 +19,8 @@ target_link_libraries(din_common_io PRIVATE lodepng nlohmann_json::nlohmann_json + pcre2-8-static + utf8proc ) set_target_properties(din_common_io PROPERTIES diff --git a/common/io/tokenizer.cpp b/common/io/tokenizer.cpp index 86dcc97..99eb7d7 100644 --- a/common/io/tokenizer.cpp +++ b/common/io/tokenizer.cpp @@ -7,15 +7,19 @@ #include #include #include +#include #include #include +#include #include #include #include #include #include +#include "unicode_regex.h" #include +#include namespace din::io { @@ -103,6 +107,8 @@ struct JsonTokenizerData std::unordered_map token_to_id; std::unordered_map bpe_ranks; std::vector special_tokens; + std::vector added_tokens; + bool normalize_nfc = false; }; void AddToken(JsonTokenizerData& data, const std::string& token, int64_t id) @@ -133,6 +139,8 @@ JsonTokenizerData LoadJsonTokenizerData(const std::string& path) } JsonTokenizerData tokenizer_data; + tokenizer_data.normalize_nfc = + data.contains("normalizer") && data["normalizer"].is_object() && data["normalizer"].value("type", "") == "NFC"; const auto& vocab = data["model"]["vocab"]; for (const auto& [token, id_val] : vocab.items()) @@ -150,6 +158,7 @@ JsonTokenizerData LoadJsonTokenizerData(const std::string& path) } const auto content = token["content"].get(); AddToken(tokenizer_data, content, token["id"].get()); + tokenizer_data.added_tokens.push_back(content); if (token.value("special", false)) { tokenizer_data.special_tokens.push_back(content); @@ -194,6 +203,11 @@ JsonTokenizerData LoadJsonTokenizerData(const std::string& path) { return left.size() > right.size(); }); + std::sort(tokenizer_data.added_tokens.begin(), tokenizer_data.added_tokens.end(), + [](const auto& left, const auto& right) + { + return left.size() > right.size(); + }); return tokenizer_data; } @@ -347,14 +361,17 @@ Tokenizer::Tokenizer(const std::string& path, TokenizerFormat format) switch (format) { case TokenizerFormat::Json: + case TokenizerFormat::ByteBpeJson: { auto tokenizer_data = LoadJsonTokenizerData(path); id_to_token_ = std::move(tokenizer_data.id_to_token); token_to_id_ = std::move(tokenizer_data.token_to_id); bpe_ranks_ = std::move(tokenizer_data.bpe_ranks); special_tokens_ = std::move(tokenizer_data.special_tokens); + added_tokens_ = std::move(tokenizer_data.added_tokens); + normalize_nfc_ = format == TokenizerFormat::ByteBpeJson && tokenizer_data.normalize_nfc; byte_encoder_ = BuildByteEncoder(); - decode_mode_ = DecodeMode::Pieces; + decode_mode_ = format == TokenizerFormat::ByteBpeJson ? DecodeMode::ByteBpe : DecodeMode::Pieces; break; } case TokenizerFormat::Vocab: @@ -434,8 +451,14 @@ bool IsContractionAt(std::string_view text, size_t pos, size_t* length) return false; } -std::vector SplitForByteLevelBpe(std::string_view text) +std::vector SplitForByteLevelBpe(std::string_view text, bool qwen = false) { + if (qwen) + { + static const UnicodeRegex pattern( + R"((?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\r\n\p{L}\p{N}]?\p{L}+|\p{N}| ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+)"); + return pattern.FindAll(text); + } std::vector pieces; size_t pos = 0; while (pos < text.size()) @@ -593,7 +616,7 @@ int64_t Tokenizer::TokenId(const std::string& token) const return it->second; } -std::vector Tokenizer::Encode(const std::string& text, bool add_special_tokens) const +std::vector Tokenizer::Encode(const std::string& input, bool add_special_tokens) const { if (bpe_ranks_.empty()) { @@ -601,10 +624,23 @@ std::vector Tokenizer::Encode(const std::string& text, bool add_special } (void)add_special_tokens; + std::string normalized; + if (normalize_nfc_) + { + utf8proc_uint8_t* data = nullptr; + const auto length = utf8proc_map(reinterpret_cast(input.data()), input.size(), &data, + static_cast(UTF8PROC_STABLE | UTF8PROC_COMPOSE)); + const std::unique_ptr owner(data, std::free); + if (length < 0) + throw std::invalid_argument(utf8proc_errmsg(length)); + normalized.assign(reinterpret_cast(data), static_cast(length)); + } + const auto& text = normalize_nfc_ ? normalized : input; + std::vector ids; const auto encode_segment = [&](std::string_view segment) { - for (const auto& piece : SplitForByteLevelBpe(segment)) + for (const auto& piece : SplitForByteLevelBpe(segment, decode_mode_ == DecodeMode::ByteBpe)) { for (const auto& token : ApplyByteLevelBpe(piece, byte_encoder_, bpe_ranks_)) { @@ -623,7 +659,7 @@ std::vector Tokenizer::Encode(const std::string& text, bool add_special while (pos < text.size()) { const std::string* matched_special = nullptr; - for (const auto& special : special_tokens_) + for (const auto& special : decode_mode_ == DecodeMode::ByteBpe ? added_tokens_ : special_tokens_) { if (StartsWith(std::string_view(text).substr(pos), special)) { @@ -667,6 +703,27 @@ std::string Tokenizer::CleanToken(int64_t id) const std::string Tokenizer::Decode(const std::vector& ids, bool skip_special_tokens, bool strip_lang_tags) const { + if (decode_mode_ == DecodeMode::ByteBpe) + { + static const auto byte_decoder = BuildByteDecoder(); + std::string text; + for (auto id : ids) + { + const auto& token = Token(id); + if (skip_special_tokens && + std::find(special_tokens_.begin(), special_tokens_.end(), token) != special_tokens_.end()) + continue; + for (size_t i = 0; i < token.size();) + { + const auto code = NextCodePoint(token, i); + const auto byte = byte_decoder.find(code); + if (byte == byte_decoder.end()) + throw std::runtime_error("Invalid byte-level vocabulary entry"); + text.push_back(static_cast(byte->second)); + } + } + return text; + } if (decode_mode_ == DecodeMode::WhisperByteBpe) { std::string text; diff --git a/common/io/tokenizer.h b/common/io/tokenizer.h index eec5b97..c79ac5d 100644 --- a/common/io/tokenizer.h +++ b/common/io/tokenizer.h @@ -17,6 +17,7 @@ enum class TokenizerFormat Json, Vocab, WhisperJson, + ByteBpeJson, // Qwen byte-level decoding and Unicode encoding. }; inline constexpr int64_t kWhisperEndOfText = 50257; @@ -48,6 +49,7 @@ class Tokenizer enum class DecodeMode { Pieces, + ByteBpe, WhisperByteBpe, }; @@ -56,6 +58,8 @@ class Tokenizer std::unordered_map bpe_ranks_; std::array byte_encoder_; std::vector special_tokens_; + std::vector added_tokens_; + bool normalize_nfc_ = false; std::vector lang_codes_; DecodeMode decode_mode_ = DecodeMode::Pieces; }; diff --git a/common/io/unicode_regex.cpp b/common/io/unicode_regex.cpp new file mode 100644 index 0000000..1397b18 --- /dev/null +++ b/common/io/unicode_regex.cpp @@ -0,0 +1,66 @@ +// SPDX-License-Identifier: Apache-2.0 +#include "unicode_regex.h" + +#include +#include + +#define PCRE2_CODE_UNIT_WIDTH 8 +#include +#include + +namespace din::io +{ +bool ValidUtf8(std::string_view text) +{ + for (size_t i = 0; i < text.size();) + { + utf8proc_int32_t code; + const auto n = + utf8proc_iterate(reinterpret_cast(text.data() + i), text.size() - i, &code); + if (n < 0) + return false; + i += n; + } + return true; +} + +UnicodeRegex::UnicodeRegex(std::string_view pattern) +{ + int error; + PCRE2_SIZE offset; + code_ = pcre2_compile(reinterpret_cast(pattern.data()), pattern.size(), PCRE2_UTF | PCRE2_UCP, &error, + &offset, nullptr); + if (!code_) + throw std::runtime_error("Invalid Unicode regex at byte " + std::to_string(offset)); +} + +UnicodeRegex::~UnicodeRegex() +{ + pcre2_code_free(code_); +} + +std::vector UnicodeRegex::FindAll(std::string_view text) const +{ + std::vector matches; + const std::unique_ptr data( + pcre2_match_data_create_from_pattern(code_, nullptr), pcre2_match_data_free); + if (!data) + throw std::bad_alloc(); + size_t offset = 0; + while (offset < text.size()) + { + const int result = pcre2_match(code_, reinterpret_cast(text.data()), text.size(), offset, + offset ? PCRE2_NO_UTF_CHECK : 0, data.get(), nullptr); + if (result == PCRE2_ERROR_NOMATCH) + break; + if (result < 0) + throw std::invalid_argument("Invalid UTF-8 text or Unicode regex match failure: " + std::to_string(result)); + const auto* span = pcre2_get_ovector_pointer(data.get()); + if (span[1] <= offset) + throw std::runtime_error("Unicode regex must consume input"); + matches.emplace_back(text.substr(span[0], span[1] - span[0])); + offset = span[1]; + } + return matches; +} +} // namespace din::io diff --git a/common/io/unicode_regex.h b/common/io/unicode_regex.h new file mode 100644 index 0000000..59f6e66 --- /dev/null +++ b/common/io/unicode_regex.h @@ -0,0 +1,27 @@ +// SPDX-License-Identifier: Apache-2.0 +#pragma once + +#include +#include +#include + +struct pcre2_real_code_8; + +namespace din::io +{ +bool ValidUtf8(std::string_view text); + +// Compiled once; matching uses per-call state and validates UTF-8. +class UnicodeRegex +{ +public: + explicit UnicodeRegex(std::string_view pattern); + ~UnicodeRegex(); + UnicodeRegex(const UnicodeRegex&) = delete; + UnicodeRegex& operator=(const UnicodeRegex&) = delete; + std::vector FindAll(std::string_view text) const; + +private: + pcre2_real_code_8* code_; +}; +} // namespace din::io diff --git a/common/ort_session.cpp b/common/ort_session.cpp index 43a3a26..17b73a5 100644 --- a/common/ort_session.cpp +++ b/common/ort_session.cpp @@ -399,51 +399,43 @@ bool IsCudaUnifiedMemoryDevice(Ort::ConstEpDevice ep_device) bool RegisterTensorRTRTXExecutionProvider(Ort::Env& env) { - static std::once_flag registration_once; - static Ort::Env* registered_env = nullptr; - std::call_once(registration_once, - [&env] - { - auto provider_library = std::filesystem::path{ONNXRUNTIME_TRT_RTX_EP_LIBRARY_PATH}; - if (!std::filesystem::is_regular_file(provider_library)) - { + static std::mutex registration_mutex; + const std::lock_guard lock(registration_mutex); + // Check the native environment, not the address of its C++ wrapper. + for (const auto& device : env.GetEpDevices()) + if (std::string_view{device.EpName()} == kDinNvTensorRTRTXExecutionProvider) + return true; + auto provider_library = std::filesystem::path{ONNXRUNTIME_TRT_RTX_EP_LIBRARY_PATH}; + if (!std::filesystem::is_regular_file(provider_library)) + { #ifdef _WIN32 - provider_library = ExecutableDirectory() / "onnxruntime_providers_nv_tensorrt_rtx.dll"; + provider_library = ExecutableDirectory() / "onnxruntime_providers_nv_tensorrt_rtx.dll"; #else - provider_library = ExecutableDirectory() / "libonnxruntime_providers_nv_tensorrt_rtx.so"; + provider_library = ExecutableDirectory() / "libonnxruntime_providers_nv_tensorrt_rtx.so"; #endif - if (!std::filesystem::is_regular_file(provider_library)) - { - throw std::runtime_error("TensorRT RTX execution provider library not found: " + - provider_library.string()); - } - } + if (!std::filesystem::is_regular_file(provider_library)) + { + throw std::runtime_error("TensorRT RTX execution provider library not found: " + provider_library.string()); + } + } #ifdef _WIN32 - const auto provider_directory = provider_library.parent_path().wstring(); - if (SetDllDirectoryW(provider_directory.c_str()) == 0) - { - throw std::runtime_error("Failed to add TensorRT RTX EP directory to the DLL search path."); - } + const auto provider_directory = provider_library.parent_path().wstring(); + if (SetDllDirectoryW(provider_directory.c_str()) == 0) + { + throw std::runtime_error("Failed to add TensorRT RTX EP directory to the DLL search path."); + } #endif - const auto provider_library_path = ToOrtPathString(provider_library); - env.RegisterExecutionProviderLibrary(kDinNvTensorRTRTXExecutionProvider, - provider_library_path.c_str()); - - const auto ep_devices = env.GetEpDevices(); - std::cout << "Execution provider devices after TRT RTX registration:\n"; - for (const auto& device : ep_devices) - { - std::cout << " " << device.EpName() << " vendor=" << device.EpVendor() - << " device_id=" << device.Device().DeviceId() << '\n'; - } - registered_env = &env; - }); + const auto provider_library_path = ToOrtPathString(provider_library); + env.RegisterExecutionProviderLibrary(kDinNvTensorRTRTXExecutionProvider, provider_library_path.c_str()); - if (registered_env != &env) + const auto ep_devices = env.GetEpDevices(); + std::cerr << "Execution provider devices after TRT RTX registration:\n"; + for (const auto& device : ep_devices) { - throw std::logic_error("TensorRT RTX was already registered on a different Ort::Env in this process."); + std::cerr << " " << device.EpName() << " vendor=" << device.EpVendor() + << " device_id=" << device.Device().DeviceId() << '\n'; } return true; } @@ -564,7 +556,8 @@ std::string CompileEpContextModel(Ort::Env& env, const std::string& model_path, const auto output_model_path = CompiledModelPath(model_path, ep_context, profile); if (fs::exists(output_model_path)) { - if (IsCompatibleEpContext(env, output_model_path)) + if (fs::last_write_time(output_model_path) >= fs::last_write_time(model_path) && + IsCompatibleEpContext(env, output_model_path)) { return output_model_path; } @@ -640,7 +633,7 @@ int ChooseCudaDeviceOrdinal(Ort::ConstEpDevice ep_device, const char* override_e const auto ort_luid = luid_text != nullptr ? ParseUint64(*luid_text) : std::optional{}; if (luid_text != nullptr && !ort_luid.has_value()) { - std::cout << "Ignoring unparsable ORT LUID metadata value: " << *luid_text << std::endl; + std::cerr << "Ignoring unparsable ORT LUID metadata value: " << *luid_text << std::endl; } #else const std::string* pci_bus_id_text = FindMetadataValue(metadata, "pci_bus_id"); @@ -648,7 +641,7 @@ int ChooseCudaDeviceOrdinal(Ort::ConstEpDevice ep_device, const char* override_e pci_bus_id_text != nullptr ? ParsePciBusId(*pci_bus_id_text) : std::optional{}; if (pci_bus_id_text != nullptr && !ort_pci_bus_id.has_value()) { - std::cout << "Ignoring unparsable ORT pci_bus_id metadata value: " << *pci_bus_id_text << std::endl; + std::cerr << "Ignoring unparsable ORT pci_bus_id metadata value: " << *pci_bus_id_text << std::endl; } #endif @@ -700,24 +693,24 @@ int ChooseCudaDeviceOrdinal(Ort::ConstEpDevice ep_device, const char* override_e if (metadata_matches.size() > 1) { - std::cout << "Multiple CUDA devices match ORT device metadata"; + std::cerr << "Multiple CUDA devices match ORT device metadata"; } else { - std::cout << "No CUDA device matches ORT device metadata"; + std::cerr << "No CUDA device matches ORT device metadata"; } #ifdef _WIN32 if (luid_text != nullptr) { - std::cout << " LUID=" << *luid_text; + std::cerr << " LUID=" << *luid_text; } #else if (pci_bus_id_text != nullptr) { - std::cout << " pci_bus_id=" << *pci_bus_id_text; + std::cerr << " pci_bus_id=" << *pci_bus_id_text; } #endif - std::cout << " for ORT hardware device_id=" << ort_hardware_device_id << "; defaulting to CUDA ordinal 0. Set " + std::cerr << " for ORT hardware device_id=" << ort_hardware_device_id << "; defaulting to CUDA ordinal 0. Set " << (override_env_var != nullptr ? override_env_var : "FLUX_CUDA_DEVICE_ID") << " to override." << std::endl;