mod dataset;
use std::time::Instant;
use topos::{
Conv2d, Linear, Module, Shape, Symbol, Tape, Tensor, Value, cross_entropy, init, max_pool,
};
use dataset::{Split, load, shuffle};
const IMAGE_SIDE: usize = 32;
const PIXELS: usize = 3 * IMAGE_SIDE * IMAGE_SIDE;
const CLASSES: usize = 10;
const BATCH_LEN: usize = 64;
const FILTERS: [usize; 3] = [16, 32, 64];
const FLAT_LEN: usize = FILTERS[2] * (IMAGE_SIDE / 8) * (IMAGE_SIDE / 8);
struct Model {
conv_1: Conv2d<f32>,
conv_2: Conv2d<f32>,
conv_3: Conv2d<f32>,
head: Linear<f32>,
}
impl Model {
fn new(tape: &Tape<f32>) -> Self {
let conv_1_weights =
init::normal(21, (2.0 / 27.0_f64).sqrt())(&Shape::new([FILTERS[0], 3, 3, 3]));
let conv_2_weights =
init::normal(22, (2.0 / 144.0_f64).sqrt())(&Shape::new([FILTERS[1], FILTERS[0], 3, 3]));
let conv_3_weights =
init::normal(23, (2.0 / 288.0_f64).sqrt())(&Shape::new([FILTERS[2], FILTERS[1], 3, 3]));
let mut head_weights = init::kaiming(24);
Self {
conv_1: Conv2d::new(
tape,
conv_1_weights,
Tensor::filled([FILTERS[0]], 0.0),
1,
1,
),
conv_2: Conv2d::new(
tape,
conv_2_weights,
Tensor::filled([FILTERS[1]], 0.0),
1,
1,
),
conv_3: Conv2d::new(
tape,
conv_3_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);
let stage_3 = max_pool(self.conv_3.express(stage_2).relu(), 2, 2);
self.head.express(stage_3.reshape([rows, FLAT_LEN]))
}
fn parameters(&self) -> [Symbol; 8] {
[
self.conv_1.weights(),
self.conv_1.bias(),
self.conv_2.weights(),
self.conv_2.bias(),
self.conv_3.weights(),
self.conv_3.bias(),
self.head.weights(),
self.head.bias(),
]
}
}
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(), 3, IMAGE_SIDE, IMAGE_SIDE], pixels),
Tensor::selection(labels, CLASSES, 1.0),
)
}
fn main() {
let route = std::env::var("TOPOS_ROUTE").unwrap_or_else(|_| "engine".to_string());
let steps: usize = std::env::var("TOPOS_STEPS")
.ok()
.and_then(|steps| steps.parse().ok())
.unwrap_or(300);
let (train, _test) = load();
println!("route {route}: {steps} steps over {} images", train.len());
let recorded = match route.as_str() {
"engine" => false,
"recorded" => true,
other => panic!("unknown TOPOS_ROUTE {other:?}; use engine or recorded"),
};
let tape = Tape::new();
let model = Model::new(&tape);
let images = tape.input(Tensor::filled(
[BATCH_LEN, 3, 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 (images, targets, loss) = (images.symbol(), targets.symbol(), loss.symbol());
let parameter_symbols = model.parameters();
let forward_nodes = tape.len();
let adjoints = recorded.then(|| {
let adjoints = tape.differentiate(loss, parameter_symbols);
println!(
"recorded the chain rule: {forward_nodes} forward nodes + {} gradient nodes",
tape.len() - forward_nodes
);
adjoints
});
let network = tape.into_network();
let mut parameters = network.parameters();
let plan = match &adjoints {
Some(adjoints) => network.entry(adjoints.roots()).lower(),
None => network.entry([loss]).backward().lower(),
};
for line in plan
.describe()
.lines()
.filter(|line| line.starts_with("plan:") || line.starts_with("live volume:"))
{
println!("{line}");
}
let mut order: Vec<usize> = (0..train.len()).collect();
let mut shuffle_state: u64 = 3;
shuffle(&mut order, &mut shuffle_state);
let fast = Tensor::new([], [0.05_f32]);
let slow = Tensor::new([], [0.005_f32]);
let mut first_loss = 0.0_f32;
let mut last_loss = 0.0_f32;
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 = plan.forward(
¶meters,
[(images, batch_images), (targets, batch_targets)],
);
let batch_loss = run.of(loss).scalar();
if step == 0 {
first_loss = batch_loss;
}
last_loss = batch_loss;
let gradients = match &adjoints {
Some(adjoints) => run.recorded_gradients(adjoints),
None => 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)
});
}
let elapsed = training.elapsed().as_secs_f64();
println!(
"loss step 0: {first_loss:.6} ({:08x}), step {}: {last_loss:.6} ({:08x})",
first_loss.to_bits(),
steps - 1,
last_loss.to_bits(),
);
println!(
"route {route}: {steps} steps in {elapsed:.3}s ({:.1} ms/step)",
elapsed * 1000.0 / steps as f64
);
}