From 62937fe30748c68fcd75d906b5bc62a404d0b1de Mon Sep 17 00:00:00 2001 From: Ruilin Gao Date: Thu, 10 Sep 2026 19:20:18 -0400 Subject: [PATCH 1/3] feat(audio): preserve channels in WAV output Add an interleaved channel contract to AudioResult and preserve it in the CLI WAV writer. Account for channels in benchmark duration, validate output buffers before opening files, and cover mono compatibility and stereo output with CPU regression tests. AudioChunkCallback and input downmix remain mono. Clients and family DSOs must be rebuilt together because the public result layout changes. Refs: #1254 Signed-off-by: Ruilin Gao --- CMakeLists.txt | 10 ++ apps/benchmark/native/benchmark_worker.cpp | 22 ++-- apps/cli/io.cpp | 39 ++++-- apps/cli/io.h | 2 + apps/cli/tests/test_audio_io.cpp | 114 ++++++++++++++++++ core/runtime/include/trtmc/task.h | 5 + .../docs/architecture/runtime-lifecycle.md | 19 +++ 7 files changed, 191 insertions(+), 20 deletions(-) create mode 100644 apps/cli/tests/test_audio_io.cpp diff --git a/CMakeLists.txt b/CMakeLists.txt index 2d653e3281..9caa417294 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -348,6 +348,16 @@ if(TRTMC_BUILD_TESTS) target_compile_options(test_cli PRIVATE -Wall -Wextra -Wpedantic) add_test(NAME cli COMMAND test_cli) + add_executable(test_audio_io apps/cli/tests/test_audio_io.cpp apps/cli/io.cpp) + target_include_directories(test_audio_io PRIVATE + ${PROJECT_SOURCE_DIR}/core/runtime/include + ${PROJECT_SOURCE_DIR}/apps + ${PROJECT_SOURCE_DIR}/third_party/stb + ) + target_compile_options(test_audio_io PRIVATE -Wall -Wextra -Wpedantic) + add_test(NAME audio_io COMMAND test_audio_io) + set_tests_properties(audio_io PROPERTIES LABELS "cpu") + set(_trtmc_test_runtime_root "${CMAKE_BINARY_DIR}/tests/runtime") add_library(trtmc_test_backend_fake SHARED core/runtime/tests/fake_backend.cpp) target_include_directories(trtmc_test_backend_fake PRIVATE ${PROJECT_SOURCE_DIR}/core/runtime/include) diff --git a/apps/benchmark/native/benchmark_worker.cpp b/apps/benchmark/native/benchmark_worker.cpp index 0abbf93e7a..fb16679681 100644 --- a/apps/benchmark/native/benchmark_worker.cpp +++ b/apps/benchmark/native/benchmark_worker.cpp @@ -421,14 +421,15 @@ Json run_generate_audio(trtmc::ITask& task, const Json& request, const Timing& t return measure( timing, [&]() { return interface.generate_audio(prompt, config); }, [](const trtmc::AudioResult& result) { - const double seconds = - result.sample_rate > 0 - ? static_cast(result.samples.size()) / result.sample_rate - : 0.0; + const double seconds = result.sample_rate > 0 && result.num_channels > 0 + ? static_cast(result.samples.size()) / + result.num_channels / result.sample_rate + : 0.0; return Json{{"output_samples", result.samples.size()}, {"num_samples", result.samples.size()}, {"output_audio_seconds", seconds}, - {"sample_rate", result.sample_rate}}; + {"sample_rate", result.sample_rate}, + {"num_channels", result.num_channels}}; }); } @@ -458,13 +459,14 @@ Json run_speak(trtmc::ITask& task, const Json& request, const Timing& timing) { [](const auto& value) { const auto& result = value.first; return Json{{"input_audio_seconds", value.second}, - {"output_audio_seconds", - result.sample_rate > 0 - ? static_cast(result.samples.size()) / result.sample_rate - : 0.0}, + {"output_audio_seconds", result.sample_rate > 0 && result.num_channels > 0 + ? static_cast(result.samples.size()) / + result.num_channels / result.sample_rate + : 0.0}, {"output_samples", result.samples.size()}, {"num_samples", result.samples.size()}, - {"sample_rate", result.sample_rate}}; + {"sample_rate", result.sample_rate}, + {"num_channels", result.num_channels}}; }); } diff --git a/apps/cli/io.cpp b/apps/cli/io.cpp index 472c86e721..04e5822c9e 100644 --- a/apps/cli/io.cpp +++ b/apps/cli/io.cpp @@ -15,6 +15,7 @@ #include #include #include +#include #include #include #include @@ -24,22 +25,37 @@ namespace trtmc::cli::io { void write_wav(const AudioResult& audio, const std::string& path) { if (audio.samples.empty()) throw std::runtime_error("write_wav: empty audio"); - - std::ofstream output(path, std::ios::binary); - if (!output) - throw std::runtime_error("write_wav: cannot open " + path); - - const auto num_samples = static_cast(audio.samples.size()); + if (audio.sample_rate <= 0 || audio.num_channels <= 0) + throw std::runtime_error("write_wav: sample rate and channel count must be positive"); + if (audio.samples.size() % static_cast(audio.num_channels) != 0) + throw std::runtime_error("write_wav: incomplete interleaved audio frame"); + if (audio.num_samples < 0 || + (audio.num_samples != 0 && + static_cast(audio.num_samples) != audio.samples.size())) + throw std::runtime_error("write_wav: sample count does not match buffer"); + + const std::uint64_t block_bytes = static_cast(audio.num_channels) * 4U; + const std::uint64_t bytes_per_second = block_bytes * audio.sample_rate; + // Keep RIFF sizes within the signed range used by the existing WAV reader. + if (block_bytes > std::numeric_limits::max() || + bytes_per_second > std::numeric_limits::max() || + audio.samples.size() > + (static_cast(std::numeric_limits::max()) - 36U) / 4U) + throw std::runtime_error("write_wav: audio exceeds supported WAV sizes"); const std::int32_t sample_rate = audio.sample_rate; - const std::int16_t num_channels = 1; + const auto num_channels = static_cast(audio.num_channels); const std::int16_t bits_per_sample = 32; - const std::int32_t byte_rate = sample_rate * num_channels * (bits_per_sample / 8); - const auto block_align = static_cast(num_channels * (bits_per_sample / 8)); - const std::int32_t data_size = num_samples * block_align; + const auto byte_rate = static_cast(bytes_per_second); + const auto block_align = static_cast(block_bytes); + const auto data_size = static_cast(audio.samples.size() * 4U); const std::int32_t chunk_size = 36 + data_size; const std::int32_t format_size = 16; const std::int16_t audio_format = 3; + std::ofstream output(path, std::ios::binary); + if (!output) + throw std::runtime_error("write_wav: cannot open " + path); + output.write("RIFF", 4); output.write(reinterpret_cast(&chunk_size), 4); output.write("WAVEfmt ", 8); @@ -53,6 +69,9 @@ void write_wav(const AudioResult& audio, const std::string& path) { output.write("data", 4); output.write(reinterpret_cast(&data_size), 4); output.write(reinterpret_cast(audio.samples.data()), data_size); + output.close(); + if (!output) + throw std::runtime_error("write_wav: failed to write " + path); } AudioResult read_wav(const std::string& path) { diff --git a/apps/cli/io.h b/apps/cli/io.h index cdce384e41..3fcfccaa61 100644 --- a/apps/cli/io.h +++ b/apps/cli/io.h @@ -21,7 +21,9 @@ struct LoadedImage { bool empty() const { return pixels.empty(); } }; +// Audio input remains mono: multichannel files are averaged across channels. AudioResult read_wav(const std::string& path); +// Preserve AudioResult's interleaved channels in a float32 WAV file. void write_wav(const AudioResult& audio, const std::string& path); LoadedImage read_image(const std::string& path); diff --git a/apps/cli/tests/test_audio_io.cpp b/apps/cli/tests/test_audio_io.cpp new file mode 100644 index 0000000000..9b9ba50df0 --- /dev/null +++ b/apps/cli/tests/test_audio_io.cpp @@ -0,0 +1,114 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "cli/io.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace { +int failures = 0; + +void check(bool value, const char* name) { + if (!value) { + std::cerr << "FAIL: " << name << '\n'; + ++failures; + } +} + +std::uint32_t little_endian(const std::vector& bytes, std::size_t offset, + std::size_t length) { + std::uint32_t value = 0; + for (std::size_t i = 0; i < length; ++i) + value |= static_cast(bytes.at(offset + i)) << (8U * i); + return value; +} + +std::vector read_bytes(const std::filesystem::path& path) { + std::ifstream input(path, std::ios::binary); + return {std::istreambuf_iterator(input), std::istreambuf_iterator()}; +} +} // namespace + +int main() { + const auto path = std::filesystem::temp_directory_path() / "trtmc-multichannel-audio.wav"; + try { + // Existing three-field construction must continue to mean mono. + const trtmc::AudioResult mono{{-1.0F, 0.25F, 1.0F}, 3, 16000}; + check(mono.num_channels == 1, "legacy aggregate defaults to mono"); + trtmc::cli::io::write_wav(mono, path.string()); + const auto mono_bytes = read_bytes(path); + check(little_endian(mono_bytes, 22, 2) == 1 && little_endian(mono_bytes, 28, 4) == 64000 && + little_endian(mono_bytes, 32, 2) == 4, + "mono WAV header remains unchanged"); + const auto restored = trtmc::cli::io::read_wav(path.string()); + check(restored.samples == mono.samples && restored.num_samples == 3 && + restored.sample_rate == 16000 && restored.num_channels == 1, + "mono round trip"); + + // Three distinct left/right frames; comparing bytes avoids a reader + // accidentally hiding a writer error by downmixing the output. + trtmc::AudioResult stereo{{1.0F, -1.0F, 0.5F, 0.0F, -0.5F, 1.0F}, 6, 48000, 2}; + trtmc::cli::io::write_wav(stereo, path.string()); + const auto bytes = read_bytes(path); + check(bytes.size() == 44 + 6 * sizeof(float), "stereo payload is not counted twice"); + check(little_endian(bytes, 4, 4) == bytes.size() - 8 && little_endian(bytes, 20, 2) == 3 && + little_endian(bytes, 22, 2) == 2 && little_endian(bytes, 24, 4) == 48000 && + little_endian(bytes, 28, 4) == 384000 && little_endian(bytes, 32, 2) == 8 && + little_endian(bytes, 34, 2) == 32 && little_endian(bytes, 40, 4) == 24, + "stereo float32 WAV header"); + check(static_cast(little_endian(bytes, 40, 4)) / little_endian(bytes, 28, 4) == + 3.0 / 48000.0, + "duration counts frames, not scalar samples"); + for (std::size_t i = 0; i < stereo.samples.size(); ++i) { + const auto bits = little_endian(bytes, 44 + 4 * i, 4); + float sample = 0.0F; + std::memcpy(&sample, &bits, sizeof(sample)); + check(sample == stereo.samples[i], "interleaved channel samples round trip"); + } + const auto downmixed = trtmc::cli::io::read_wav(path.string()); + check(downmixed.num_channels == 1 && downmixed.num_samples == 3 && + downmixed.samples == std::vector({0.0F, 0.25F, 0.25F}), + "existing input downmix is preserved"); + + stereo.num_samples = 0; + trtmc::cli::io::write_wav(stereo, path.string()); + check(read_bytes(path) == bytes, "unspecified sample count uses buffer size"); + + const auto rejects = [&](trtmc::AudioResult invalid) { + bool rejected = false; + try { + trtmc::cli::io::write_wav(invalid, path.string()); + } catch (const std::runtime_error&) { + rejected = true; + } + check(rejected, "invalid audio is rejected"); + check(read_bytes(path) == bytes, "validation does not truncate an existing file"); + }; + rejects({}); + rejects({{1.0F, 2.0F}, 2, 48000, 0}); + rejects({{1.0F, 2.0F}, 2, 48000, -1}); + rejects({{1.0F, 2.0F}, 2, 0, 2}); + rejects({{1.0F, 2.0F}, 2, -1, 2}); + rejects({{1.0F, 2.0F, 3.0F}, 3, 48000, 2}); + rejects({{1.0F, 2.0F}, 1, 48000, 2}); + rejects({{1.0F, 2.0F}, -1, 48000, 2}); + rejects({{1.0F, 2.0F}, 2, std::numeric_limits::max(), 2}); + rejects({std::vector(8192), 8192, 48000, 8192}); + } catch (const std::exception& error) { + std::cerr << "FAIL: " << error.what() << '\n'; + ++failures; + } + std::filesystem::remove(path); + std::cerr << (failures == 0 ? "ALL PASSED\n" : "SOME FAILED\n"); + return failures == 0 ? 0 : 1; +} diff --git a/core/runtime/include/trtmc/task.h b/core/runtime/include/trtmc/task.h index 087c9f02c5..a11e86699b 100644 --- a/core/runtime/include/trtmc/task.h +++ b/core/runtime/include/trtmc/task.h @@ -104,9 +104,13 @@ struct ImageResult { }; struct AudioResult { + // Interleaved PCM: frame 0 channel 0, frame 0 channel 1, ..., frame 1 channel 0. std::vector samples; + // Total scalar samples, not frames per channel. Zero may mean unspecified. std::int32_t num_samples{0}; std::int32_t sample_rate{24000}; + // Appended to preserve existing three-field aggregate initialization. + std::int32_t num_channels{1}; }; struct TranscriptionStreamConfig { @@ -401,6 +405,7 @@ struct AudioGenerationConfig { std::int32_t seed{-1}; }; +// Legacy mono streaming callback; multichannel output is represented by AudioResult. using AudioChunkCallback = std::function; diff --git a/website/docs/architecture/runtime-lifecycle.md b/website/docs/architecture/runtime-lifecycle.md index b82b68eb40..61c1a3f75c 100644 --- a/website/docs/architecture/runtime-lifecycle.md +++ b/website/docs/architecture/runtime-lifecycle.md @@ -33,6 +33,25 @@ contract. The backend owns TensorRT runtime objects, not model policy. returning or copy the lightweight `BundleReader` into the pipeline for deferred reads; it must not retain a reference to the temporary factory context. +## Audio output contract + +`AudioResult` carries interleaved float PCM. `num_channels` defaults to one; +for stereo the buffer order is `L0, R0, L1, R1, ...`. `num_samples` counts total +scalar samples, not frames per channel; zero leaves the count unspecified and +consumers use `samples.size()`. A nonzero count must match the buffer size. +The sample rate is in frames per second, so duration is +`samples.size() / num_channels / sample_rate`. Buffers must contain whole frames. + +The CLI WAV writer preserves the channel count and writes float32 PCM. Its WAV +reader intentionally still averages input channels to mono for existing speech +tasks. `AudioChunkCallback` also remains mono; this result contract does not add +multichannel streaming support. + +The appended field preserves existing three-field aggregate initialization and +mono behavior at source level. It changes the C++ struct layout: rebuild clients +and family DSOs together rather than mixing binaries compiled against different +versions of `task.h`. + ## Optional load settings Runtime-sized KV capacity is passed directly to compatible families. From e55be2c03395e4ba8ad72a1bec7837dcb6d354d7 Mon Sep 17 00:00:00 2001 From: Ruilin Gao <229880965+ruiling-smartbear@users.noreply.github.com> Date: Mon, 14 Sep 2026 16:38:45 -0400 Subject: [PATCH 2/3] fix(benchmark): validate audio sample counts Signed-off-by: Ruilin Gao <229880965+ruiling-smartbear@users.noreply.github.com> --- CMakeLists.txt | 10 ++++ apps/benchmark/native/audio_observation.h | 27 ++++++++++ apps/benchmark/native/benchmark_worker.cpp | 27 ++-------- .../tests/native/test_audio_observation.cpp | 50 +++++++++++++++++++ 4 files changed, 92 insertions(+), 22 deletions(-) create mode 100644 apps/benchmark/native/audio_observation.h create mode 100644 apps/benchmark/tests/native/test_audio_observation.cpp diff --git a/CMakeLists.txt b/CMakeLists.txt index 9caa417294..22535f7e3d 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -425,6 +425,16 @@ if(TRTMC_BUILD_TESTS) set_tests_properties(trt_module_dynamic_input PROPERTIES SKIP_RETURN_CODE 77 LABELS gpu) if(TRTMC_BUILD_EXAMPLES) + add_executable(test_audio_observation apps/benchmark/tests/native/test_audio_observation.cpp) + target_include_directories(test_audio_observation PRIVATE + ${PROJECT_SOURCE_DIR}/apps/benchmark/native + ${PROJECT_SOURCE_DIR}/core/runtime/include + ) + target_link_libraries(test_audio_observation PRIVATE nlohmann_json::nlohmann_json) + target_compile_options(test_audio_observation PRIVATE -Wall -Wextra -Wpedantic) + add_test(NAME audio_observation COMMAND test_audio_observation) + set_tests_properties(audio_observation PROPERTIES LABELS "cpu") + add_executable(test_dataset_answer apps/benchmark/tests/native/test_dataset_answer.cpp) target_include_directories(test_dataset_answer PRIVATE ${PROJECT_SOURCE_DIR}/apps/benchmark/native diff --git a/apps/benchmark/native/audio_observation.h b/apps/benchmark/native/audio_observation.h new file mode 100644 index 0000000000..49bf4913d6 --- /dev/null +++ b/apps/benchmark/native/audio_observation.h @@ -0,0 +1,27 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#pragma once + +#include "trtmc/task.h" + +#include +#include + +inline nlohmann::json audio_observation(const trtmc::AudioResult& result) { + const auto count = result.samples.size(); + if (result.num_samples < 0 || + (result.num_samples != 0 && static_cast(result.num_samples) != count)) + throw std::runtime_error("audio benchmark: sample count does not match buffer"); + if (result.sample_rate <= 0 || result.num_channels <= 0) + throw std::runtime_error("audio benchmark: sample rate and channel count must be positive"); + if (count % result.num_channels != 0) + throw std::runtime_error("audio benchmark: incomplete interleaved audio frame"); + return {{"output_samples", count}, + {"num_samples", count}, + {"output_audio_seconds", + static_cast(count) / result.num_channels / result.sample_rate}, + {"sample_rate", result.sample_rate}, + {"num_channels", result.num_channels}}; +} diff --git a/apps/benchmark/native/benchmark_worker.cpp b/apps/benchmark/native/benchmark_worker.cpp index fb16679681..a0bde2e93a 100644 --- a/apps/benchmark/native/benchmark_worker.cpp +++ b/apps/benchmark/native/benchmark_worker.cpp @@ -3,6 +3,7 @@ * SPDX-License-Identifier: Apache-2.0 */ +#include "audio_observation.h" #include "trtmc/runtime/family_loader.h" #include "trtmc/task.h" @@ -419,18 +420,7 @@ Json run_generate_audio(trtmc::ITask& task, const Json& request, const Timing& t config.seed = optional_value(request, "seed", -1); const std::string prompt = request.at("prompt").get(); return measure( - timing, [&]() { return interface.generate_audio(prompt, config); }, - [](const trtmc::AudioResult& result) { - const double seconds = result.sample_rate > 0 && result.num_channels > 0 - ? static_cast(result.samples.size()) / - result.num_channels / result.sample_rate - : 0.0; - return Json{{"output_samples", result.samples.size()}, - {"num_samples", result.samples.size()}, - {"output_audio_seconds", seconds}, - {"sample_rate", result.sample_rate}, - {"num_channels", result.num_channels}}; - }); + timing, [&]() { return interface.generate_audio(prompt, config); }, audio_observation); } Json run_speak(trtmc::ITask& task, const Json& request, const Timing& timing) { @@ -457,16 +447,9 @@ Json run_speak(trtmc::ITask& task, const Json& request, const Timing& timing) { static_cast(audio.samples.size()) / audio.sample_rate}; }, [](const auto& value) { - const auto& result = value.first; - return Json{{"input_audio_seconds", value.second}, - {"output_audio_seconds", result.sample_rate > 0 && result.num_channels > 0 - ? static_cast(result.samples.size()) / - result.num_channels / result.sample_rate - : 0.0}, - {"output_samples", result.samples.size()}, - {"num_samples", result.samples.size()}, - {"sample_rate", result.sample_rate}, - {"num_channels", result.num_channels}}; + auto summary = audio_observation(value.first); + summary["input_audio_seconds"] = value.second; + return summary; }); } diff --git a/apps/benchmark/tests/native/test_audio_observation.cpp b/apps/benchmark/tests/native/test_audio_observation.cpp new file mode 100644 index 0000000000..662c91ebc7 --- /dev/null +++ b/apps/benchmark/tests/native/test_audio_observation.cpp @@ -0,0 +1,50 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#include "audio_observation.h" + +#include + +int main() { + int failures = 0; + const auto check = [&](bool ok, const char* name) { + if (!ok) { + std::cerr << "FAIL: " << name << '\n'; + ++failures; + } + }; + for (const int channels : {1, 2}) { + trtmc::AudioResult audio{std::vector(48000 * channels), 0, 48000, channels}; + for (const int count : {0, 48000 * channels}) { + audio.num_samples = count; + const auto summary = audio_observation(audio); + check(summary.at("num_samples") == audio.samples.size(), "resolved sample count"); + check(summary.at("output_samples") == audio.samples.size(), "output sample count"); + check(summary.at("output_audio_seconds") == 1.0, "channel-aware duration"); + check(summary.at("num_channels") == channels, "channel metadata"); + } + for (const int count : {-1, 1, 48000 * channels + 1}) { + audio.num_samples = count; + bool rejected = false; + try { + audio_observation(audio); + } catch (const std::runtime_error&) { + rejected = true; + } + check(rejected, "reject inconsistent sample count"); + } + } + for (const auto& audio : + {trtmc::AudioResult{{0.0F}, 0, 48000, 2}, trtmc::AudioResult{{0.0F}, 0, 0, 1}, + trtmc::AudioResult{{0.0F}, 0, 48000, 0}}) { + bool rejected = false; + try { + audio_observation(audio); + } catch (const std::runtime_error&) { + rejected = true; + } + check(rejected, "reject invalid frame or metadata"); + } + return failures ? 1 : 0; +} From 1145906ed178e15882c94c4026862106ef7755dd Mon Sep 17 00:00:00 2001 From: Ruilin Gao Date: Tue, 15 Sep 2026 01:15:46 -0400 Subject: [PATCH 3/3] fix(benchmark): exclude observations from timing Record task-call latency before metadata validation and summary construction so candidate observations do not inflate the comparison timing. Signed-off-by: Ruilin Gao --- apps/benchmark/native/benchmark_worker.cpp | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/apps/benchmark/native/benchmark_worker.cpp b/apps/benchmark/native/benchmark_worker.cpp index 38c543e8e2..2df7ad7f64 100644 --- a/apps/benchmark/native/benchmark_worker.cpp +++ b/apps/benchmark/native/benchmark_worker.cpp @@ -283,8 +283,9 @@ Json measure(const Timing& timing, Invoke&& invoke, Observe&& observe) { for (int index = 0; index < timing.iterations; ++index) { const auto started = Clock::now(); last = invoke(); + const double runtime_e2e_wall_ms = elapsed_ms(started); Json observation = observe(*last); - observation["runtime_e2e_wall_ms"] = elapsed_ms(started); + observation["runtime_e2e_wall_ms"] = runtime_e2e_wall_ms; observations.push_back(std::move(observation)); } return {{"observations", std::move(observations)},