use crate::core::encoder::Encoder;
use super::merge::{
merge_and_collect, merge_and_collect_ids_into, 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_merge_ids_into(
piece: &[u8],
merge_ranks: RankLookup<'_>,
id_encoder: &Encoder,
seeding: Seeding,
out: &mut Vec<u32>,
) {
if piece.len() <= 1 {
return;
}
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");
return 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,
);
});
}
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 byte_pair_merge_ids_into(
piece,
merge_ranks.without_ids(),
id_encoder,
seeding,
out,
);
}
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 {
byte_pair_merge_ids_into(piece, merge_ranks.without_ids(), id_encoder, seeding, out);
}
}
}
#[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_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,
)
})
}