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
62 changes: 43 additions & 19 deletions apps/cli/cli.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -988,34 +988,58 @@ int dispatch(const Command& command, ITask& task, std::ostream& output) {
const std::string path = require_option(command, "--output");
if (streaming) {
const int32_t chunk_frames = int_option(command, "--chunk-frames", 32, 1);
auto* multichannel = dynamic_cast<IMultichannelStreamingAudioGeneration*>(&task);
auto* mono = dynamic_cast<IStreamingAudioGeneration*>(&task);
if (multichannel == nullptr && mono == nullptr)
throw std::invalid_argument("task does not support streaming audio generation");
std::ofstream stream(path, std::ios::binary | std::ios::trunc);
if (!stream)
throw std::runtime_error("unable to open streaming audio output: " + path);
int32_t sample_rate = 0;
const int32_t total =
require_interface<IStreamingAudioGeneration>(task).generate_audio_streaming(
require_option(command, "--prompt"), config,
[&](const float* samples, int32_t num_samples, int32_t rate) {
if (samples == nullptr || num_samples <= 0 || rate <= 0)
throw std::runtime_error(
"streaming audio family returned an invalid chunk");
if (sample_rate != 0 && sample_rate != rate)
throw std::runtime_error(
"streaming audio sample rate changed mid-stream");
sample_rate = rate;
stream.write(reinterpret_cast<const char*>(samples),
static_cast<std::streamsize>(num_samples) * sizeof(float));
if (!stream)
throw std::runtime_error("failed to write streaming audio output: " +
path);
},
chunk_frames);
int32_t num_channels = 0;
std::int64_t observed_samples = 0;
const auto write_chunk = [&](const AudioChunkView& chunk) {
if (chunk.samples == nullptr || chunk.num_samples <= 0 || chunk.sample_rate <= 0 ||
chunk.num_channels <= 0 || chunk.num_samples % chunk.num_channels != 0)
throw std::runtime_error("streaming audio family returned an invalid chunk");
if ((sample_rate != 0 && sample_rate != chunk.sample_rate) ||
(num_channels != 0 && num_channels != chunk.num_channels))
throw std::runtime_error("streaming audio format changed mid-stream");
for (int32_t index = 0; index < chunk.num_samples; ++index)
if (!std::isfinite(chunk.samples[index]))
throw std::runtime_error("streaming audio contains non-finite samples");
if (observed_samples > std::numeric_limits<std::int64_t>::max() - chunk.num_samples)
throw std::runtime_error("streaming audio sample count overflow");
sample_rate = chunk.sample_rate;
num_channels = chunk.num_channels;
stream.write(reinterpret_cast<const char*>(chunk.samples),
static_cast<std::streamsize>(chunk.num_samples) * sizeof(float));
if (!stream)
throw std::runtime_error("failed to write streaming audio output: " + path);
observed_samples += chunk.num_samples;
};
const auto& prompt = require_option(command, "--prompt");
const std::int64_t total =
multichannel != nullptr
? multichannel->generate_audio_streaming(prompt, config, write_chunk,
chunk_frames)
: mono->generate_audio_streaming(
prompt, config,
[&](const float* samples, int32_t count, int32_t rate) {
write_chunk({samples, count, rate, 1});
},
chunk_frames);
stream.close();
if (total <= 0 || sample_rate <= 0)
if (!stream)
throw std::runtime_error("failed to close streaming audio output: " + path);
if (total <= 0 || observed_samples == 0)
throw std::runtime_error("streaming audio family produced no samples");
if (total != observed_samples)
throw std::runtime_error("streaming audio sample count does not match chunks");
write_json(output, {{"output", path},
{"format", "float32le"},
{"sample_rate", sample_rate},
{"num_channels", num_channels},
{"num_samples", total}});
return EXIT_SUCCESS;
}
Expand Down
93 changes: 93 additions & 0 deletions apps/cli/tests/test_cli.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
#include <filesystem>
#include <fstream>
#include <iostream>
#include <limits>
#include <memory>
#include <sstream>
#include <stdexcept>
Expand Down Expand Up @@ -125,6 +126,56 @@ class FakeStreamingAudio final : public trtmc::IAudioGeneration,
}
};

class FakeMultichannelAudio final : public trtmc::IAudioGeneration,
public trtmc::IStreamingAudioGeneration,
public trtmc::IMultichannelStreamingAudioGeneration {
public:
std::string fault;
trtmc::AudioGenerationConfig seen;
std::int32_t seen_chunk_frames{0};

trtmc::AudioResult generate_audio(const std::string&,
const trtmc::AudioGenerationConfig&) override {
throw std::logic_error("synchronous audio path was selected");
}
std::int32_t generate_audio_streaming(const std::string&, const trtmc::AudioGenerationConfig&,
trtmc::AudioChunkCallback, std::int32_t) override {
throw std::logic_error("mono capability was selected instead of multichannel");
}
std::int64_t generate_audio_streaming(const std::string&,
const trtmc::AudioGenerationConfig& config,
trtmc::MultichannelAudioChunkCallback callback,
std::int32_t chunk_frames) override {
seen = config;
seen_chunk_frames = chunk_frames;
if (fault == "empty")
return 0;
float first[] = {0.25F, -0.5F, 0.75F, -1.0F};
trtmc::AudioChunkView chunk{first, 4, 48000, 2};
if (fault == "null")
chunk.samples = nullptr;
if (fault == "partial frame")
chunk.num_samples = 3;
if (fault == "empty chunk")
chunk.num_samples = 0;
if (fault == "negative count")
chunk.num_samples = -2;
if (fault == "invalid rate")
chunk.sample_rate = 0;
if (fault == "invalid channels")
chunk.num_channels = 0;
if (fault == "nonfinite")
first[0] = std::numeric_limits<float>::infinity();
callback(chunk);
if (fault == "producer error")
throw std::runtime_error("synthetic producer failure");
const float last[] = {0.125F, -0.25F};
callback(
{last, 2, fault == "rate change" ? 24000 : 48000, fault == "channel change" ? 1 : 2});
return fault == "wrong total" ? 3 : 6;
}
};

class FakeTranscriptionStream final : public trtmc::ITranscriptionStream {
public:
explicit FakeTranscriptionStream(trtmc::TranscriptionStreamConfig config)
Expand Down Expand Up @@ -587,6 +638,48 @@ int main() {
"streaming audio output format is explicit");
std::filesystem::remove(audio_path);

check(audio_output.str().find("\"num_channels\":1") != std::string::npos,
"legacy streaming reports one channel");
FakeMultichannelAudio stereo_audio;
audio_output.str("");
audio_output.clear();
check(trtmc::cli::dispatch(audio_command, stereo_audio, audio_output) == 0,
"multichannel capability takes precedence over mono");
check(stereo_audio.seen.max_new_tokens == 9 && stereo_audio.seen.seed == 17 &&
stereo_audio.seen_chunk_frames == 4,
"multichannel options reach the Task API");
check(audio_output.str().find("\"num_channels\":2") != std::string::npos &&
audio_output.str().find("\"sample_rate\":48000") != std::string::npos &&
audio_output.str().find("\"num_samples\":6") != std::string::npos,
"streaming metadata describes interleaved stereo");
std::vector<float> streamed_samples(6);
{
std::ifstream input(audio_path, std::ios::binary);
input.read(reinterpret_cast<char*>(streamed_samples.data()), 6 * sizeof(float));
check(static_cast<bool>(input) && input.peek() == std::char_traits<char>::eof(),
"stream contains exactly the delivered samples");
}
check(streamed_samples == std::vector<float>({0.25F, -0.5F, 0.75F, -1.0F, 0.125F, -0.25F}),
"chunk concatenation preserves left and right sample order");

for (const std::string fault :
{"empty", "null", "partial frame", "empty chunk", "negative count", "invalid rate",
"invalid channels", "nonfinite", "producer error", "rate change", "channel change",
"wrong total"}) {
stereo_audio.fault = fault;
audio_output.str("");
audio_output.clear();
bool rejected = false;
try {
trtmc::cli::dispatch(audio_command, stereo_audio, audio_output);
} catch (const std::runtime_error&) {
rejected = true;
}
check(rejected, ("reject streaming fault: " + fault).c_str());
check(audio_output.str().empty(), "failed stream does not report successful output");
}
std::filesystem::remove(audio_path);

const std::filesystem::path transcription_path = "/tmp/trtmc-cli-transcription-stream.wav";
trtmc::AudioResult transcription_audio;
transcription_audio.samples = {0.25F, -0.5F, 0.75F};
Expand Down
26 changes: 26 additions & 0 deletions core/runtime/include/trtmc/task.h
Original file line number Diff line number Diff line change
Expand Up @@ -404,6 +404,17 @@ struct AudioGenerationConfig {
using AudioChunkCallback =
std::function<void(const float* samples, std::int32_t num_samples, std::int32_t sample_rate)>;

// Borrowed interleaved float PCM, valid only during the callback. num_samples
// counts scalar samples (not frames); each nonempty chunk contains whole frames.
struct AudioChunkView {
const float* samples{nullptr};
std::int32_t num_samples{0};
std::int32_t sample_rate{0};
std::int32_t num_channels{1};
};

using MultichannelAudioChunkCallback = std::function<void(const AudioChunkView& chunk)>;

struct SpeechToSpeechConfig {
std::int32_t max_new_tokens{128};
std::int32_t seed{-1};
Expand Down Expand Up @@ -571,6 +582,21 @@ class IStreamingAudioGeneration {
std::int32_t chunk_frames) = 0;
};

// Optional channel-aware capability; existing mono implementations need not
// change. Callbacks are synchronous, ordered, and non-concurrent. Rate/channel
// count stay fixed for a call. Return signals completion and reports the total
// scalar samples delivered. Callback exceptions must stop generation and escape
// the call; implementations must not retain the callback after returning.
class IMultichannelStreamingAudioGeneration {
public:
static constexpr const char* kTask = IAudioGeneration::kTask;
virtual ~IMultichannelStreamingAudioGeneration() = default;
virtual std::int64_t generate_audio_streaming(const std::string& prompt,
const AudioGenerationConfig& config,
MultichannelAudioChunkCallback callback,
std::int32_t chunk_frames) = 0;
};

class ITranscription : public virtual ITask {
public:
static constexpr const char* kTask = "transcription";
Expand Down
1 change: 1 addition & 0 deletions core/runtime/tests/test_task_api.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@ static_assert(std::is_abstract_v<trtmc::IImageBatchGeneration>);
static_assert(std::is_abstract_v<trtmc::IWorldModelGeneration>);
static_assert(std::is_abstract_v<trtmc::IAudioGeneration>);
static_assert(std::is_abstract_v<trtmc::IStreamingAudioGeneration>);
static_assert(std::is_abstract_v<trtmc::IMultichannelStreamingAudioGeneration>);
static_assert(std::is_abstract_v<trtmc::ITranscription>);
static_assert(std::is_abstract_v<trtmc::IBatchTranscription>);
static_assert(std::is_abstract_v<trtmc::IStreamingTranscription>);
Expand Down
26 changes: 26 additions & 0 deletions website/docs/architecture/runtime-lifecycle.md
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,32 @@ 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.

## Multichannel streaming audio

Families can opt into `IMultichannelStreamingAudioGeneration` without changing
the existing mono `IStreamingAudioGeneration` interface. A family implementing
both is dispatched through the multichannel capability by the CLI.

Each `AudioChunkView` borrows interleaved float PCM (`L0, R0, L1, R1, ...` for
stereo), with an explicit channel count and sample rate. `num_samples` counts
scalar samples, not frames per channel. Chunks must be nonempty whole frames;
the sample rate and channel count stay constant within a call. Callbacks are
synchronous, ordered, and non-concurrent. Their pointers are valid only during
the callback. Normal return ends the stream and reports the sum of delivered
scalar samples; callback exceptions must stop generation and propagate.

`trtmc generate-audio ... --stream true --output audio.raw` writes interleaved
float32 samples and reports `format`, `sample_rate`, `num_channels`, and
`num_samples` in its success JSON. Playback duration is
`num_samples / num_channels / sample_rate`. This is raw PCM, not a WAV file.
Invalid chunks, format changes, inconsistent totals, and file-write errors fail
the command without success JSON. A failed stream can leave a partial output
file; callers must not treat file existence alone as success.

This capability does not add HTTP transport, encoded formats, or streaming
support to models that do not already produce incremental audio. Existing mono
families remain unchanged and report `num_channels: 1`.

## Optional load settings

Runtime-sized KV capacity is passed directly to compatible families.
Expand Down
Loading