Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 20 additions & 0 deletions CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down
27 changes: 27 additions & 0 deletions apps/benchmark/native/audio_observation.h
Original file line number Diff line number Diff line change
@@ -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 <nlohmann/json.hpp>
#include <stdexcept>

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<std::size_t>(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<double>(count) / result.num_channels / result.sample_rate},
{"sample_rate", result.sample_rate},
{"num_channels", result.num_channels}};
}
28 changes: 7 additions & 21 deletions apps/benchmark/native/benchmark_worker.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
* SPDX-License-Identifier: Apache-2.0
*/

#include "audio_observation.h"
#include "trtmc/runtime/family_loader.h"
#include "trtmc/task.h"

Expand Down Expand Up @@ -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)},
Expand Down Expand Up @@ -419,17 +421,7 @@ Json run_generate_audio(trtmc::ITask& task, const Json& request, const Timing& t
config.seed = optional_value<std::int32_t>(request, "seed", -1);
const std::string prompt = request.at("prompt").get<std::string>();
return measure(
timing, [&]() { return interface.generate_audio(prompt, config); },
[](const trtmc::AudioResult& result) {
const double seconds =
result.sample_rate > 0
? static_cast<double>(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) {
Expand All @@ -456,15 +448,9 @@ Json run_speak(trtmc::ITask& task, const Json& request, const Timing& timing) {
static_cast<double>(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<double>(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;
});
}

Expand Down
50 changes: 50 additions & 0 deletions apps/benchmark/tests/native/test_audio_observation.cpp
Original file line number Diff line number Diff line change
@@ -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 <iostream>

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<float>(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;
}
39 changes: 29 additions & 10 deletions apps/cli/io.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
#include <cstddef>
#include <cstdint>
#include <fstream>
#include <limits>
#include <stdexcept>
#include <string>
#include <vector>
Expand All @@ -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<std::int32_t>(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<std::size_t>(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<std::size_t>(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<std::uint64_t>(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<std::int16_t>::max() ||
bytes_per_second > std::numeric_limits<std::int32_t>::max() ||
audio.samples.size() >
(static_cast<std::size_t>(std::numeric_limits<std::int32_t>::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<std::int16_t>(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<std::int16_t>(num_channels * (bits_per_sample / 8));
const std::int32_t data_size = num_samples * block_align;
const auto byte_rate = static_cast<std::int32_t>(bytes_per_second);
const auto block_align = static_cast<std::int16_t>(block_bytes);
const auto data_size = static_cast<std::int32_t>(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<const char*>(&chunk_size), 4);
output.write("WAVEfmt ", 8);
Expand All @@ -53,6 +69,9 @@ void write_wav(const AudioResult& audio, const std::string& path) {
output.write("data", 4);
output.write(reinterpret_cast<const char*>(&data_size), 4);
output.write(reinterpret_cast<const char*>(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) {
Expand Down
2 changes: 2 additions & 0 deletions apps/cli/io.h
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand Down
114 changes: 114 additions & 0 deletions apps/cli/tests/test_audio_io.cpp
Original file line number Diff line number Diff line change
@@ -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 <cstdint>
#include <cstring>
#include <filesystem>
#include <fstream>
#include <iostream>
#include <iterator>
#include <limits>
#include <stdexcept>
#include <vector>

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<unsigned char>& 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<std::uint32_t>(bytes.at(offset + i)) << (8U * i);
return value;
}

std::vector<unsigned char> read_bytes(const std::filesystem::path& path) {
std::ifstream input(path, std::ios::binary);
return {std::istreambuf_iterator<char>(input), std::istreambuf_iterator<char>()};
}
} // 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<double>(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<float>({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<std::int32_t>::max(), 2});
rejects({std::vector<float>(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;
}
5 changes: 5 additions & 0 deletions core/runtime/include/trtmc/task.h
Original file line number Diff line number Diff line change
Expand Up @@ -104,9 +104,13 @@ struct ImageResult {
};

struct AudioResult {
// Interleaved PCM: frame 0 channel 0, frame 0 channel 1, ..., frame 1 channel 0.
std::vector<float> 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 {
Expand Down Expand Up @@ -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<void(const float* samples, std::int32_t num_samples, std::int32_t sample_rate)>;

Expand Down
Loading
Loading