mod chart;
mod dataset;
use std::time::Instant;
use topos::{
Conv2d, Linear, Module, Parameters, Plan, Shape, Symbol, Tape, Tensor, Value, cross_entropy,
init, max_pool,
};
use chart::loss_chart;
use dataset::{Split, load, shuffle};
const IMAGE_SIDE: usize = 28;
const PIXELS: usize = IMAGE_SIDE * IMAGE_SIDE;
const CLASSES: usize = 10;
const BATCH_LEN: usize = 64;
const PROBE_LEN: usize = 1000;
const FILTERS_1: usize = 8;
const FILTERS_2: usize = 16;
const FLAT_LEN: usize = FILTERS_2 * (IMAGE_SIDE / 4) * (IMAGE_SIDE / 4);
const STEPS: usize = 2000;
struct Model {
conv_1: Conv2d<f32>,
conv_2: Conv2d<f32>,
head: Linear<f32>,
}
impl Model {
fn new(tape: &Tape<f32>) -> Self {
let conv_1_weights =
init::normal(11, (2.0 / 9.0_f64).sqrt())(&Shape::new([FILTERS_1, 1, 3, 3]));
let conv_2_weights =
init::normal(12, (2.0 / 72.0_f64).sqrt())(&Shape::new([FILTERS_2, FILTERS_1, 3, 3]));
let mut head_weights = init::kaiming(13);
Self {
conv_1: Conv2d::new(tape, conv_1_weights, Tensor::filled([FILTERS_1], 0.0), 1, 1),
conv_2: Conv2d::new(tape, conv_2_weights, Tensor::filled([FILTERS_2], 0.0), 1, 1),
head: Linear::new(
tape,
head_weights(&Shape::new([FLAT_LEN, CLASSES])),
head_weights(&Shape::new([CLASSES])),
),
}
}
fn express<'tape>(&self, images: Value<'tape, f32>, rows: usize) -> Value<'tape, f32> {
let stage_1 = max_pool(self.conv_1.express(images).relu(), 2, 2);
let stage_2 = max_pool(self.conv_2.express(stage_1).relu(), 2, 2);
self.head.express(stage_2.reshape([rows, FLAT_LEN]))
}
}
fn batch_payloads(split: &Split, indices: &[usize]) -> (Tensor<f32>, Tensor<f32>) {
let mut pixels = Vec::with_capacity(indices.len() * PIXELS);
for &index in indices {
pixels.extend_from_slice(&split.pixels[index * PIXELS..(index + 1) * PIXELS]);
}
let labels: Vec<usize> = indices.iter().map(|&index| split.labels[index]).collect();
(
Tensor::new([indices.len(), 1, IMAGE_SIDE, IMAGE_SIDE], pixels),
Tensor::selection(labels, CLASSES, 1.0),
)
}
fn probe_correct(
parameters: &Parameters<f32>,
probe_plan: &Plan<f32>,
images_symbol: Symbol,
logits_symbol: Symbol,
test: &Split,
start: usize,
) -> usize {
let indices: Vec<usize> = (start..start + PROBE_LEN).collect();
let (images, _) = batch_payloads(test, &indices);
let run = probe_plan.forward(parameters, [(images_symbol, images)]);
let logits = run.of(logits_symbol).to_vec();
let mut correct = 0;
for (row, &index) in logits.chunks(CLASSES).zip(&indices) {
let predicted = row
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).expect("logits are finite"))
.map(|(class, _)| class)
.expect("a logit row is never empty");
if predicted == test.labels[index] {
correct += 1;
}
}
correct
}
fn main() {
let (train, test) = load();
println!(
"loaded {} training and {} test images",
train.len(),
test.len()
);
let tape = Tape::new();
let model = Model::new(&tape);
let images = tape.input(Tensor::filled(
[BATCH_LEN, 1, IMAGE_SIDE, IMAGE_SIDE],
0.0_f32,
));
let targets = tape.input(Tensor::selection(vec![0; BATCH_LEN], CLASSES, 1.0));
let loss = cross_entropy(model.express(images, BATCH_LEN), targets);
let probe_images = tape.input(Tensor::filled(
[PROBE_LEN, 1, IMAGE_SIDE, IMAGE_SIDE],
0.0_f32,
));
let probe_logits = model.express(probe_images, PROBE_LEN);
let (images, targets, loss, probe_images, probe_logits) = (
images.symbol(),
targets.symbol(),
loss.symbol(),
probe_images.symbol(),
probe_logits.symbol(),
);
let recorded_nodes = tape.len();
println!("recorded {recorded_nodes} nodes for both expressions");
let network = tape.into_network();
let mut parameters = network.parameters();
let training_plan = network.entry([loss]).backward().lower();
let probe_plan = network.entry([probe_logits]).lower();
for line in training_plan.describe().lines().filter(|line| {
line.starts_with("plan:") || line.starts_with("live volume:") || line.starts_with("fused")
}) {
println!("training {line}");
}
for line in probe_plan.describe().lines().filter(|line| {
line.starts_with("plan:") || line.starts_with("live volume:") || line.starts_with("fused")
}) {
println!("probe {line}");
}
let mut order: Vec<usize> = (0..train.len()).collect();
let mut shuffle_state: u64 = 5;
shuffle(&mut order, &mut shuffle_state);
let fast = Tensor::new([], [0.1_f32]);
let slow = Tensor::new([], [0.01_f32]);
let mut losses = Vec::new();
let training = Instant::now();
for step in 0..STEPS {
let start = (step * BATCH_LEN) % (train.len() - BATCH_LEN);
let batch = &order[start..start + BATCH_LEN];
let (batch_images, batch_targets) = batch_payloads(&train, batch);
let run = training_plan.forward(
¶meters,
[(images, batch_images), (targets, batch_targets)],
);
let batch_loss = run.of(loss).scalar();
losses.push(batch_loss);
if step == 0 {
println!(
"step 0: minibatch loss = {batch_loss:.4} (a uniform model costs ln 10 ~ 2.30)"
);
}
let gradients = run.backward(loss).parameters(¶meters);
let learning_rate = if step < STEPS * 3 / 4 { &fast } else { &slow };
parameters = parameters.step(&gradients, |parameter, gradient| {
parameter.clone() - gradient.clone() * learning_rate.broadcast_like(gradient)
});
if (step + 1) % 250 == 0 {
let correct = probe_correct(
¶meters,
&probe_plan,
probe_images,
probe_logits,
&test,
0,
);
println!(
"step {:4}: minibatch loss = {batch_loss:.4}, probe accuracy = {:.1}%",
step + 1,
correct as f64 * 100.0 / PROBE_LEN as f64
);
}
}
let elapsed = training.elapsed().as_secs_f64();
println!(
"trained {STEPS} steps in {elapsed:.3}s ({:.1} ms/step)",
elapsed * 1000.0 / STEPS as f64
);
let correct: usize = (0..test.len() / PROBE_LEN)
.map(|chunk| {
probe_correct(
¶meters,
&probe_plan,
probe_images,
probe_logits,
&test,
chunk * PROBE_LEN,
)
})
.sum();
println!(
"test accuracy: {:.2}% over {} images",
correct as f64 * 100.0 / test.len() as f64,
test.len()
);
assert_eq!(network.len(), recorded_nodes);
println!("the tape held {recorded_nodes} nodes through every step");
println!("{}", loss_chart("mnist convnet training", &losses));
}