use tritium_cpu as _;
#[cfg(feature = "cuda")]
use tritium_cuda as _;
use std::path::PathBuf;
use clap::{Parser, Subcommand, ValueEnum};
mod backends;
mod campaign;
#[cfg(feature = "cuda")]
mod campaign_artifact;
#[cfg(feature = "nccl")]
mod campaign_world;
mod generate;
mod inspect;
#[cfg(feature = "cuda")]
mod nvml_probe;
mod pull;
mod quantize;
mod release;
mod repack;
mod report;
mod salt;
const DEFAULT_EOS: u32 = 128_001;
#[derive(Parser, Debug)]
#[command(
name = "tritium",
about = "Inspect ternary GGUF models and list available compute backends.",
version
)]
struct Cli {
#[command(subcommand)]
command: Command,
}
#[derive(Subcommand, Debug)]
enum Command {
Inspect {
path: PathBuf,
},
ListBackends,
Campaign {
#[command(subcommand)]
campaign: campaign::CampaignCommand,
},
Salt {
#[command(subcommand)]
salt: salt::SaltCommand,
},
Pull {
repo: String,
#[arg(long)]
file: Option<String>,
#[arg(long, default_value = "main")]
revision: String,
},
Generate {
#[arg(long)]
model: PathBuf,
#[arg(long)]
tokens: PathBuf,
#[arg(long, default_value_t = 16)]
max_new: usize,
#[arg(long, default_value_t = true, action = clap::ArgAction::Set)]
greedy: bool,
#[arg(long, default_value_t = DEFAULT_EOS)]
eos: u32,
},
Report {
#[command(subcommand)]
report: ReportCommand,
},
Release {
#[command(subcommand)]
release: release::ReleaseCommand,
},
Repack {
#[arg(long)]
input: PathBuf,
#[arg(long)]
output: PathBuf,
#[arg(long, value_enum)]
to: repack::RepackTarget,
},
Quantize {
#[arg(long)]
input: PathBuf,
#[arg(long)]
output: PathBuf,
#[arg(long, default_value_t = 2.0)]
bpw: f64,
#[arg(long, value_enum, default_value_t = quantize::ScaleGroupArg::Block)]
scale_group: quantize::ScaleGroupArg,
#[arg(long, value_enum, default_value_t = SaltSensitivityArg::Uniform)]
sensitivity: SaltSensitivityArg,
#[arg(long)]
fisher: Option<PathBuf>,
#[arg(long, value_enum, default_value_t = quantize::OutputFormat::Sidecar)]
format: quantize::OutputFormat,
},
}
#[derive(Subcommand, Debug)]
enum ReportCommand {
Sparsity {
#[arg(long)]
model: PathBuf,
},
Decode {
#[arg(long)]
model: PathBuf,
#[arg(long)]
tokens: PathBuf,
#[arg(long, default_value = "cpu")]
backend: String,
#[arg(long, default_value_t = 8)]
decode_steps: usize,
#[arg(long, default_value_t = 1)]
warmup: usize,
#[arg(long, value_enum, default_value_t = ReportFormat::Both)]
format: ReportFormat,
},
Ttft {
#[arg(long)]
model: PathBuf,
#[arg(long)]
tokens: PathBuf,
#[arg(long, default_value = "cpu")]
backend: String,
#[arg(long, default_value_t = 1)]
runs: usize,
#[arg(long, value_enum, default_value_t = ReportFormat::Both)]
format: ReportFormat,
},
Compare {
#[arg(long)]
model: PathBuf,
#[arg(long)]
tokens: PathBuf,
#[arg(long, default_value = "cuda")]
backend: String,
#[arg(long, default_value_t = 512)]
prompt_len: usize,
#[arg(long, default_value_t = 256)]
decode_steps: usize,
#[arg(long, default_value_t = 16)]
warmup: usize,
#[arg(long, default_value_t = 3)]
reps: usize,
#[arg(long, default_value_t = 5)]
runs: usize,
#[arg(long, value_enum, default_value_t = ReportFormat::Both)]
format: ReportFormat,
},
Parity {
#[arg(long)]
model: PathBuf,
#[arg(long)]
tokens: PathBuf,
#[arg(long, default_value_t = 16)]
max_new: usize,
#[arg(long, default_value_t = DEFAULT_EOS)]
eos: u32,
#[arg(long, value_enum, default_value_t = ReportFormat::Both)]
format: ReportFormat,
},
Salt {
#[arg(long)]
input: PathBuf,
#[arg(long)]
rows: usize,
#[arg(long)]
k: usize,
#[arg(long)]
budgets: String,
#[arg(long, value_enum, default_value_t = SaltSensitivityArg::Uniform)]
sensitivity: SaltSensitivityArg,
#[arg(long, value_enum, default_value_t = ReportFormat::Both)]
format: ReportFormat,
},
SaltModel {
#[arg(long)]
input: PathBuf,
#[arg(long)]
budgets: String,
#[arg(long, value_enum, default_value_t = SaltSensitivityArg::Uniform)]
sensitivity: SaltSensitivityArg,
#[arg(long, value_enum, default_value_t = quantize::ScaleGroupArg::Block)]
scale_group: quantize::ScaleGroupArg,
#[arg(long, default_value_t = 0)]
limit: usize,
#[arg(long, default_value_t = false)]
per_tensor: bool,
#[arg(long, value_enum, default_value_t = ReportFormat::Both)]
format: ReportFormat,
},
}
#[derive(Clone, Copy, Debug, ValueEnum)]
pub(crate) enum ReportFormat {
Both,
Json,
Table,
}
#[derive(Clone, Copy, Debug, ValueEnum)]
pub(crate) enum SaltSensitivityArg {
Uniform,
Energy,
}
fn main() -> anyhow::Result<()> {
let cli = Cli::parse();
match cli.command {
Command::Inspect { path } => inspect::run(&path)?,
Command::ListBackends => backends::run(),
Command::Campaign { campaign: command } => campaign::run(command)?,
Command::Salt { salt: command } => salt::run(command)?,
Command::Pull {
repo,
file,
revision,
} => pull::run(&repo, file.as_deref(), &revision)?,
Command::Generate {
model,
tokens,
max_new,
greedy,
eos,
} => {
let ids = generate::read_token_file(&tokens)?;
generate::run(&model, &ids, max_new, greedy, eos)?;
}
Command::Repack { input, output, to } => repack::run(&input, &output, to)?,
Command::Release { release: command } => release::run(command)?,
Command::Report { report: command } => match command {
ReportCommand::Sparsity { model } => report::sparsity(&model)?,
ReportCommand::Decode {
model,
tokens,
backend,
decode_steps,
warmup,
format,
} => {
let ids = generate::read_token_file(&tokens)?;
report::decode(&model, &ids, &backend, decode_steps, warmup, format)?;
}
ReportCommand::Compare {
model,
tokens,
backend,
prompt_len,
decode_steps,
warmup,
reps,
runs,
format,
} => {
let ids = generate::read_token_file(&tokens)?;
report::compare(
&model,
&ids,
&backend,
prompt_len,
decode_steps,
warmup,
reps,
runs,
format,
)?
}
ReportCommand::Ttft {
model,
tokens,
backend,
runs,
format,
} => {
let ids = generate::read_token_file(&tokens)?;
report::ttft(&model, &ids, &backend, runs, format)?;
}
ReportCommand::Parity {
model,
tokens,
max_new,
eos,
format,
} => {
let ids = generate::read_token_file(&tokens)?;
report::parity(&model, &ids, max_new, eos, format)?;
}
ReportCommand::Salt {
input,
rows,
k,
budgets,
sensitivity,
format,
} => report::salt(&input, rows, k, &budgets, sensitivity, format)?,
ReportCommand::SaltModel {
input,
budgets,
sensitivity,
scale_group,
limit,
per_tensor,
format,
} => report::salt_model(
&input,
&budgets,
sensitivity,
scale_group,
limit,
per_tensor,
format,
)?,
},
Command::Quantize {
input,
output,
bpw,
scale_group,
sensitivity,
fisher,
format,
} => quantize::run(
&input,
&output,
bpw,
scale_group,
sensitivity,
fisher.as_deref(),
format,
)?,
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn campaign_teacher_cache_cli_parses_fixed_window_inputs() {
let cli = Cli::try_parse_from([
"tritium",
"campaign",
"teacher-cache",
"--model-dir",
"model",
"--corpus",
"tokens.json",
"--seq-len",
"32",
"--output",
"teacher.ttpr",
])
.expect("teacher-cache CLI");
assert!(matches!(
cli.command,
Command::Campaign {
campaign: campaign::CampaignCommand::TeacherCache { seq_len: 32, .. }
}
));
}
#[test]
fn campaign_run_cli_parses_config() {
let cli = Cli::try_parse_from(["tritium", "campaign", "run", "--config", "campaign.json"])
.expect("campaign run CLI");
assert!(matches!(
cli.command,
Command::Campaign {
campaign: campaign::CampaignCommand::Run { config }
} if config == std::path::Path::new("campaign.json")
));
}
#[test]
fn qwen36_preflight_cli_parses_immutable_candidate_output() {
let cli = Cli::try_parse_from([
"tritium",
"salt",
"qwen36-preflight",
"--model-dir",
"model",
"--work-root",
"work",
"--output",
"candidate.json",
])
.expect("Qwen3.6 preflight CLI");
assert!(matches!(
cli.command,
Command::Salt {
salt: salt::SaltCommand::Qwen36Preflight {
model_dir,
work_root,
output,
}
} if model_dir == std::path::Path::new("model")
&& work_root == std::path::Path::new("work")
&& output == std::path::Path::new("candidate.json")
));
}
}