use rustc_hash::FxHashMap;
use std::collections::BinaryHeap;
#[derive(Debug, Clone, Copy)]
struct Node {
prev: usize,
next: usize,
start: usize,
len: usize,
}
#[derive(Debug, Clone, Copy)]
struct Merge {
left: usize,
right: usize,
rank: u32,
len: usize,
}
impl PartialEq for Merge {
fn eq(&self, other: &Self) -> bool {
self.rank == other.rank && self.left == other.left
}
}
impl Eq for Merge {}
impl Ord for Merge {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
other
.rank
.cmp(&self.rank)
.then_with(|| other.left.cmp(&self.left))
}
}
impl PartialOrd for Merge {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
pub fn byte_pair_encode(piece: &[u8], encoder: &FxHashMap<Vec<u8>, u32>) -> Vec<u32> {
byte_pair_encode_with_ranks(piece, encoder, encoder)
}
pub fn byte_pair_encode_with_ranks(
piece: &[u8],
merge_ranks: &FxHashMap<Vec<u8>, u32>,
id_encoder: &FxHashMap<Vec<u8>, u32>,
) -> Vec<u32> {
byte_pair_encode_pieces(piece, merge_ranks, id_encoder)
.into_iter()
.filter_map(|p| match p {
Piece::Token(id) => Some(id),
Piece::Unresolved { .. } => None,
})
.collect()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum Piece {
Token(u32),
Unresolved {
start: usize,
len: usize,
},
}
pub(super) fn byte_pair_encode_pieces(
piece: &[u8],
merge_ranks: &FxHashMap<Vec<u8>, u32>,
id_encoder: &FxHashMap<Vec<u8>, u32>,
) -> Vec<Piece> {
byte_pair_encode_pieces_seeded(piece, merge_ranks, id_encoder, false)
}
pub(super) fn byte_pair_encode_pieces_seeded(
piece: &[u8],
merge_ranks: &FxHashMap<Vec<u8>, u32>,
id_encoder: &FxHashMap<Vec<u8>, u32>,
char_granular: bool,
) -> Vec<Piece> {
if piece.is_empty() {
return vec![];
}
if piece.len() == 1 {
return match id_encoder.get(piece) {
Some(&r) => vec![Piece::Token(r)],
None => vec![Piece::Unresolved { start: 0, len: 1 }],
};
}
if let Some(&id) = id_encoder.get(piece) {
return vec![Piece::Token(id)];
}
let nodes: Vec<Node> = match char_granular
.then(|| std::str::from_utf8(piece).ok())
.flatten()
{
Some(text) => Vec::from_iter(text.char_indices().map(|(start, c)| Node {
prev: 0,
next: 0,
start,
len: c.len_utf8(),
})),
None => Vec::from_iter((0..piece.len()).map(|start| Node {
prev: 0,
next: 0,
start,
len: 1,
})),
};
merge_and_collect(piece, nodes, merge_ranks, id_encoder, None)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) struct Seed {
pub(super) start: usize,
pub(super) len: usize,
pub(super) id: Option<u32>,
}
pub(super) fn byte_pair_encode_pieces_presegmented(
piece: &[u8],
seeds: &[Seed],
merge_ranks: &FxHashMap<Vec<u8>, u32>,
id_encoder: &FxHashMap<Vec<u8>, u32>,
) -> Vec<Piece> {
if seeds.is_empty() {
return vec![];
}
let nodes = Vec::from_iter(seeds.iter().map(|seed| Node {
prev: 0,
next: 0,
start: seed.start,
len: seed.len,
}));
merge_and_collect(piece, nodes, merge_ranks, id_encoder, Some(seeds))
}
fn merge_and_collect(
piece: &[u8],
mut nodes: Vec<Node>,
merge_ranks: &FxHashMap<Vec<u8>, u32>,
id_encoder: &FxHashMap<Vec<u8>, u32>,
seeds: Option<&[Seed]>,
) -> Vec<Piece> {
let count = nodes.len();
for (i, node) in nodes.iter_mut().enumerate() {
node.prev = if i == 0 { usize::MAX } else { i - 1 };
node.next = if i + 1 == count { usize::MAX } else { i + 1 };
}
let push = |queue: &mut BinaryHeap<Merge>, left: usize, right: usize, nodes: &[Node]| {
if left == usize::MAX || right == usize::MAX {
return;
}
let (l, r) = (&nodes[left], &nodes[right]);
let len = l.len + r.len;
let slice = &piece[l.start..l.start + len];
if let Some(rank) = merge_ranks.get(slice).copied().filter(|&r| r != u32::MAX) {
queue.push(Merge {
left,
right,
rank,
len,
});
}
};
let mut queue: BinaryHeap<Merge> = BinaryHeap::new();
for i in 0..nodes.len().saturating_sub(1) {
push(&mut queue, i, i + 1, &nodes);
}
while let Some(candidate) = queue.pop() {
let (li, ri) = (candidate.left, candidate.right);
let (left, right) = (nodes[li], nodes[ri]);
if left.len == 0 || left.next != ri || left.len + right.len != candidate.len {
continue;
}
nodes[li].len = left.len + right.len;
nodes[ri].len = 0;
let new_next = right.next;
nodes[li].next = new_next;
if new_next != usize::MAX {
nodes[new_next].prev = li;
}
push(&mut queue, nodes[li].prev, li, &nodes);
push(&mut queue, li, new_next, &nodes);
}
let mut result: Vec<Piece> = Vec::with_capacity(piece.len());
let mut curr = 0;
let mut unresolved_run: Option<(usize, usize)> = None;
while curr != usize::MAX {
let node = &nodes[curr];
let slice = &piece[node.start..node.start + node.len];
let seeded = seeds
.and_then(|seeds| seeds.get(curr))
.filter(|seed| seed.len == node.len)
.and_then(|seed| seed.id);
if let Some(id) = seeded.or_else(|| id_encoder.get(slice).copied()) {
if let Some((start, len)) = unresolved_run.take() {
result.push(Piece::Unresolved { start, len });
}
result.push(Piece::Token(id));
} else {
for (offset, &byte) in slice.iter().enumerate() {
if let Some(&id) = id_encoder.get(&[byte][..]) {
if let Some((start, len)) = unresolved_run.take() {
result.push(Piece::Unresolved { start, len });
}
result.push(Piece::Token(id));
} else {
let byte_start = node.start + offset;
match &mut unresolved_run {
Some((start, len)) if *start + *len == byte_start => *len += 1,
_ => {
if let Some((start, len)) = unresolved_run.take() {
result.push(Piece::Unresolved { start, len });
}
unresolved_run = Some((byte_start, 1));
}
}
}
}
}
curr = nodes[curr].next;
}
if let Some((start, len)) = unresolved_run.take() {
result.push(Piece::Unresolved { start, len });
}
result
}
pub(super) fn merge_ranks<'a>(
merged: &[String],
vocab_in_id_order: impl Iterator<Item = &'a str>,
) -> FxHashMap<Vec<u8>, u32> {
let merge_set: std::collections::HashSet<&str> = merged.iter().map(String::as_str).collect();
let mut ranks: FxHashMap<Vec<u8>, u32> = FxHashMap::default();
for token in vocab_in_id_order.filter(|t| !merge_set.contains(t)) {
let next = ranks.len() as u32;
ranks.entry(token.as_bytes().to_vec()).or_insert(next);
}
let base_count = ranks.len() as u32;
for (i, token) in merged.iter().enumerate() {
ranks
.entry(token.as_bytes().to_vec())
.or_insert(base_count + i as u32);
}
ranks
}
#[cfg(test)]
mod tests {
use super::*;
use proptest::prelude::*;
#[derive(Debug, Clone, Copy)]
struct RefNode {
prev: usize,
next: usize,
rank: u32,
start: usize,
len: usize,
}
fn byte_pair_encode_reference(
piece: &[u8],
merge_ranks: &FxHashMap<Vec<u8>, u32>,
id_encoder: &FxHashMap<Vec<u8>, u32>,
) -> Vec<u32> {
byte_pair_encode_reference_seeded(piece, merge_ranks, id_encoder, false)
}
fn byte_pair_encode_reference_seeded(
piece: &[u8],
merge_ranks: &FxHashMap<Vec<u8>, u32>,
id_encoder: &FxHashMap<Vec<u8>, u32>,
char_granular: bool,
) -> Vec<u32> {
if piece.is_empty() {
return vec![];
}
if piece.len() == 1 {
return id_encoder.get(piece).copied().map_or(vec![], |r| vec![r]);
}
if let Some(&id) = id_encoder.get(piece) {
return vec![id];
}
let spans: Vec<(usize, usize)> = match char_granular
.then(|| std::str::from_utf8(piece).ok())
.flatten()
{
Some(text) => text
.char_indices()
.map(|(start, c)| (start, c.len_utf8()))
.collect(),
None => (0..piece.len()).map(|start| (start, 1)).collect(),
};
let mut nodes: Vec<RefNode> = Vec::with_capacity(spans.len());
for (i, &(start, len)) in spans.iter().enumerate() {
nodes.push(RefNode {
prev: if i == 0 { usize::MAX } else { i - 1 },
next: if i + 1 == spans.len() {
usize::MAX
} else {
i + 1
},
rank: u32::MAX,
start,
len,
});
}
let get_rank = |left_idx: usize, right_idx: usize, nodes: &[RefNode]| -> u32 {
if left_idx == usize::MAX || right_idx == usize::MAX {
return u32::MAX;
}
let left = &nodes[left_idx];
let right = &nodes[right_idx];
let start = left.start;
let len = left.len + right.len;
let slice = &piece[start..start + len];
merge_ranks.get(slice).copied().unwrap_or(u32::MAX)
};
for i in 0..nodes.len() - 1 {
nodes[i].rank = get_rank(i, nodes[i].next, &nodes);
}
loop {
let mut min_rank = u32::MAX;
let mut min_idx = usize::MAX;
let mut curr = 0;
while nodes[curr].prev != usize::MAX {
curr = nodes[curr].prev;
}
while curr != usize::MAX {
let r = nodes[curr].rank;
if r < min_rank {
min_rank = r;
min_idx = curr;
}
curr = nodes[curr].next;
}
if min_rank == u32::MAX {
break;
}
let next_idx = nodes[min_idx].next;
nodes[min_idx].len += nodes[next_idx].len;
let new_next = nodes[next_idx].next;
nodes[min_idx].next = new_next;
if new_next != usize::MAX {
nodes[new_next].prev = min_idx;
}
if nodes[min_idx].prev != usize::MAX {
let prev = nodes[min_idx].prev;
nodes[prev].rank = get_rank(prev, min_idx, &nodes);
}
nodes[min_idx].rank = get_rank(min_idx, nodes[min_idx].next, &nodes);
}
let mut result = Vec::new();
let mut curr = 0;
while nodes[curr].prev != usize::MAX {
curr = nodes[curr].prev;
}
while curr != usize::MAX {
let node = &nodes[curr];
let slice = &piece[node.start..node.start + node.len];
if let Some(&id) = id_encoder.get(slice) {
result.push(id);
} else {
for &byte in slice {
if let Some(&id) = id_encoder.get(&[byte][..]) {
result.push(id);
}
}
}
curr = nodes[curr].next;
}
result
}
fn make_encoder() -> FxHashMap<Vec<u8>, u32> {
let mut encoder = FxHashMap::default();
encoder.insert(b"a".to_vec(), 0);
encoder.insert(b"b".to_vec(), 1);
encoder.insert(b"c".to_vec(), 2);
encoder.insert(b"ab".to_vec(), 3);
encoder.insert(b"bc".to_vec(), 4);
encoder.insert(b"abc".to_vec(), 5);
encoder
}
#[test]
fn test_single_byte() {
let encoder = make_encoder();
assert_eq!(byte_pair_encode(b"a", &encoder), vec![0]);
}
#[test]
fn test_simple_merge() {
let encoder = make_encoder();
assert_eq!(byte_pair_encode(b"ab", &encoder), vec![3]);
}
#[test]
fn test_chain_merge() {
let encoder = make_encoder();
assert_eq!(byte_pair_encode(b"abc", &encoder), vec![5]);
}
#[test]
fn test_empty() {
let encoder = make_encoder();
let empty: Vec<u32> = vec![];
assert_eq!(byte_pair_encode(b"", &encoder), empty);
}
#[test]
fn test_no_merge_possible() {
let encoder = make_encoder();
assert_eq!(byte_pair_encode(b"ac", &encoder), vec![0, 2]);
}
fn tie_encoder() -> FxHashMap<Vec<u8>, u32> {
let mut encoder = FxHashMap::default();
encoder.insert(b"a".to_vec(), 0);
encoder.insert(b"aa".to_vec(), 1);
encoder
}
#[test]
fn test_tiebreak_leftmost_wins() {
let encoder = tie_encoder();
assert_eq!(byte_pair_encode(b"aaa", &encoder), vec![1, 0]);
assert_eq!(byte_pair_encode(b"aaaaa", &encoder), vec![1, 1, 0]);
assert_eq!(
byte_pair_encode_reference(b"aaa", &encoder, &encoder),
vec![1, 0]
);
assert_eq!(
byte_pair_encode_reference(b"aaaaa", &encoder, &encoder),
vec![1, 1, 0]
);
}
fn tokens_only(pieces: Vec<Piece>) -> Vec<u32> {
pieces
.into_iter()
.filter_map(|p| match p {
Piece::Token(id) => Some(id),
Piece::Unresolved { .. } => None,
})
.collect()
}
#[test]
fn test_tiebreak_leftmost_wins_multibyte_chars() {
let mut encoder = FxHashMap::default();
encoder.insert("▁".as_bytes().to_vec(), 0);
encoder.insert("▁▁".as_bytes().to_vec(), 1);
let piece = "▁▁▁".as_bytes();
assert_eq!(
tokens_only(byte_pair_encode_pieces_seeded(
piece, &encoder, &encoder, true
)),
vec![1, 0]
);
assert_eq!(
byte_pair_encode_reference_seeded(piece, &encoder, &encoder, true),
vec![1, 0]
);
}
fn prop_encoder() -> FxHashMap<Vec<u8>, u32> {
let mut encoder = FxHashMap::default();
let tokens: [&[u8]; 18] = [
b"a", b"b", b"c", b"d", b"aa", b"ab", b"ba", b"bb", b"cd", b"dc", b"cc", b"aaa",
b"aab", b"abab", b"aaaa", b"bcd", b"abcd", b"abc",
];
for (i, token) in tokens.iter().enumerate() {
encoder.insert(token.to_vec(), i as u32);
}
encoder
}
fn prop_two_maps() -> (FxHashMap<Vec<u8>, u32>, FxHashMap<Vec<u8>, u32>) {
let ranked: [(&[u8], u32); 13] = [
(b"aa", 1),
(b"ab", 1),
(b"bb", 1),
(b"ba", 2),
(b"cc", 2),
(b"cd", 3),
(b"dc", 3),
(b"aaa", 4),
(b"aab", 4),
(b"abb", 4),
(b"abab", 5),
(b"aaaa", 5),
(b"abcd", 6),
];
let mut merge_ranks = FxHashMap::default();
for (token, rank) in ranked {
merge_ranks.insert(token.to_vec(), rank);
}
let ids: [&[u8]; 17] = [
b"abcd", b"aaaa", b"abab", b"abb", b"aab", b"aaa", b"dc", b"cd", b"cc", b"ba", b"bb",
b"ab", b"aa", b"d", b"c", b"b", b"a",
];
let mut id_encoder = FxHashMap::default();
for (i, token) in ids.iter().enumerate() {
id_encoder.insert(token.to_vec(), i as u32);
}
(merge_ranks, id_encoder)
}
fn char_prop_maps() -> (FxHashMap<Vec<u8>, u32>, FxHashMap<Vec<u8>, u32>) {
let ranked: [(&str, u32); 10] = [
("aa", 1),
("ab", 1),
("▁a", 1),
("b▁", 2),
("中😀", 2),
("😀😀", 2),
("aab", 3),
("ab▁", 3),
("中😀中", 4),
("aaaa", 5),
];
let mut merge_ranks = FxHashMap::default();
for (token, rank) in ranked {
merge_ranks.insert(token.as_bytes().to_vec(), rank);
}
let ids: [&str; 14] = [
"a",
"b",
"▁",
"中",
"😀",
"aaaa",
"中😀中",
"ab▁",
"aab",
"😀😀",
"b▁",
"▁a",
"ab",
"aa",
];
let mut id_encoder = FxHashMap::default();
for (i, token) in ids.iter().enumerate() {
id_encoder.insert(token.as_bytes().to_vec(), i as u32);
}
(merge_ranks, id_encoder)
}
#[test]
fn test_long_single_char_run() {
let encoder = tie_encoder();
let piece = vec![b'a'; 4096];
let expected = vec![1u32; 2048];
assert_eq!(byte_pair_encode(&piece, &encoder), expected);
assert_eq!(
byte_pair_encode_reference(&piece, &encoder, &encoder),
expected
);
}
#[test]
fn test_repeated_ab() {
let encoder = prop_encoder();
let piece = b"ab".repeat(512);
assert_eq!(
byte_pair_encode(&piece, &encoder),
byte_pair_encode_reference(&piece, &encoder, &encoder)
);
}
#[test]
fn test_whole_piece_is_one_token() {
let encoder = prop_encoder();
assert_eq!(byte_pair_encode(b"abcd", &encoder), vec![16]);
}
#[test]
fn test_single_byte_and_empty_agree_with_reference() {
let encoder = prop_encoder();
let pieces: [&[u8]; 3] = [b"", b"a", b"z"];
for piece in pieces {
assert_eq!(
byte_pair_encode(piece, &encoder),
byte_pair_encode_reference(piece, &encoder, &encoder),
"piece {piece:?}"
);
}
}
#[test]
fn test_bytes_absent_from_vocab() {
let encoder = prop_encoder();
let pieces: [&[u8]; 4] = [b"azb", b"\xff\xfe", b"abzcd", b"zzz"];
for piece in pieces {
assert_eq!(
byte_pair_encode(piece, &encoder),
byte_pair_encode_reference(piece, &encoder, &encoder),
"piece {piece:?}"
);
}
}
fn ac_only_encoder() -> FxHashMap<Vec<u8>, u32> {
let mut encoder = FxHashMap::default();
encoder.insert(b"a".to_vec(), 0);
encoder.insert(b"c".to_vec(), 1);
encoder
}
#[test]
fn test_pieces_reports_unresolved_span() {
let encoder = ac_only_encoder();
assert_eq!(
byte_pair_encode_pieces(b"abc", &encoder, &encoder),
vec![
Piece::Token(0),
Piece::Unresolved { start: 1, len: 1 },
Piece::Token(1),
]
);
assert_eq!(byte_pair_encode(b"abc", &encoder), vec![0, 1]);
}
#[test]
fn test_pieces_coalesce_consecutive_unresolved() {
let mut encoder = FxHashMap::default();
encoder.insert(b"a".to_vec(), 0);
encoder.insert(b"d".to_vec(), 1);
assert_eq!(
byte_pair_encode_pieces(b"abcd", &encoder, &encoder),
vec![
Piece::Token(0),
Piece::Unresolved { start: 1, len: 2 },
Piece::Token(1),
]
);
}
#[test]
fn test_pieces_unresolved_at_start_and_end() {
let encoder = ac_only_encoder();
assert_eq!(
byte_pair_encode_pieces(b"bab", &encoder, &encoder),
vec![
Piece::Unresolved { start: 0, len: 1 },
Piece::Token(0),
Piece::Unresolved { start: 2, len: 1 },
]
);
}
#[test]
fn test_pieces_full_coverage_has_no_unresolved() {
let encoder = prop_encoder();
for piece in [&b""[..], b"a", b"abcd", b"aabbccdd", b"dcbadcba"] {
let pieces = byte_pair_encode_pieces(piece, &encoder, &encoder);
assert!(
pieces.iter().all(|p| matches!(p, Piece::Token(_))),
"piece {piece:?} produced {pieces:?}"
);
}
}
#[test]
fn test_pieces_single_byte_and_empty_fast_paths() {
let encoder = ac_only_encoder();
let empty: Vec<Piece> = vec![];
assert_eq!(byte_pair_encode_pieces(b"", &encoder, &encoder), empty);
assert_eq!(
byte_pair_encode_pieces(b"a", &encoder, &encoder),
vec![Piece::Token(0)]
);
assert_eq!(
byte_pair_encode_pieces(b"b", &encoder, &encoder),
vec![Piece::Unresolved { start: 0, len: 1 }]
);
}
proptest! {
#[test]
fn prop_matches_reference_single_map(
piece in prop::collection::vec(prop::sample::select(vec![b'a', b'b', b'c', b'd']), 0..64)
) {
let encoder = prop_encoder();
prop_assert_eq!(
byte_pair_encode(&piece, &encoder),
byte_pair_encode_reference(&piece, &encoder, &encoder)
);
}
#[test]
fn prop_matches_reference_arbitrary_bytes(
piece in prop::collection::vec(any::<u8>(), 0..48)
) {
let encoder = prop_encoder();
prop_assert_eq!(
byte_pair_encode(&piece, &encoder),
byte_pair_encode_reference(&piece, &encoder, &encoder)
);
}
#[test]
fn prop_matches_reference_two_maps(
piece in prop::collection::vec(prop::sample::select(vec![b'a', b'b', b'c', b'd', b'z']), 0..64)
) {
let (merge_ranks, id_encoder) = prop_two_maps();
prop_assert_eq!(
byte_pair_encode_with_ranks(&piece, &merge_ranks, &id_encoder),
byte_pair_encode_reference(&piece, &merge_ranks, &id_encoder)
);
}
#[test]
fn prop_pieces_tokens_match_with_ranks(
piece in prop::collection::vec(any::<u8>(), 0..48)
) {
let encoder = prop_encoder();
let tokens_only: Vec<u32> = byte_pair_encode_pieces(&piece, &encoder, &encoder)
.into_iter()
.filter_map(|p| match p {
Piece::Token(id) => Some(id),
Piece::Unresolved { .. } => None,
})
.collect();
prop_assert_eq!(
tokens_only,
byte_pair_encode_with_ranks(&piece, &encoder, &encoder)
);
}
#[test]
fn prop_char_seeded_matches_reference(
chars in prop::collection::vec(
prop::sample::select(vec!['a', 'b', '▁', '中', '😀']), 0..32)
) {
let (merge_ranks, id_encoder) = char_prop_maps();
let text: String = chars.into_iter().collect();
let piece = text.as_bytes();
prop_assert_eq!(
tokens_only(byte_pair_encode_pieces_seeded(piece, &merge_ranks, &id_encoder, true)),
byte_pair_encode_reference_seeded(piece, &merge_ranks, &id_encoder, true)
);
}
}
}