mod chart;
mod corpus;
use std::time::Instant;
use rayon::prelude::*;
use topos::{Activation, Mlp, Module, Parameters, Shape, Tape, Tensor, cross_entropy, init};
use chart::loss_chart;
use corpus::{VOCABULARY_LEN, draw, from_token, load_names, shuffle, training_samples};
const CONTEXT_LEN: usize = 3;
const EMBED_DIM: usize = 10;
const HIDDEN_LEN: usize = 100;
const SHARD_COUNT: usize = 8;
const SHARD_LEN: usize = 8;
fn tree_sum(mut layer: Vec<Parameters<f32>>) -> Parameters<f32> {
while layer.len() > 1 {
layer = layer
.par_chunks(2)
.map(|pair| match pair {
[left, right] => left + right,
[single] => single.clone(),
_ => unreachable!("chunks of two hold one or two tables"),
})
.collect();
}
layer.into_iter().next().expect("at least one shard ran")
}
fn main() {
let names = load_names();
let mut samples = training_samples::<CONTEXT_LEN>(&names);
let mut shuffle_state: u64 = 9;
shuffle(&mut samples, &mut shuffle_state);
println!("loaded {} names, {} samples", names.len(), samples.len());
let tape: Tape<f32> = Tape::new();
let embeddings = tape.parameter(init::normal(8, 1.0)(&Shape::new([
VOCABULARY_LEN,
EMBED_DIM,
])));
let mlp = Mlp::new(
&tape,
&[CONTEXT_LEN * EMBED_DIM, HIDDEN_LEN, VOCABULARY_LEN],
Activation::Tanh,
init::xavier(7),
);
let contexts = tape.input(Tensor::selection(
vec![0; SHARD_LEN * CONTEXT_LEN],
VOCABULARY_LEN,
1.0,
));
let targets = tape.input(Tensor::selection(vec![0; SHARD_LEN], VOCABULARY_LEN, 1.0));
let embedded = embeddings
.gather(contexts)
.reshape([SHARD_LEN, CONTEXT_LEN * EMBED_DIM]);
let loss = cross_entropy(mlp.express(embedded), targets);
let sample_context = tape.input(Tensor::selection(vec![0; CONTEXT_LEN], VOCABULARY_LEN, 1.0));
let sample_embedded = embeddings
.gather(sample_context)
.reshape([1, CONTEXT_LEN * EMBED_DIM]);
let sample_probabilities = mlp.express(sample_embedded).softmax(1);
let contexts_symbol = contexts.symbol();
let targets_symbol = targets.symbol();
let loss_symbol = loss.symbol();
let sample_context_symbol = sample_context.symbol();
let sample_probabilities_symbol = sample_probabilities.symbol();
let network = tape.into_network();
let recorded_nodes = network.len();
let batch_len = SHARD_COUNT * SHARD_LEN;
let shard_inverse = Tensor::new([], [1.0 / SHARD_COUNT as f32]);
let fast = Tensor::new([], [0.1]);
let slow = Tensor::new([], [0.01]);
let mut parameters = network.parameters();
let mut window_loss = 0.0;
let mut losses = Vec::new();
let training = Instant::now();
for step in 0..5000 {
let start = (step * batch_len) % (samples.len() - batch_len);
let batch = &samples[start..start + batch_len];
let shard_results: Vec<(f32, Parameters<f32>)> = (0..SHARD_COUNT)
.into_par_iter()
.map(|shard| {
let rows = &batch[shard * SHARD_LEN..(shard + 1) * SHARD_LEN];
let shard_contexts: Vec<usize> = rows
.iter()
.flat_map(|(context, _)| context.iter().copied())
.collect();
let shard_targets: Vec<usize> = rows.iter().map(|&(_, next)| next).collect();
let run = network.entry([loss_symbol]).interpret(
¶meters,
[
(
contexts_symbol,
Tensor::selection(shard_contexts, VOCABULARY_LEN, 1.0),
),
(
targets_symbol,
Tensor::selection(shard_targets, VOCABULARY_LEN, 1.0),
),
],
);
let shard_loss = run.of(loss_symbol).scalar();
(
shard_loss,
run.backward(loss_symbol).parameters(¶meters),
)
})
.collect();
let (shard_losses, shard_gradients): (Vec<f32>, Vec<Parameters<f32>>) =
shard_results.into_iter().unzip();
let batch_loss = shard_losses.iter().sum::<f32>() / SHARD_COUNT as f32;
let gradients = tree_sum(shard_gradients)
.map(|gradient| gradient.clone() * shard_inverse.broadcast_like(gradient));
losses.push(batch_loss);
if step == 0 {
println!(
"step 0: minibatch loss = {batch_loss:.4} (a uniform model costs ln 27 ~ 3.30)"
);
}
window_loss += batch_loss;
if (step + 1) % 500 == 0 {
println!(
"steps {:4}..{:4}: mean minibatch loss = {:.4}",
step + 1 - 500,
step + 1,
window_loss / 500.0
);
window_loss = 0.0;
}
let learning_rate = if step < 4000 { &fast } else { &slow };
parameters = parameters.step(&gradients, |parameter, gradient| {
parameter.clone() - gradient.clone() * learning_rate.broadcast_like(gradient)
});
}
println!(
"trained {} steps on {SHARD_COUNT} shards 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("mlp (data parallel) training", &losses));
println!("sampled names:");
let mut state: u64 = 7;
for _ in 0..10 {
let mut window = [0usize; CONTEXT_LEN];
let mut name = String::new();
loop {
let run = network.entry([sample_probabilities_symbol]).interpret(
¶meters,
[(
sample_context_symbol,
Tensor::selection(window.to_vec(), VOCABULARY_LEN, 1.0),
)],
);
let row = run.of(sample_probabilities_symbol).to_vec();
let token = draw(&row, &mut state);
if token == 0 {
break;
}
name.push(from_token(token));
window.rotate_left(1);
window[CONTEXT_LEN - 1] = token;
}
println!(" {name}");
}
}