diff --git a/CMakeLists.txt b/CMakeLists.txt index 1098879680..079a5b2944 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) @@ -415,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 35d7a9344c..2df7ad7f64 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" @@ -282,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)}, @@ -419,17 +421,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 - ? static_cast(result.samples.size()) / 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}}; - }); + timing, [&]() { return interface.generate_audio(prompt, config); }, audio_observation); } Json run_speak(trtmc::ITask& task, const Json& request, const Timing& timing) { @@ -456,15 +448,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 - ? static_cast(result.samples.size()) / result.sample_rate - : 0.0}, - {"output_samples", result.samples.size()}, - {"num_samples", result.samples.size()}, - {"sample_rate", result.sample_rate}}; + 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; +} 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 933d5d7114..41dd287b27 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 { @@ -402,6 +406,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.