Skip to content
Merged
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
7 changes: 3 additions & 4 deletions AGENTS.md
Original file line number Diff line number Diff line change
Expand Up @@ -50,10 +50,9 @@ Python/PyTorch is conversion- and development-time only; normal inference is Rus

## Verification

CI installs the current stable Rust toolchain on every run. Before handing off changes, update the
local stable toolchain (`rustup update stable`), confirm `rustc --version` matches current stable,
and run the exact CI sequence below. Do not substitute `cargo check` for Clippy, filter the training
tests, or omit `cargo build`:
CI installs the current stable Rust toolchain on every run. Before handing off changes, run the
exact CI sequence below. Do not substitute `cargo check` for Clippy, filter the training tests, or
omit `cargo build`:

```console
cargo fmt --check
Expand Down
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -105,7 +105,7 @@ the training-set size (capped at 1024), and uses 80% of that verified maximum fo
## Export ONNX

```console
montgomery export-onnx --model yolo26n.bpk
montgomery export --model yolo26n.bpk --format onnx
```

This reads the explicit Burnpack and writes `yolo26n.onnx`; use `--output` to select another path.
Expand Down
33 changes: 25 additions & 8 deletions src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,13 @@ enum DeviceSelection {
Gpu,
}

#[cfg(feature = "onnx")]
#[derive(Debug, Clone, Copy, PartialEq, Eq, ValueEnum)]
enum ExportFormat {
/// Open Neural Network Exchange format.
Onnx,
}

#[derive(Debug, Parser)]
#[command(
name = "montgomery",
Expand All @@ -60,9 +67,9 @@ enum Command {
Bench(BenchArgs),
/// Pack an imported tensor-only state into a versioned native Burnpack artifact.
PackWeights(PackWeightsArgs),
/// Export the exact loaded Burn model weights to a validated portable ONNX artifact.
/// Export a Burnpack model to another artifact format.
#[cfg(feature = "onnx")]
ExportOnnx(ExportOnnxArgs),
Export(ExportArgs),
/// Train a model with the native Burn/WGPU trainer.
#[cfg(feature = "training")]
Train(TrainArgs),
Expand All @@ -71,7 +78,7 @@ enum Command {
Val(ValArgs),
/// Export a native training checkpoint to the existing inference Burnpack format.
#[cfg(feature = "training")]
Export(ExportTrainingArgs),
ExportCheckpoint(ExportCheckpointArgs),
/// Internal isolated worker used by automatic batch-size discovery.
#[cfg(feature = "training")]
#[command(name = "__batch-probe", hide = true)]
Expand Down Expand Up @@ -196,7 +203,7 @@ struct ValArgs {

#[cfg(feature = "training")]
#[derive(Debug, ClapArgs)]
struct ExportTrainingArgs {
struct ExportCheckpointArgs {
#[arg(long)]
checkpoint: PathBuf,
#[arg(long)]
Expand All @@ -213,10 +220,13 @@ struct BatchProbeArgs {

#[cfg(feature = "onnx")]
#[derive(Debug, ClapArgs)]
struct ExportOnnxArgs {
struct ExportArgs {
/// Montgomery .bpk model to export; architecture is read from artifact metadata.
#[arg(long, value_name = "MODEL.bpk")]
model: PathBuf,
/// Artifact format to export.
#[arg(long, value_enum)]
format: ExportFormat,
/// Final ONNX path (defaults to <model>.onnx). A missing suffix is added explicitly.
#[arg(long)]
output: Option<PathBuf>,
Expand Down Expand Up @@ -597,7 +607,7 @@ fn main() -> montgomery::Result<()> {
Ok(())
}
#[cfg(feature = "onnx")]
Command::ExportOnnx(args) => export_onnx_command(args),
Command::Export(args) => export_command(args),
#[cfg(feature = "training")]
Command::Train(args) => {
let initialization = match (args.architecture, args.model, args.resume) {
Expand Down Expand Up @@ -697,7 +707,7 @@ fn main() -> montgomery::Result<()> {
Ok(())
}
#[cfg(feature = "training")]
Command::Export(args) => {
Command::ExportCheckpoint(args) => {
let output = export_training(args.checkpoint, args.output)?;
eprintln!("Exported inference artifact to {}", output.display());
Ok(())
Expand All @@ -706,7 +716,14 @@ fn main() -> montgomery::Result<()> {
}

#[cfg(feature = "onnx")]
fn export_onnx_command(args: ExportOnnxArgs) -> montgomery::Result<()> {
fn export_command(args: ExportArgs) -> montgomery::Result<()> {
match args.format {
ExportFormat::Onnx => export_onnx_command(args),
}
}

#[cfg(feature = "onnx")]
fn export_onnx_command(args: ExportArgs) -> montgomery::Result<()> {
let weights = args.model;
let model = ModelId::from_burnpack(&weights)?;
let output = args
Expand Down
34 changes: 33 additions & 1 deletion tests/integration.rs
Original file line number Diff line number Diff line change
Expand Up @@ -361,10 +361,17 @@ fn cli_help_exposes_the_supported_workflows_and_coordinate_contract() {
let output = montgomery(&["--help"]);
assert!(output.status.success(), "{}", stderr(&output));
let help = String::from_utf8_lossy(&output.stdout);
for command in ["predict", "bench", "pack-weights", "export-onnx", "train"] {
for command in ["predict", "bench", "pack-weights", "export", "train"] {
assert!(help.contains(command), "missing {command} in:\n{help}");
}

let output = montgomery(&["export", "--help"]);
assert!(output.status.success(), "{}", stderr(&output));
let help = String::from_utf8_lossy(&output.stdout);
assert!(help.contains("--model <MODEL.bpk>"), "{help}");
assert!(help.contains("--format <FORMAT>"), "{help}");
assert!(help.contains("onnx"), "{help}");

let output = montgomery(&["predict", "--help"]);
assert!(output.status.success(), "{}", stderr(&output));
let help = String::from_utf8_lossy(&output.stdout);
Expand All @@ -380,6 +387,25 @@ fn cli_help_exposes_the_supported_workflows_and_coordinate_contract() {
assert!(help.contains("--state <STATE.pt>"), "{help}");
}

#[test]
fn cli_export_requires_a_supported_format_before_loading_the_model() {
let output = montgomery(&["export", "--model", "missing.bpk"]);
assert!(!output.status.success());
assert!(stderr(&output).contains("--format <FORMAT>"));

let output = montgomery(&[
"export",
"--model",
"missing.bpk",
"--format",
"unsupported",
]);
assert!(!output.status.success());
let error = stderr(&output);
assert!(error.contains("invalid value 'unsupported'"), "{error}");
assert!(error.contains("onnx"), "{error}");
}

#[cfg(feature = "training")]
#[test]
fn training_cli_requires_one_explicit_initialization_mode() {
Expand Down Expand Up @@ -409,6 +435,12 @@ fn training_cli_requires_one_explicit_initialization_mode() {
let output = montgomery(&["train", "--model", "upstream.pt", "--data", "missing.yaml"]);
assert!(!output.status.success());
assert!(stderr(&output).contains("native .bpk artifact"));

let output = montgomery(&["export-checkpoint", "--help"]);
assert!(output.status.success(), "{}", stderr(&output));
let help = String::from_utf8_lossy(&output.stdout);
assert!(help.contains("--checkpoint <CHECKPOINT>"), "{help}");
assert!(help.contains("--output <OUTPUT>"), "{help}");
}

#[test]
Expand Down
2 changes: 1 addition & 1 deletion tools/onnx/README.md
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
# ONNX export environment

The Rust `export-onnx` command loads the checkpoint and writes the exact loaded parameters to a
The Rust `export --format onnx` command loads the checkpoint and writes the exact loaded parameters to a
private SafeTensors snapshot. These scripts reconstruct the graph from the pinned sibling source,
export it, run ONNX checker and strict shape inference, execute it with ONNX Runtime CPU, compare
the named outputs, and write the sidecar. They do not download weights, install packages, or use a
Expand Down
Loading