mod model;
mod tokenizer;
mod weights;
use std::io::{Read, Write};
use std::process::{Child, ChildStdin, ChildStdout, Command, Stdio};
use std::time::Instant;
use topos::{Bf16, Element, Emittable, Module, Plan, Run, Symbol, Tape, Tensor, checkpoint};
use model::{CONTEXT_LEN, EMBED_DIM, Gpt2, VOCABULARY_LEN, load};
use tokenizer::Tokenizer;
use weights::{Weights, cached_text};
const END_OF_TEXT: usize = 50256;
enum Engine {
Decode,
Full,
Xla,
}
struct Sampler {
stream: Symbol,
extraction: Symbol,
logits: Symbol,
}
fn record<E: Element + From<f32> + 'static>(tape: &Tape<E>, model: &Gpt2<E>) -> Sampler {
let embedded = tape.input(Tensor::filled([CONTEXT_LEN, EMBED_DIM], E::from(0.0)));
let extraction = tape.input(Tensor::selection(vec![0], CONTEXT_LEN, E::from(1.0)));
let last = model.express(embedded).gather(extraction);
let logits = last.matmul(tape.resolve(model.embeddings()).transpose());
Sampler {
stream: embedded.symbol(),
extraction: extraction.symbol(),
logits: logits.symbol(),
}
}
struct LayerCache {
keys_in: Symbol,
keys_out: Symbol,
values_in: Symbol,
values_out: Symbol,
}
struct Decoder {
stream: Symbol,
position: Symbol,
mask: Symbol,
logits: Symbol,
caches: Vec<LayerCache>,
}
fn record_decode<E: Element + From<f32> + 'static>(tape: &Tape<E>, model: &Gpt2<E>) -> Decoder {
let zeros = |shape: [usize; 2]| Tensor::filled(shape, E::from(0.0));
let stream = tape.input(zeros([1, EMBED_DIM]));
let position = tape.input(Tensor::selection(vec![0], CONTEXT_LEN, E::from(1.0)));
let mask = tape.input(zeros([1, CONTEXT_LEN]));
let pairs: Vec<_> = (0..model.layers())
.map(|_| {
(
tape.input(zeros([CONTEXT_LEN, EMBED_DIM])),
tape.input(zeros([CONTEXT_LEN, EMBED_DIM])),
)
})
.collect();
let (last, updated) = model.express_decode(tape, stream, &pairs, position, mask);
let logits = tape
.resolve(model.embeddings())
.matmul(last.transpose())
.transpose();
let caches = pairs
.iter()
.zip(&updated)
.map(
|((keys_in, values_in), (keys_out, values_out))| LayerCache {
keys_in: keys_in.symbol(),
keys_out: keys_out.symbol(),
values_in: values_in.symbol(),
values_out: values_out.symbol(),
},
)
.collect();
Decoder {
stream: stream.symbol(),
position: position.symbol(),
mask: mask.symbol(),
logits: logits.symbol(),
caches,
}
}
struct Carry<E> {
entries: Vec<(Symbol, Tensor<E>)>,
}
impl<E: Element + From<f32>> Carry<E> {
fn new(caches: &[LayerCache]) -> Self {
let zeros = || Tensor::filled([CONTEXT_LEN, EMBED_DIM], E::from(0.0));
let entries = caches
.iter()
.flat_map(|cache| [(cache.keys_in, zeros()), (cache.values_in, zeros())])
.collect();
Self { entries }
}
fn feeds(&self) -> impl Iterator<Item = (Symbol, Tensor<E>)> + '_ {
self.entries
.iter()
.map(|(symbol, payload)| (*symbol, payload.clone()))
}
fn advanced(run: &Run<E>, caches: &[LayerCache]) -> Self {
let entries = caches
.iter()
.flat_map(|cache| {
[
(cache.keys_in, run.of(cache.keys_out).clone()),
(cache.values_in, run.of(cache.values_out).clone()),
]
})
.collect();
Self { entries }
}
}
struct XlaServer {
child: Child,
requests: ChildStdin,
responses: ChildStdout,
}
impl XlaServer {
fn new<E>(plan: &Plan<E>, arguments: &[Tensor<E>]) -> Self
where
E: Element + Emittable + Copy,
f32: From<E>,
{
let directory = weights::cache_directory();
let module_path = directory.join("gpt2-plan.mlir");
let static_path = directory.join("gpt2-static.bin");
let manifest_path = directory.join("gpt2-manifest.json");
std::fs::write(&module_path, plan.emit_stablehlo().expect("the plan emits"))
.expect("the module writes");
let mut staged = Vec::new();
for tensor in arguments {
let axes = tensor.shape().axes().to_vec();
staged.extend((axes.len() as u32).to_le_bytes());
for extent in axes {
staged.extend((extent as u32).to_le_bytes());
}
for element in tensor.to_vec() {
staged.extend(f32::from(element).to_le_bytes());
}
}
std::fs::write(&static_path, staged).expect("the arguments write");
std::fs::write(
&manifest_path,
format!("{{\"dynamic\": [[{CONTEXT_LEN}, {EMBED_DIM}], [1, {CONTEXT_LEN}]]}}"),
)
.expect("the manifest writes");
let python = std::env::var("TOPOS_XLA_PYTHON").unwrap_or_else(|_| "python3".to_string());
let mut command: Vec<String> = python.split_whitespace().map(str::to_string).collect();
command.push("tools/serve-stablehlo-xla.py".to_string());
let mut child = Command::new(&command[0])
.args(&command[1..])
.arg(&module_path)
.arg(&static_path)
.arg(&manifest_path)
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.spawn()
.expect("the serving process starts; is `jax` installed for it?");
let requests = child.stdin.take().expect("the server's input pipes");
let responses = child.stdout.take().expect("the server's output pipes");
Self {
child,
requests,
responses,
}
}
fn step(&mut self, stream: &[f32], extraction: &[f32]) -> Vec<f32> {
let mut request = Vec::with_capacity(4 * (stream.len() + extraction.len()));
for &value in stream.iter().chain(extraction) {
request.extend(value.to_le_bytes());
}
self.requests
.write_all(&request)
.expect("the request writes");
self.requests.flush().expect("the request flushes");
let mut response = vec![0u8; 4 * VOCABULARY_LEN];
self.responses
.read_exact(&mut response)
.expect("the server answers; see its standard error");
response
.chunks_exact(4)
.map(|chunk| f32::from_le_bytes(chunk.try_into().expect("four bytes")))
.collect()
}
}
impl Drop for XlaServer {
fn drop(&mut self) {
let _ = self.child.kill();
let _ = self.child.wait();
}
}
fn unit(state: &mut u64) -> f64 {
*state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
let bits = (*state >> 11) as f64;
bits / (1u64 << 53) as f64
}
fn draw(logits: &[f32], temperature: f64, top: usize, state: &mut u64) -> usize {
let mut ranked: Vec<(usize, f32)> = logits.iter().copied().enumerate().collect();
ranked.sort_by(|a, b| b.1.total_cmp(&a.1));
ranked.truncate(top);
let peak = ranked[0].1 as f64;
let weights: Vec<f64> = ranked
.iter()
.map(|&(_, logit)| libm::exp((logit as f64 - peak) / temperature))
.collect();
let total: f64 = weights.iter().sum();
let mut remaining = unit(state) * total;
for (&(id, _), weight) in ranked.iter().zip(&weights) {
if remaining < *weight {
return id;
}
remaining -= weight;
}
ranked[0].0
}
fn run<E>(prompt: &str, count: usize, engine: Engine, label: &str)
where
E: Element + Emittable + From<f32> + Copy + 'static,
f32: From<E>,
{
let loading = Instant::now();
let tokenizer = Tokenizer::new(&cached_text("vocab.json"), &cached_text("merges.txt"));
let weights = Weights::load();
let tape = Tape::new();
let gpt2 = Gpt2::<E>::new(&tape);
let sampler = record(&tape, &gpt2);
let decoder = record_decode(&tape, &gpt2);
let network = tape.into_network();
let parameters = load(&network.parameters(), &gpt2, &weights);
drop(weights);
println!(
"loaded the checkpoint in {:.1}s",
loading.elapsed().as_secs_f64()
);
let mut window = vec![END_OF_TEXT];
window.extend(tokenizer.encode(prompt));
assert!(
window.len() + count <= CONTEXT_LEN,
"prompt and generation must fit the {CONTEXT_LEN}-token context"
);
assert_eq!(
tokenizer.decode(&window[1..]),
prompt,
"the tokenizer round-trips the prompt"
);
if let Engine::Decode = engine {
let compiling = Instant::now();
let outputs: Vec<Symbol> = decoder
.caches
.iter()
.flat_map(|cache| [cache.keys_out, cache.values_out])
.collect();
let plan: Plan<E> = network.entry([decoder.logits]).observe(outputs).lower();
println!(
"recorded {} nodes and compiled the decode plan in {:.1}s",
network.len(),
compiling.elapsed().as_secs_f64()
);
let table = parameters.of(gpt2.embeddings()).to_vec();
let row = |token: usize| {
Tensor::new(
[1, EMBED_DIM],
table[token * EMBED_DIM..(token + 1) * EMBED_DIM].to_vec(),
)
};
let mask_row = |until: usize| {
let elements: Vec<E> = (0..CONTEXT_LEN)
.map(|at| {
if at <= until {
E::from(0.0)
} else {
E::from(f32::NEG_INFINITY)
}
})
.collect();
Tensor::new([1, CONTEXT_LEN], elements)
};
let step = |carry: Carry<E>, token: usize, at: usize| -> (Vec<f32>, Carry<E>) {
let run = plan.forward(
¶meters,
carry.feeds().chain([
(decoder.stream, row(token)),
(
decoder.position,
Tensor::selection(vec![at], CONTEXT_LEN, E::from(1.0)),
),
(decoder.mask, mask_row(at)),
]),
);
let logits = run.of(decoder.logits).to_vec();
(
logits.iter().map(|&element| f32::from(element)).collect(),
Carry::advanced(&run, &decoder.caches),
)
};
let prefilling = Instant::now();
let mut carry = Carry::new(&decoder.caches);
let mut logits = Vec::new();
for (at, &token) in window.iter().enumerate() {
(logits, carry) = step(carry, token, at);
}
println!(
"prefilled {} tokens in {:.2}s",
window.len(),
prefilling.elapsed().as_secs_f64()
);
print!("{prompt}");
let mut state: u64 = 7;
let generation = Instant::now();
for index in 0..count {
let token = draw(&logits, 0.9, 40, &mut state);
if token == END_OF_TEXT {
break;
}
window.push(token);
print!("{}", tokenizer.decode(&[token]));
std::io::stdout().flush().expect("stdout flushes");
if index + 1 < count {
(logits, carry) = step(carry, token, window.len() - 1);
}
}
let elapsed = generation.elapsed().as_secs_f64();
let generated = window.len() - 1 - tokenizer.encode(prompt).len();
println!();
println!(
"generated {generated} tokens on the {label} engine in {elapsed:.1}s ({:.0} ms/token)",
elapsed / generated.max(1) as f64 * 1e3
);
return;
}
let compiling = Instant::now();
let plan: Plan<E> = network.entry([sampler.logits]).lower();
println!(
"recorded {} nodes and compiled the plan in {:.1}s",
network.len(),
compiling.elapsed().as_secs_f64()
);
let table = parameters.of(gpt2.embeddings()).to_vec();
let embedded = |window: &[usize]| {
let mut stream = vec![E::from(0.0); CONTEXT_LEN * EMBED_DIM];
for (row, &token) in window.iter().enumerate() {
stream[row * EMBED_DIM..(row + 1) * EMBED_DIM]
.copy_from_slice(&table[token * EMBED_DIM..(token + 1) * EMBED_DIM]);
}
stream
};
let widened = |elements: &[E]| -> Vec<f32> {
elements.iter().map(|&element| f32::from(element)).collect()
};
let mut server = match engine {
Engine::Decode | Engine::Full => None,
Engine::Xla => {
println!("starting the XLA server (compiling the emitted plan) ...");
let arguments = checkpoint::snapshot(¶meters, &gpt2);
let mut server = XlaServer::new(&plan, &arguments);
let extraction = Tensor::selection(vec![0], CONTEXT_LEN, 1.0_f32);
server.step(&widened(&embedded(&window)), &extraction.to_vec());
Some(server)
}
};
print!("{prompt}");
let mut state: u64 = 7;
let generation = Instant::now();
for _ in 0..count {
let stream = embedded(&window);
let extraction = Tensor::selection(vec![window.len() - 1], CONTEXT_LEN, E::from(1.0));
let logits = match &mut server {
Some(server) => server.step(&widened(&stream), &widened(&extraction.to_vec())),
None => {
let run = plan.forward(
¶meters,
[
(
sampler.stream,
Tensor::new([CONTEXT_LEN, EMBED_DIM], stream),
),
(sampler.extraction, extraction),
],
);
widened(&run.of(sampler.logits).to_vec())
}
};
let token = draw(&logits, 0.9, 40, &mut state);
if token == END_OF_TEXT {
break;
}
window.push(token);
print!("{}", tokenizer.decode(&[token]));
std::io::stdout().flush().expect("stdout flushes");
}
let elapsed = generation.elapsed().as_secs_f64();
let generated = window.len() - 1 - tokenizer.encode(prompt).len();
println!();
println!(
"generated {generated} tokens on the {label} engine in {elapsed:.1}s ({:.0} ms/token)",
elapsed / generated.max(1) as f64 * 1e3
);
}
fn main() {
let prompt = std::env::args()
.nth(1)
.unwrap_or_else(|| "The library of this place holds one book".to_string());
let count: usize = std::env::args()
.nth(2)
.map(|argument| argument.parse().expect("a token count"))
.unwrap_or(40);
let engine = std::env::args()
.nth(3)
.unwrap_or_else(|| "tape".to_string());
match engine.as_str() {
"tape" => run::<f32>(&prompt, count, Engine::Decode, "tape"),
"full" => run::<f32>(&prompt, count, Engine::Full, "full"),
"xla" => run::<f32>(&prompt, count, Engine::Xla, "xla"),
"bf16" => run::<Bf16>(&prompt, count, Engine::Decode, "bf16"),
other => panic!("unknown engine `{other}`; use `tape`, `full`, `xla`, or `bf16`"),
}
}