#[allow(dead_code)]
mod corpus;
use std::time::Instant;
use malevich::stat::Window;
use malevich::{Frame, Line, Plot, Rule};
use topos::{Adam, Optimizer, Sgd, Shape, Tape, Tensor, Value, cross_entropy, init};
use corpus::{VOCABULARY_LEN, 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;
const STEPS: usize = 5000;
const BIGRAM_LIMIT: f64 = 2.45;
struct Model<'tape> {
embeddings: Value<'tape, f32>,
hidden_weights: Value<'tape, f32>,
hidden_bias: Value<'tape, f32>,
output_weights: Value<'tape, f32>,
output_bias: Value<'tape, f32>,
}
impl<'tape> Model<'tape> {
fn new(tape: &'tape Tape<f32>) -> Self {
let mut weights = init::xavier(7);
Self {
embeddings: tape.parameter(init::normal(8, 1.0)(&Shape::new([
VOCABULARY_LEN,
EMBED_DIM,
]))),
hidden_weights: tape
.parameter(weights(&Shape::new([CONTEXT_LEN * EMBED_DIM, HIDDEN_LEN]))),
hidden_bias: tape.parameter(weights(&Shape::new([HIDDEN_LEN]))),
output_weights: tape.parameter(weights(&Shape::new([HIDDEN_LEN, VOCABULARY_LEN]))),
output_bias: tape.parameter(weights(&Shape::new([VOCABULARY_LEN]))),
}
}
fn parameters(&self) -> [Value<'tape, f32>; 5] {
[
self.embeddings,
self.hidden_weights,
self.hidden_bias,
self.output_weights,
self.output_bias,
]
}
fn express(&self, contexts: Value<'tape, f32>, rows: usize) -> Value<'tape, 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_like(0, product)).tanh();
let product = hidden.matmul(self.output_weights);
product + self.output_bias.broadcast_along_like(0, product)
}
}
fn train(
samples: &[([usize; CONTEXT_LEN], usize)],
recorded: bool,
optimizer: &mut dyn Optimizer<f32>,
learning_rate: impl Fn(usize) -> Tensor<f32>,
) -> (Vec<f32>, f64) {
let tape = Tape::new();
let model = Model::new(&tape);
let contexts = tape.input(Tensor::selection(
vec![0; BATCH_LEN * CONTEXT_LEN],
VOCABULARY_LEN,
1.0,
));
let targets = tape.input(Tensor::selection(vec![0; BATCH_LEN], VOCABULARY_LEN, 1.0));
let loss = cross_entropy(model.express(contexts, BATCH_LEN), targets);
let (contexts, targets, loss) = (contexts.symbol(), targets.symbol(), loss.symbol());
let parameter_symbols = model.parameters().map(|parameter| parameter.symbol());
let adjoints = recorded.then(|| tape.differentiate(loss, parameter_symbols));
let network = tape.into_network();
let plan = adjoints
.as_ref()
.map(|adjoints| network.entry(adjoints.roots()).lower());
let mut parameters = network.parameters();
let mut losses = Vec::new();
let training = Instant::now();
for step in 0..STEPS {
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 feeds = [
(
contexts,
Tensor::selection(batch_contexts, VOCABULARY_LEN, 1.0),
),
(
targets,
Tensor::selection(batch_targets, VOCABULARY_LEN, 1.0),
),
];
let (batch_loss, gradients) = if let (Some(plan), Some(adjoints)) = (&plan, &adjoints) {
let run = plan.forward(¶meters, feeds);
let gradients = run.recorded_gradients(adjoints);
(run.of(loss).scalar(), gradients)
} else {
let run = network.forward(¶meters, feeds);
(
run.of(loss).scalar(),
run.backward(loss).parameters(¶meters),
)
};
losses.push(batch_loss);
parameters = optimizer.step(¶meters, &gradients, &learning_rate(step));
}
(losses, training.elapsed().as_secs_f64())
}
fn comparison_chart(sgd: &[f32], adam: &[f32]) -> String {
let window_len = (STEPS / 20).max(2);
let smooth = |losses: &[f32]| {
let losses: Vec<f64> = losses.iter().copied().map(f64::from).collect();
Window::new(window_len).mean(&losses)
};
let sgd = smooth(sgd);
let adam = smooth(adam);
Plot::new()
.layer(Line::y(&sgd[..]).label("sgd"))
.layer(Line::y(&adam[..]).label("adam"))
.layer(Rule::h(BIGRAM_LIMIT).label("bigram limit"))
.title("sgd vs adam, rolling mean")
.x_label("step")
.y_label("loss")
.render_best(&Frame::detect())
}
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 recorded = match std::env::var("TOPOS_GRADIENTS").as_deref() {
Ok("engine") => false,
Ok("recorded") | Err(_) => true,
Ok(other) => panic!("unknown TOPOS_GRADIENTS {other:?}; use recorded or engine"),
};
println!(
"gradients: {}",
if recorded {
"recorded (differentiate + one forward-only plan)"
} else {
"engine (interpreter backward)"
}
);
let fast = Tensor::new([], [0.1_f32]);
let slow = Tensor::new([], [0.01_f32]);
let (sgd_losses, sgd_seconds) = train(&samples, recorded, &mut Sgd, |step| {
if step < 4000 {
fast.clone()
} else {
slow.clone()
}
});
let mut adam = Adam::new(
Tensor::new([], [0.9_f32]),
Tensor::new([], [0.999_f32]),
Tensor::new([], [1e-8_f32]),
);
let flat = Tensor::new([], [0.005_f32]);
let (adam_losses, adam_seconds) = train(&samples, recorded, &mut adam, |_| flat.clone());
for (label, losses, seconds) in [
("sgd", &sgd_losses, sgd_seconds),
("adam", &adam_losses, adam_seconds),
] {
let window: f32 = losses[STEPS - 500..].iter().sum::<f32>() / 500.0;
println!(
"{label:5} {STEPS} steps in {seconds:.3}s ({:.2} ms/step), last-500 mean loss {window:.4}",
seconds * 1000.0 / STEPS as f64
);
}
println!("{}", comparison_chart(&sgd_losses, &adam_losses));
}