Skip to content
Draft
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
25 changes: 22 additions & 3 deletions apps/cli/cli.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -91,7 +91,7 @@ const std::unordered_map<std::string, CommandSpec>& command_specs() {
{"generate-audio",
{CommandKind::kGenerateAudio,
{"--prompt", "--output", "--max-new-tokens", "--talker-max-new-tokens", "--seed",
"--stream", "--chunk-frames"}}},
"--stream", "--chunk-frames", "--reference-audio", "--reference-text"}}},
{"transcribe",
{CommandKind::kTranscribe,
{"--input", "--max-output-tokens", "--source-language", "--target-language",
Expand Down Expand Up @@ -983,6 +983,11 @@ int dispatch(const Command& command, ITask& task, std::ostream& output) {
config.seed = int_option(command, "--seed", -1);
const bool streaming = has_option(command, "--stream") &&
parse_bool(command.options.at("--stream"), "--stream");
const bool reference_audio = has_option(command, "--reference-audio");
if (has_option(command, "--reference-text") && !reference_audio)
throw std::invalid_argument("--reference-text requires --reference-audio");
if (streaming && reference_audio)
throw std::invalid_argument("reference-conditioned streaming is not supported");
if (!streaming && has_option(command, "--chunk-frames"))
throw std::invalid_argument("--chunk-frames requires --stream true");
const std::string path = require_option(command, "--output");
Expand Down Expand Up @@ -1019,8 +1024,22 @@ int dispatch(const Command& command, ITask& task, std::ostream& output) {
{"num_samples", total}});
return EXIT_SUCCESS;
}
const auto result = require_interface<IAudioGeneration>(task).generate_audio(
require_option(command, "--prompt"), config);
AudioResult result;
if (reference_audio) {
auto& generator = require_interface<IReferenceAudioGeneration>(task);
auto decoded = read_audio(require_option(command, "--reference-audio"));
require_finite(decoded.samples, "reference audio");
AudioReference reference;
reference.samples = std::move(decoded.samples);
reference.sample_rate = decoded.sample_rate;
if (has_option(command, "--reference-text"))
reference.transcript = command.options.at("--reference-text");
result = generator.generate_audio_with_reference(require_option(command, "--prompt"),
reference, config);
} else {
result = require_interface<IAudioGeneration>(task).generate_audio(
require_option(command, "--prompt"), config);
}
require_finite(result.samples, "generated audio");
io::write_wav(result, path);
write_json(output, {{"output", path},
Expand Down
60 changes: 60 additions & 0 deletions apps/cli/tests/test_cli.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -102,6 +102,30 @@ class EmptyImageWorker final : public trtmc::IImageGeneration {
}
};

class FakeReferenceAudio final : public trtmc::IAudioGeneration,
public trtmc::IReferenceAudioGeneration {
public:
trtmc::AudioReference seen;
trtmc::AudioGenerationConfig config;
std::string prompt;
trtmc::AudioResult generate_audio(const std::string&,
const trtmc::AudioGenerationConfig&) override {
throw std::logic_error("reference request was silently ignored");
}
trtmc::AudioResult
generate_audio_with_reference(const std::string& text, const trtmc::AudioReference& reference,
const trtmc::AudioGenerationConfig& options) override {
seen = reference;
config = options;
prompt = text;
trtmc::AudioResult result;
result.samples = reference.samples;
result.num_samples = static_cast<int32_t>(result.samples.size());
result.sample_rate = reference.sample_rate;
return result;
}
};

class FakeStreamingAudio final : public trtmc::IAudioGeneration,
public trtmc::IStreamingAudioGeneration {
public:
Expand Down Expand Up @@ -594,6 +618,42 @@ int main() {
transcription_audio.sample_rate = 16000;
trtmc::cli::io::write_wav(transcription_audio, transcription_path.string());

auto reference_command =
parse({"trtmc", "generate-audio", "model.bundle", "--runtime-root", "lib", "--prompt",
"new speech", "--reference-audio", transcription_path.string(), "--reference-text",
"reference speech", "--max-new-tokens", "20", "--seed", "7", "--output",
audio_path.string()});
FakeReferenceAudio reference_task;
std::ostringstream reference_output;
check(trtmc::cli::dispatch(reference_command, reference_task, reference_output) == 0,
"reference audio dispatch succeeds");
check(reference_task.seen.sample_rate == 16000 && reference_task.seen.samples.size() == 3 &&
reference_task.seen.transcript == "reference speech" &&
reference_task.prompt == "new speech" && reference_task.config.seed == 7 &&
reference_task.config.max_new_tokens == 20,
"reference request reaches the optional Task API");
std::filesystem::remove(audio_path);
auto dispatch_rejected = [&](const trtmc::cli::Command& command, trtmc::ITask& task) {
try {
std::ostringstream ignored;
trtmc::cli::dispatch(command, task, ignored);
return false;
} catch (const std::exception&) {
return true;
}
};
check(dispatch_rejected(reference_command, audio),
"family without reference capability rejects reference input");
auto invalid_reference = reference_command;
invalid_reference.options.erase("--reference-audio");
check(dispatch_rejected(invalid_reference, reference_task),
"reference transcript without audio is rejected");
invalid_reference = reference_command;
invalid_reference.options["--stream"] = "true";
check(dispatch_rejected(invalid_reference, reference_task),
"reference streaming is explicitly rejected");
check(!std::filesystem::exists(audio_path), "invalid reference requests do not create output");

const auto offline_transcription_command = parse({"trtmc",
"transcribe",
"model.bundle",
Expand Down
3 changes: 3 additions & 0 deletions core/runtime/include/trtmc/runtime/tensor.h
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@ enum class DType {
kBFloat16,
kInt32,
kInt8,
kBool,
};

inline std::size_t dtype_size(DType dt) {
Expand All @@ -36,6 +37,8 @@ inline std::size_t dtype_size(DType dt) {
return 4;
case DType::kInt8:
return 1;
case DType::kBool:
return 1;
}
return 0;
}
Expand Down
19 changes: 19 additions & 0 deletions core/runtime/include/trtmc/task.h
Original file line number Diff line number Diff line change
Expand Up @@ -558,6 +558,25 @@ class IAudioGeneration : public virtual ITask {
const AudioGenerationConfig& config = {}) = 0;
};

// Optional capability for request-time voice conditioning. The caller supplies
// decoded mono PCM; resampling, features and speaker semantics belong to the
// family. Implementations must not retain references to this request after return.
struct AudioReference {
std::vector<float> samples;
std::int32_t sample_rate{0};
// Exact transcript when supplied; empty means no transcript was provided.
std::string transcript;
};

class IReferenceAudioGeneration {
public:
static constexpr const char* kTask = IAudioGeneration::kTask;
virtual ~IReferenceAudioGeneration() = default;
virtual AudioResult generate_audio_with_reference(const std::string& prompt,
const AudioReference& reference,
const AudioGenerationConfig& config = {}) = 0;
};

// Optional task capability for families that produce audio before an utterance
// is complete. It is separate from IAudioGeneration so non-streaming families
// never need an adapter or a fake implementation.
Expand Down
2 changes: 2 additions & 0 deletions core/runtime/tensorrt/trt_module_impl.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,8 @@ DType TrtModuleImpl::from_trt_dtype(nvinfer1::DataType dt) {
return DType::kInt32;
case nvinfer1::DataType::kINT8:
return DType::kInt8;
case nvinfer1::DataType::kBOOL:
return DType::kBool;
default:
return DType::kFloat32;
}
Expand Down
9 changes: 9 additions & 0 deletions families/cosyvoice3/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,9 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Native offline CosyVoice3 with request-time reference-voice conditioning.

Use the standard build API for native bundles. The family module CLI exposes
individual component commands.
Importing this package does not load TensorRT, PyTorch, or checkpoint code.
"""
173 changes: 173 additions & 0 deletions families/cosyvoice3/__main__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,173 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Individual component commands for CosyVoice3 development."""

from __future__ import annotations

import argparse
from dataclasses import asdict
import json
from pathlib import Path

from .artifacts import write_component
from .config import MODEL_ID, MODEL_REVISION, SOURCE_REVISION, ShapeProfile, read_config


def build(args):
from . import trt_compat
from .checkpoint_mapper import load_flow_weights
from .flow_builder import build_flow_engine

output = args.output.resolve()
if output.exists():
raise ValueError(f"Output already exists; choose a new directory: {output}")
cfg = read_config(args.model_dir)
profile = ShapeProfile(args.min_frames, args.opt_frames, args.max_frames)
weights = load_flow_weights(args.model_dir, cfg)
plan = build_flow_engine(weights, cfg, profile, workspace_mib=args.workspace_mib)
metadata = {
"schema_version": 1, "component": "cosyvoice3_flow_estimator",
"status": "component_built",
"target_model_id": MODEL_ID, "target_model_revision": MODEL_REVISION,
"equations_source_revision": SOURCE_REVISION,
"local_checkpoint_revision_verified": False,
"architecture": asdict(cfg), "profile": asdict(profile),
"precision": "fp32", "tf32": False, "streaming": False,
"tensorrt_version": trt_compat.module_version(),
"workspace_mib": args.workspace_mib,
}
write_component(output, "flow.plan", plan, metadata)
print(json.dumps(metadata, indent=2))


def main(argv=None):
parser = argparse.ArgumentParser(description="CosyVoice3 component tools. Individual component builds are not full TTS bundles; use the standard build API for native per-request voice conditioning.")
commands = parser.add_subparsers(dest="command", required=True)
inspect = commands.add_parser("inspect", help="Read the model config safely, without executing YAML constructors")
inspect.add_argument("--model-dir", type=Path, required=True)
build_parser = commands.add_parser("build-flow", help="Build a native FP32 offline DiT component engine")
build_parser.add_argument("--model-dir", type=Path, required=True)
build_parser.add_argument("--output", type=Path, required=True)
build_parser.add_argument("--min-frames", type=int, default=4)
build_parser.add_argument("--opt-frames", type=int, default=64)
build_parser.add_argument("--max-frames", type=int, default=256)
build_parser.add_argument("--workspace-mib", type=int, default=512)
conditioner = commands.add_parser("build-conditioner", help="Build offline token/speaker preprocessing (not full TTS)")
conditioner.add_argument("--model-dir", type=Path, required=True)
conditioner.add_argument("--output", type=Path, required=True)
conditioner.add_argument("--min-tokens", type=int, default=2)
conditioner.add_argument("--opt-tokens", type=int, default=32)
conditioner.add_argument("--max-tokens", type=int, default=128)
conditioner.add_argument("--workspace-mib", type=int, default=64)
for component in ("llm", "hift", "campplus", "speech-tokenizer"):
command = commands.add_parser(f"build-{component}", help=f"Build native offline {component} component")
command.add_argument("--model-dir", type=Path, required=True)
command.add_argument("--output", type=Path, required=True)
command.add_argument("--workspace-mib", type=int, default=256)
if component == "llm":
command.add_argument("--max-context", type=int, default=1024)
command.add_argument("--opt-tokens", type=int, default=64)
else:
command.add_argument("--min-frames", type=int, default=4)
command.add_argument("--opt-frames", type=int, default=64 if component == "hift" else 500)
command.add_argument("--max-frames", type=int, default=256 if component == "hift" else 3000)
voice = commands.add_parser("prepare-voice", help="Reference WAV -> native TensorRT frontend -> voice NPZ")
for name in ("audio", "campplus", "speech-tokenizer", "output"):
voice.add_argument(f"--{name}", type=Path, required=True)
tts = commands.add_parser("synthesize", help="Offline text + prepared voice NPZ -> WAV (experimental, unqualified)")
for name in ("model-dir", "llm", "conditioner", "flow", "hift", "voice", "output"):
tts.add_argument(f"--{name}", type=Path, required=True)
tts.add_argument("--text", required=True)
tts.add_argument("--instruction", default="You are a helpful assistant.")
tts.add_argument("--prompt-text", default="", help="Exact reference transcript for zero-shot; omit for instruction mode")
tts.add_argument("--max-tokens", type=int, default=100)
tts.add_argument("--seed", type=int, default=2512)
tts.add_argument("--greedy", action="store_true", help="Deterministic argmax, not the official default RAS sampler")
args = parser.parse_args(argv)
if args.command == "inspect":
print(json.dumps({"target": MODEL_ID, "flow": asdict(read_config(args.model_dir)), "status": "component_only"}, indent=2))
elif args.command == "build-flow":
build(args)
elif args.command == "build-conditioner":
build_conditioner(args)
elif args.command == "synthesize":
from .tts import synthesize

synthesize(args)
elif args.command == "prepare-voice":
from .tts import prepare_voice

prepare_voice(args)
else:
build_speech_component(args)


def build_conditioner(args):
from . import trt_compat
from .conditioning import TokenProfile, build_engine, load_weights

output = args.output.resolve()
if output.exists():
raise FileExistsError(output)
read_config(args.model_dir)
profile = TokenProfile(args.min_tokens, args.opt_tokens, args.max_tokens)
plan = build_engine(load_weights(args.model_dir), profile, workspace_mib=args.workspace_mib)
manifest = {
"schema_version": 1, "component": "cosyvoice3_conditioner",
"status": "component_built",
"target_model_id": MODEL_ID, "target_model_revision": MODEL_REVISION,
"equations_source_revision": SOURCE_REVISION,
"local_checkpoint_revision_verified": False,
"precision": "fp32", "tf32": False, "streaming": False,
"profile": asdict(profile),
"tensorrt_version": trt_compat.module_version(), "workspace_mib": args.workspace_mib,
}
write_component(output, "conditioning.plan", plan, manifest)
print(json.dumps(manifest, indent=2))


def build_speech_component(args):
from . import trt_compat

component = args.command.removeprefix("build-").replace("-", "_")
frontend = component in ("campplus", "speech_tokenizer")
output = args.output.resolve()
if output.exists():
raise FileExistsError(output)
metadata = {}
if component == "llm":
from .llm import build_engine, load_weights

cfg, weights = load_weights(args.model_dir)
plan = build_engine(weights, cfg, max_context=args.max_context, opt_tokens=args.opt_tokens,
workspace_mib=args.workspace_mib)
metadata.update(architecture=asdict(cfg), max_context=args.max_context, opt_tokens=args.opt_tokens,
compact_kv_cache=True, checkpoint="llm.pt")
elif frontend:
from .frontend import FrontendProfile, build_engine

profile = FrontendProfile(args.min_frames, args.opt_frames, args.max_frames)
plan = build_engine(args.model_dir, component, profile, workspace_mib=args.workspace_mib)
metadata.update(profile=asdict(profile), batch_size=1, padded_input=False,
learned_execution="native_tensorrt", checkpoint_structure_validated=True)
else:
from .hift import build_engine, load_weights

profile = ShapeProfile(args.min_frames, args.opt_frames, args.max_frames)
plan = build_engine(load_weights(args.model_dir), profile, workspace_mib=args.workspace_mib)
metadata.update(profile=asdict(profile), sample_rate=24000, samples_per_frame=480,
f0_precision="fp32_reference_uses_fp64", finalize=True, noise="explicit_uniform_0_1",
phase_accumulation="chronological_fp32_recurrence")
metadata.update(schema_version=1, component=f"cosyvoice3_{component}",
status="component_built",
target_model_id=MODEL_ID, target_model_revision=MODEL_REVISION,
equations_source_revision=SOURCE_REVISION, local_checkpoint_revision_verified=False,
precision="fp32", tf32=False, streaming=False,
tensorrt_version=trt_compat.module_version(), workspace_mib=args.workspace_mib)
write_component(output, f"{component}.plan", plan, metadata)
print(json.dumps(metadata, indent=2))


if __name__ == "__main__":
main()
26 changes: 26 additions & 0 deletions families/cosyvoice3/artifacts.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Family-local atomic component publication."""

import json
import os
from pathlib import Path
import tempfile

def write_component(output, plan_name, plan, manifest):
"""Publish a complete component without updating an existing directory."""
output = Path(output)
if output.exists():
raise FileExistsError(output)
if Path(plan_name).name != plan_name or not plan_name.endswith(".plan"):
raise ValueError("plan_name must be a .plan filename")
output.parent.mkdir(parents=True, exist_ok=True)
with tempfile.TemporaryDirectory(prefix=".cosyvoice3-", dir=output.parent) as tmp:
stage = Path(tmp) / "component"
stage.mkdir()
(stage / plan_name).write_bytes(plan)
(stage / "manifest.json").write_text(json.dumps(manifest, indent=2) + "\n", encoding="utf-8")
if output.exists():
raise FileExistsError(output)
os.rename(stage, output)
Loading
Loading