ML in Rust Basics
10 examples to get you started with machine learning in Rust - 7 basic and 3 intermediate.
Search across all documentation pages
10 examples to get you started with machine learning in Rust - 7 basic and 3 intermediate.
cargo new ml-playground && cd ml-playground
cargo add candle-core candle-nn --features mkl
cargo add tokenizers
cargo add ort --features download-binaries
cargo add axum tokio serde --features deriveTooling: Examples target Rust 1.97.0 (edition 2024). Pin ML crate versions in production - APIs evolve quickly.
Run exported models without a Python runtime at inference time.
use ort::session::Session;
fn main() -> ort::Result<()> {
let session = Session::builder()?.commit_from_file("model.onnx")?;
println!("inputs: {:?}", session.inputs);
Ok(())
}ort bundles ONNX Runtime for CPU/GPU execution.Related: ONNX Runtime (ort) - full inference loop
Convert prompts to token IDs with the Hugging Face tokenizers crate.
use tokenizers::Tokenizer;
fn main() -> tokenizers::Result<()> {
let tok = Tokenizer::from_file("tokenizer.json")?;
let enc = tok.encode("Hello, Rust ML!", true)?;
println!("ids: {:?}", enc.get_ids());
Ok(())
}tokenizer.json ships with most HF model repos.add_special_tokens controls BOS/EOS insertion for decoder models.Related: Tokenizers - batching and truncation
Create and multiply tensors on CPU (or GPU with features enabled).
use candle_core::{Device, Tensor};
fn main() -> candle_core::Result<()> {
let device = Device::Cpu;
let a = Tensor::new(&[[1f32, 2.0], [3.0, 4.0]], &device)?;
let b = Tensor::new(&[[2f32, 0.0], [1.0, 2.0]], &device)?;
let c = a.matmul(&b)?;
println!("{}", c);
Ok(())
}Device::Cpu is the portable default; CUDA/Metal need feature flags and drivers.Result - validate dimensions early.Related: candle - nn modules and training
Feed f32 input and read logits.
use ndarray::Array;
use ort::{inputs, session::Session, value::Tensor};
fn main() -> ort::Result<()> {
let session = Session::builder()?.commit_from_file("mnist.onnx")?;
let input = Array::from_shape_vec((1, 1, 28, 28), vec![0.0f32; 784])?;
let outputs = session.run(inputs![Tensor::from_array(input)?])?;
let logits = outputs[0].try_extract_tensor::<f32>()?;
println!("logits len: {}", logits.len());
Ok(())
}nhwd conventions from the export script.1 still matters for graph compatibility.Wrap inference behind HTTP with a minimal router.
use axum::{routing::get, Json, Router};
use serde_json::json;
#[tokio::main]
async fn main() {
let app = Router::new().route("/health", get(|| async { Json(json!({"ok": true})) }));
let listener = tokio::net::TcpListener::bind("127.0.0.1:8080").await.unwrap();
axum::serve(listener, app).await.unwrap();
}/predict with serde-validated JSON bodies next.Related: Serving Models with Axum - production handlers
Typed API contracts prevent shape mistakes at the edge.
use serde::{Deserialize, Serialize};
#[derive(Deserialize)]
struct PredictRequest {
text: String,
}
#[derive(Serialize)]
struct PredictResponse {
label: String,
score: f32,
}Initialize a backend explicitly when experimenting with burn.
use burn::backend::NdArray;
use burn::tensor::Tensor;
type B = NdArray<f32>;
fn main() {
let device = <B as burn::backend::Backend>::Device::default();
let t: Tensor<B, 2> = Tensor::from_data([[1.0, 2.0], [3.0, 4.0]], &device);
let _ = t.matmul(t.transpose());
}Related: burn - training loops and checkpoints
Compute vectors for multiple strings in one forward pass.
// Conceptual layout - pair tokenizers + ort or candle model:
// 1. tokenize with padding/truncation
// 2. build input tensors [batch, seq]
// 3. run session, take pooled output [batch, dim]Related: Embeddings & Vector Search
Load weights once, stream tokens with a Rust LLM crate (llama.cpp bindings or mistral.rs).
// Pseudocode structure for local LLM services:
// let model = LlamaModel::load("weights.gguf", ¶ms)?;
// let mut ctx = model.new_context()?;
// for token in ctx.generate(prompt, max_tokens) { print!("{token}") }max_seq_len per deployment.Related: LLM Inference & Serving
Enable CUDA or Metal only in release builds that target GPU hosts.
[dependencies]
candle-core = { version = "0.8", features = ["cuda"] }gpu vs default CPU builds.Device::cuda_if_available() fails.Related: GPU & Acceleration
Stack versions: This page was written for Rust 1.97.0 (edition 2024), Tokio 1.x, Axum 0.8, serde 1.0, sqlx 0.8, clap 4, and Polars 0.46+.
Reviewed by Chris St. John·Last updated Jul 19, 2026