mod chart;
mod corpus;
use std::time::Instant;
use topos::{Compile, Network, Shape, Tensor, Tensorial, Value, 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 BATCH_LEN: usize = 64;
struct Model<'network> {
embeddings: Value<'network, Tensor<f32>>,
hidden_weights: Value<'network, Tensor<f32>>,
hidden_bias: Value<'network, Tensor<f32>>,
output_weights: Value<'network, Tensor<f32>>,
output_bias: Value<'network, Tensor<f32>>,
}
impl<'network> Model<'network> {
fn new(network: &'network Network<Tensor<f32>>) -> Self {
let mut weights = init::xavier(7);
Self {
embeddings: network.parameter(init::normal(8, 1.0)(&Shape::new([
VOCABULARY_LEN,
EMBED_DIM,
]))),
hidden_weights: network
.parameter(weights(&Shape::new([CONTEXT_LEN * EMBED_DIM, HIDDEN_LEN]))),
hidden_bias: network.parameter(weights(&Shape::new([HIDDEN_LEN]))),
output_weights: network.parameter(weights(&Shape::new([HIDDEN_LEN, VOCABULARY_LEN]))),
output_bias: network.parameter(weights(&Shape::new([VOCABULARY_LEN]))),
}
}
fn parameters(&self) -> [Value<'network, Tensor<f32>>; 5] {
[
self.embeddings,
self.hidden_weights,
self.hidden_bias,
self.output_weights,
self.output_bias,
]
}
fn express(
&self,
contexts: Value<'network, Tensor<f32>>,
rows: usize,
) -> Value<'network, Tensor<f32>> {
let embedded = self
.embeddings
.gather(contexts)
.reshape([rows, CONTEXT_LEN * EMBED_DIM]);
let product = embedded.matmul(self.hidden_weights);
let hidden = (product + self.hidden_bias.broadcast_along(0, product)).tanh();
let product = hidden.matmul(self.output_weights);
product + self.output_bias.broadcast_along(0, product)
}
}
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 network = Network::new();
let model = Model::new(&network);
let contexts = network.input(Tensor::selection(
vec![0; BATCH_LEN * CONTEXT_LEN],
VOCABULARY_LEN,
1.0,
));
let targets = network.input(Tensor::selection(vec![0; BATCH_LEN], VOCABULARY_LEN, 1.0));
let loss = cross_entropy(model.express(contexts, BATCH_LEN), targets);
let sample_context =
network.input(Tensor::selection(vec![0; CONTEXT_LEN], VOCABULARY_LEN, 1.0));
let sample_probabilities = model.express(sample_context, 1).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 forward_nodes = network.len();
let parameter_symbols = model.parameters().map(|parameter| parameter.symbol());
let gradient_symbols = network.differentiate(loss_symbol, parameter_symbols);
println!(
"recorded the chain rule: {} forward nodes + {} gradient nodes",
forward_nodes,
network.len() - forward_nodes
);
let plan = network.compile(Compile::roots(
std::iter::once(loss_symbol).chain(gradient_symbols.iter().copied()),
));
let fast = Tensor::new([], [0.1]);
let slow = Tensor::new([], [0.01]);
let mut network = network;
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 batch_contexts: Vec<usize> = batch
.iter()
.flat_map(|(context, _)| context.iter().copied())
.collect();
let batch_targets: Vec<usize> = batch.iter().map(|&(_, next)| next).collect();
let run = plan.forward(
&network,
[
(
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(network.resolve(loss_symbol)).to_vec()[0];
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 gradients =
run.recorded_gradients(parameter_symbols.iter().zip(&gradient_symbols).map(
|(¶meter, &gradient)| (network.resolve(parameter), network.resolve(gradient)),
));
let learning_rate = if step < 4000 { &fast } else { &slow };
network = network.update(&gradients, |parameter, gradient| {
parameter.clone() - gradient.clone() * learning_rate.broadcast_like(gradient)
});
}
println!(
"trained {} steps in {:.3}s, one plan run per step, no backward pass",
losses.len(),
training.elapsed().as_secs_f64()
);
println!("{}", loss_chart("compiled mlp 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.forward_for(
[sample_probabilities_symbol],
[(
sample_context_symbol,
Tensor::selection(window.to_vec(), VOCABULARY_LEN, 1.0),
)],
);
let row = run
.of(network.resolve(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}");
}
}