#![allow(dead_code)]
use fluxbench::{flux, Bencher};
use std::hint::black_box;
use splintr::pretrained::{llama3_special_tokens, LLAMA3_VOCAB_PACKED};
use splintr::{Tokenizer, LLAMA3_PATTERN, NO_SPLIT_PATTERN};
struct Rng(u64);
impl Rng {
fn next(&mut self) -> usize {
self.0 = self
.0
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(self.0 >> 33) as usize
}
}
fn tokenizer() -> Tokenizer {
Tokenizer::from_packed_chain(
LLAMA3_VOCAB_PACKED,
&[LLAMA3_PATTERN],
llama3_special_tokens(),
)
.expect("bundled llama3 vocabulary must load")
}
fn unsplit_tokenizer() -> Tokenizer {
Tokenizer::from_packed_chain(
LLAMA3_VOCAB_PACKED,
&[NO_SPLIT_PATTERN],
llama3_special_tokens(),
)
.expect("bundled llama3 vocabulary must load under a no-split pattern")
}
fn novel_text(words: usize, seed: u64) -> String {
let alphabet = b"abcdefghijklmnopqrstuvwxyz";
let mut rng = Rng(seed | 1);
let mut out = String::new();
for _ in 0..words {
if !out.is_empty() {
out.push(' ');
}
let len = 4 + rng.next() % 9;
for _ in 0..len {
out.push(alphabet[rng.next() % alphabet.len()] as char);
}
}
out
}
fn repetitive_text(words: usize, seed: u64) -> String {
const WORDS: [&str; 12] = [
"the",
"tokenizer",
"encodes",
"text",
"into",
"tokens",
"and",
"then",
"decodes",
"them",
"again",
"quickly",
];
let mut rng = Rng(seed | 1);
let mut out = String::new();
for _ in 0..words {
if !out.is_empty() {
out.push(' ');
}
out.push_str(WORDS[rng.next() % WORDS.len()]);
}
out
}
#[flux::bench(group = "short_chunks", args = [64, 256, 1024])]
fn encode_novel(b: &mut Bencher, words: usize) {
let tok = tokenizer();
let text = novel_text(words, 0x51D3);
b.iter(|| black_box(tok.encode_ordinary(black_box(&text))));
}
#[flux::bench(group = "cached_chunks", args = [64, 256, 1024])]
fn encode_repetitive(b: &mut Bencher, words: usize) {
let tok = tokenizer();
let text = repetitive_text(words, 0xA71E);
black_box(tok.encode_ordinary(&text));
b.iter(|| black_box(tok.encode_ordinary(black_box(&text))));
}
#[flux::bench(group = "unsplit_piece", args = [500, 1000, 2000, 4000, 8000])]
fn encode_one_long_piece(b: &mut Bencher, chars: usize) {
let tok = unsplit_tokenizer();
let text: String = novel_text(4000, 0x3311).chars().take(chars).collect();
b.iter(|| black_box(tok.encode_ordinary(black_box(&text))));
}
#[flux::bench(group = "batch_scaling", args = [512])]
fn batch_sequential(b: &mut Bencher, texts: usize) {
let tok = tokenizer();
let corpus: Vec<String> = (0..texts)
.map(|i| novel_text(24, 0x9E37 ^ i as u64))
.collect();
b.iter(|| {
for text in &corpus {
black_box(tok.encode(text));
}
});
}
#[flux::bench(group = "batch_scaling", args = [512])]
fn batch_parallel(b: &mut Bencher, texts: usize) {
let tok = tokenizer();
let corpus: Vec<String> = (0..texts)
.map(|i| novel_text(24, 0x9E37 ^ i as u64))
.collect();
b.iter(|| black_box(tok.encode_batch(black_box(&corpus))));
}
#[flux::synthetic(
id = "batch_speedup_512",
formula = "batch_sequential@512 / batch_parallel@512",
unit = "x"
)]
struct BatchSpeedup512;
fn main() {
fluxbench::run().unwrap();
}