use crate::core::encoder::Encoder;
use crate::core::precompiled::utf8_len;
use super::merge::{
merge_and_collect, merge_and_collect_ids_into, merge_and_count, prefers_scan, QueueScratch,
SCAN_SYMBOL_LIMIT,
};
use super::nodes::Node;
use super::ranks::{PairRanks, RankLookup};
use super::scratch::with_merge_scratch;
pub fn byte_pair_encode(piece: &[u8], encoder: &Encoder) -> Vec<u32> {
byte_pair_encode_with_ranks(piece, encoder, encoder)
}
pub fn byte_pair_encode_with_ranks(
piece: &[u8],
merge_ranks: &Encoder,
id_encoder: &Encoder,
) -> Vec<u32> {
let mut out = Vec::new();
byte_pair_encode_ids_seeded_into(
piece,
RankLookup::new(merge_ranks),
id_encoder,
Seeding::Bytes,
&mut out,
);
out
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum Seeding {
Bytes,
Chars,
RawBytes,
RawChars,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum Piece {
Token(u32),
Unresolved {
start: usize,
len: usize,
},
}
pub(crate) fn byte_pair_encode_pieces_seeded(
piece: &[u8],
merge_ranks: RankLookup<'_>,
id_encoder: &Encoder,
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)];
}
with_merge_scratch(|s| {
seed_nodes_reusing(
piece,
match char_granular {
true => Seeding::Chars,
false => Seeding::Bytes,
},
symbol_count(piece, char_granular),
&mut s.nodes,
None,
);
merge_and_collect(
piece,
&mut s.nodes,
merge_ranks.without_ids(),
id_encoder,
None,
&mut s.queue,
)
})
}
pub(crate) fn byte_pair_encode_ids_seeded_into(
piece: &[u8],
merge_ranks: RankLookup<'_>,
id_encoder: &Encoder,
seeding: Seeding,
out: &mut Vec<u32>,
) {
if piece.is_empty() {
return;
}
if piece.len() == 1 {
out.extend(id_encoder.get(piece));
return;
}
if let Some(id) = id_encoder.get(piece) {
out.push(id);
return;
}
byte_pair_merge_ids_into(piece, merge_ranks, id_encoder, seeding, out)
}
pub(crate) fn byte_pair_encode_ids_or_pieces(
piece: &[u8],
merge_ranks: RankLookup<'_>,
id_encoder: &Encoder,
seeding: Seeding,
out: &mut Vec<u32>,
) -> Option<Vec<Piece>> {
let Some(table) = merge_ranks.by_id() else {
return Some(byte_pair_encode_pieces_seeded(
piece,
merge_ranks,
id_encoder,
seeding == Seeding::Chars,
));
};
if piece.is_empty() {
return None;
}
if piece.len() == 1 {
return match id_encoder.get(piece) {
Some(id) => {
out.push(id);
None
}
None => Some(vec![Piece::Unresolved { start: 0, len: 1 }]),
};
}
if let Some(id) = id_encoder.get(piece) {
out.push(id);
return None;
}
with_merge_scratch(|s| {
if !seed_nodes_pushing(piece, seeding, &mut s.nodes, table) {
return Some(merge_and_collect(
piece,
&mut s.nodes,
merge_ranks.without_ids(),
id_encoder,
None,
&mut s.queue,
));
}
merge_and_collect_ids_into(
piece,
&mut s.nodes,
merge_ranks,
id_encoder,
out,
&mut s.queue,
);
None
})
}
pub(crate) fn merges_to_whole(
piece: &[u8],
merge_ranks: RankLookup<'_>,
char_granular: bool,
) -> bool {
with_merge_scratch(|s| {
s.nodes.clear();
s.nodes.reserve(piece.len());
let mut at = 0;
while at < piece.len() {
let len = match byte_fallback_symbol(&piece[at..], merge_ranks) {
true => BYTE_FALLBACK_SPELLING_LEN,
false => match char_granular {
true => utf8_len(piece[at]).min(piece.len() - at),
false => 1,
},
};
let index = s.nodes.len();
s.nodes.push(Node {
prev: index.wrapping_sub(1),
next: index + 1,
start: at,
len,
id: u32::MAX,
});
at += len;
}
if s.nodes.len() <= 1 {
return true;
}
s.nodes
.last_mut()
.expect("at least two nodes, checked above")
.next = usize::MAX;
merge_and_count(piece, &mut s.nodes, merge_ranks, &mut s.queue) == 1
})
}
const BYTE_FALLBACK_SPELLING_LEN: usize = 6;
#[inline]
fn byte_fallback_symbol(rest: &[u8], merge_ranks: RankLookup<'_>) -> bool {
rest.len() >= BYTE_FALLBACK_SPELLING_LEN
&& crate::core::vocab::is_byte_fallback_piece(&rest[..BYTE_FALLBACK_SPELLING_LEN])
&& merge_ranks.get(&rest[..BYTE_FALLBACK_SPELLING_LEN]) != u32::MAX
}
pub(crate) fn byte_pair_merge_ids_into(
piece: &[u8],
merge_ranks: RankLookup<'_>,
id_encoder: &Encoder,
seeding: Seeding,
out: &mut Vec<u32>,
) {
if !byte_pair_merge_ids_attempt(piece, merge_ranks, id_encoder, seeding, out) {
byte_pair_merge_ids_into(piece, merge_ranks.without_ids(), id_encoder, seeding, out);
}
}
#[inline(always)]
fn byte_pair_merge_ids_attempt(
piece: &[u8],
merge_ranks: RankLookup<'_>,
id_encoder: &Encoder,
seeding: Seeding,
out: &mut Vec<u32>,
) -> bool {
if piece.len() <= 1 {
return true;
}
let seeding = match seeding {
Seeding::Bytes if merge_ranks.by_id().is_some_and(|t| t.seeds_by_char()) => Seeding::Chars,
other => other,
};
if seeding == Seeding::RawChars {
let table = merge_ranks
.by_id()
.expect("RawChars is only chosen when the id table vouched for a character");
with_merge_scratch(|s| {
s.nodes.clear();
s.nodes.reserve(piece.len());
walk_raw_chars(piece, table, |mut node| {
let index = s.nodes.len();
node.prev = index.wrapping_sub(1);
node.next = index + 1;
s.nodes.push(node);
});
if let Some(tail) = s.nodes.last_mut() {
tail.next = usize::MAX;
}
merge_and_collect_ids_into(
piece,
&mut s.nodes,
merge_ranks,
id_encoder,
out,
&mut s.queue,
);
});
return true;
}
let symbols = symbol_count(piece, seeding == Seeding::Chars);
if prefers_scan(piece, symbols) {
let mut buf = [Node::PLACEHOLDER; SCAN_SYMBOL_LIMIT];
let nodes = &mut buf[..symbols];
if !seed_nodes_into(piece, seeding, nodes, merge_ranks.by_id()) {
return false;
}
merge_and_collect_ids_into(
piece,
nodes,
merge_ranks,
id_encoder,
out,
&mut QueueScratch::default(),
);
} else {
let seeded = with_merge_scratch(|s| {
if !seed_nodes_reusing(piece, seeding, symbols, &mut s.nodes, merge_ranks.by_id()) {
return false;
}
merge_and_collect_ids_into(
piece,
&mut s.nodes,
merge_ranks,
id_encoder,
out,
&mut s.queue,
);
true
});
if !seeded {
return false;
}
}
true
}
#[inline]
fn symbol_count(piece: &[u8], char_granular: bool) -> usize {
if !char_granular || std::str::from_utf8(piece).is_err() {
return piece.len();
}
let mut count = 0usize;
let mut chunks = piece.chunks_exact(8);
for chunk in &mut chunks {
let word = u64::from_le_bytes(chunk.try_into().expect("chunks_exact(8) is 8 bytes"));
let continuation = word & !(word << 1) & 0x8080_8080_8080_8080;
count += 8 - continuation.count_ones() as usize;
}
count
+ chunks
.remainder()
.iter()
.filter(|&&b| b & 0xC0 != 0x80)
.count()
}
fn walk_raw_chars(piece: &[u8], table: &PairRanks, mut emit: impl FnMut(Node)) {
let mut at = 0;
while at < piece.len() {
let len = match piece[at] {
b if b < 0x80 => 1,
b if b >> 5 == 0b110 => 2,
b if b >> 4 == 0b1110 => 3,
b if b >> 3 == 0b11110 => 4,
_ => 1,
}
.min(piece.len() - at);
let id = table.char_seed(&piece[at..at + len]);
if len > 1 && id != u32::MAX {
emit(Node {
prev: 0,
next: 0,
start: at,
len,
id,
});
} else {
for (offset, &byte) in piece[at..at + len].iter().enumerate() {
emit(Node {
prev: 0,
next: 0,
start: at + offset,
len: 1,
id: table.raw_byte_id(byte),
});
}
}
at += len;
}
}
fn seed_nodes_pushing(
piece: &[u8],
seeding: Seeding,
nodes: &mut Vec<Node>,
table: &PairRanks,
) -> bool {
nodes.clear();
nodes.reserve(piece.len());
let mut push = |start: usize, len: usize, id: u32| {
let index = nodes.len();
nodes.push(Node {
prev: index.wrapping_sub(1),
next: index + 1,
start,
len,
id,
});
};
let mut resolved = true;
match (seeding == Seeding::Chars)
.then(|| std::str::from_utf8(piece).ok())
.flatten()
{
Some(text) => {
for (start, c) in text.char_indices() {
let len = c.len_utf8();
let id = match resolved {
true => table.seed_id(&piece[start..start + len]),
false => u32::MAX,
};
resolved &= id != u32::MAX;
push(start, len, id);
}
}
None => {
for (start, &byte) in piece.iter().enumerate() {
let id = match (resolved, seeding == Seeding::RawBytes) {
(false, _) => u32::MAX,
(true, true) => table.raw_byte_id(byte),
(true, false) => table.byte_id(byte),
};
resolved &= id != u32::MAX;
push(start, 1, id);
}
}
}
if let Some(tail) = nodes.last_mut() {
tail.next = usize::MAX;
}
resolved
}
fn seed_nodes_reusing(
piece: &[u8],
seeding: Seeding,
symbols: usize,
nodes: &mut Vec<Node>,
table: Option<&PairRanks>,
) -> bool {
nodes.clear();
nodes.resize(symbols, Node::PLACEHOLDER);
seed_nodes_into(piece, seeding, nodes, table)
}
fn seed_nodes_into(
piece: &[u8],
seeding: Seeding,
nodes: &mut [Node],
table: Option<&PairRanks>,
) -> bool {
match (seeding == Seeding::Chars)
.then(|| std::str::from_utf8(piece).ok())
.flatten()
{
Some(text) => {
let count = nodes.len();
for (index, (node, (start, c))) in nodes.iter_mut().zip(text.char_indices()).enumerate()
{
node.start = start;
node.len = c.len_utf8();
node.prev = index.wrapping_sub(1);
node.next = match index + 1 == count {
true => usize::MAX,
false => index + 1,
};
if let Some(table) = table {
node.id = table.seed_id(&piece[start..start + node.len]);
if node.id == u32::MAX {
return false;
}
}
}
}
None => {
let raw = seeding == Seeding::RawBytes;
let count = nodes.len();
for (start, node) in nodes.iter_mut().enumerate() {
node.start = start;
node.len = 1;
node.prev = start.wrapping_sub(1);
node.next = match start + 1 == count {
true => usize::MAX,
false => start + 1,
};
if let Some(table) = table {
node.id = match raw {
true => table.raw_byte_id(piece[start]),
false => table.byte_id(piece[start]),
};
if node.id == u32::MAX {
return false;
}
}
}
}
}
true
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct Seed {
pub(crate) start: usize,
pub(crate) len: usize,
pub(crate) id: Option<u32>,
}
pub(crate) fn byte_pair_encode_pieces_presegmented(
piece: &[u8],
seeds: &[Seed],
merge_ranks: RankLookup<'_>,
id_encoder: &Encoder,
) -> Vec<Piece> {
if seeds.is_empty() {
return vec![];
}
with_merge_scratch(|s| {
let count = seeds.len();
s.nodes
.extend(seeds.iter().enumerate().map(|(index, seed)| Node {
prev: index.wrapping_sub(1),
next: match index + 1 == count {
true => usize::MAX,
false => index + 1,
},
start: seed.start,
len: seed.len,
id: u32::MAX,
}));
merge_and_collect(
piece,
&mut s.nodes,
merge_ranks.without_ids(),
id_encoder,
Some(seeds),
&mut s.queue,
)
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::bpe::ranks::BytePairRanks;
#[test]
fn a_supplied_segmentation_merges_over_multi_byte_symbols() {
let ranks: Encoder = [
(b"he".as_slice(), 0u32),
(b"hel".as_slice(), 1),
(b"hell".as_slice(), 2),
(b"hello</w>".as_slice(), 3),
]
.into_iter()
.collect();
let ids: Encoder = [
(b"h".as_slice(), 10u32),
(b"e".as_slice(), 11),
(b"l".as_slice(), 12),
(b"o</w>".as_slice(), 13),
(b"he".as_slice(), 14),
(b"hel".as_slice(), 15),
(b"hell".as_slice(), 16),
(b"hello</w>".as_slice(), 17),
]
.into_iter()
.collect();
let pairs = BytePairRanks::build(&ranks);
let buf = b"hello</w>";
let mut seeds: Vec<Seed> = (0..5)
.map(|start| Seed {
start,
len: 1,
id: None,
})
.collect();
seeds[4].len = 5;
let out = byte_pair_encode_pieces_presegmented(
buf,
&seeds,
RankLookup::with_pairs(&ranks, &pairs),
&ids,
);
assert_eq!(out, vec![Piece::Token(17)]);
}
#[test]
fn an_unabsorbed_seed_keeps_its_own_id() {
let ranks: Encoder = Encoder::default();
let ids: Encoder = [(b"b".as_slice(), 5u32)].into_iter().collect();
let pairs = BytePairRanks::build(&ranks);
let seeds = [
Seed {
start: 0,
len: 1,
id: Some(99),
},
Seed {
start: 1,
len: 1,
id: None,
},
Seed {
start: 2,
len: 1,
id: None,
},
];
let out = byte_pair_encode_pieces_presegmented(
b"abz",
&seeds,
RankLookup::with_pairs(&ranks, &pairs),
&ids,
);
assert_eq!(
out,
vec![
Piece::Token(99),
Piece::Token(5),
Piece::Unresolved { start: 2, len: 1 },
]
);
}
}