mod chart;
mod corpus;
use std::time::Instant;
use malevich::{Cells, Frame, Plot, Scale};
use topos::{Network, Shape, Tensor, Tensorial, cross_entropy, init};
use chart::loss_chart;
use corpus::{VOCABULARY_LEN, draw, from_token, load_names, shuffle, training_samples};
const BATCH_LEN: usize = 1024;
fn main() {
let names = load_names();
let mut samples = training_samples::<1>(&names);
let mut shuffle_state: u64 = 9;
shuffle(&mut samples, &mut shuffle_state);
println!(
"loaded {} names, {} bigram pairs",
names.len(),
samples.len()
);
let network: Network<Tensor<f32>> = Network::new();
let table = network.parameter(init::normal(7, 0.01)(&Shape::new([
VOCABULARY_LEN,
VOCABULARY_LEN,
])));
let contexts = network.input(Tensor::selection(vec![0; BATCH_LEN], VOCABULARY_LEN, 1.0));
let targets = network.input(Tensor::selection(vec![0; BATCH_LEN], VOCABULARY_LEN, 1.0));
let logits = table.gather(contexts);
let loss = cross_entropy(logits, targets);
let table_symbol = table.symbol();
let contexts_symbol = contexts.symbol();
let targets_symbol = targets.symbol();
let loss_symbol = loss.symbol();
let recorded_nodes = network.len();
let learning_rate = Tensor::new([], [10.0]);
let mut network = network;
let mut losses = Vec::new();
let training = Instant::now();
for step in 0..1000 {
let start = (step * BATCH_LEN) % (samples.len() - BATCH_LEN);
let batch = &samples[start..start + BATCH_LEN];
let batch_contexts: Vec<usize> = batch.iter().map(|&(context, _)| context[0]).collect();
let batch_targets: Vec<usize> = batch.iter().map(|&(_, next)| next).collect();
let loss_value = network.resolve(loss_symbol);
let run = network.forward_for(
[loss_symbol],
[
(
contexts_symbol,
Tensor::selection(batch_contexts, VOCABULARY_LEN, 1.0),
),
(
targets_symbol,
Tensor::selection(batch_targets, VOCABULARY_LEN, 1.0),
),
],
);
let batch_loss = run.of(loss_value).to_vec()[0];
losses.push(batch_loss);
if step % 100 == 0 {
println!("step {step:4}: minibatch loss = {batch_loss:.4}");
}
let gradients = run.backward(loss_value);
network = network.update(&gradients, |parameter, gradient| {
parameter.clone() - gradient.clone() * learning_rate.broadcast_like(gradient)
});
}
println!(
"trained {} steps in {:.3}s",
losses.len(),
training.elapsed().as_secs_f64()
);
assert_eq!(network.len(), recorded_nodes);
println!("the tape held {recorded_nodes} nodes through every step");
println!("{}", loss_chart("bigram training", &losses));
let probabilities = network.resolve(table_symbol).softmax(1);
let run = network.forward_for([probabilities.symbol()], std::iter::empty());
let probabilities = run
.of(probabilities)
.as_slice()
.expect("a computed softmax is contiguous")
.to_vec();
let span = (-0.5, VOCABULARY_LEN as f64 - 0.5);
let letters = (0..VOCABULARY_LEN).map(|token| String::from(from_token(token)));
let mut frame = Frame::detect();
frame.height = VOCABULARY_LEN + 6;
println!(
"{}",
Plot::new()
.layer(Cells::matrix(VOCABULARY_LEN, &probabilities[..]).extents(span, span))
.colorbar()
.x_scale(Scale::bands(letters))
.title("bigram transition probabilities")
.x_label("next token")
.y_label("current token")
.render_best(&frame)
);
println!("sampled names:");
let mut state: u64 = 7;
for _ in 0..10 {
let mut token = 0;
let mut name = String::new();
loop {
let row = &probabilities[token * VOCABULARY_LEN..(token + 1) * VOCABULARY_LEN];
token = draw(row, &mut state);
if token == 0 {
break;
}
name.push(from_token(token));
}
println!(" {name}");
}
}