use crate::bitfield::BitField;
use crate::byte_pair_encoding::BytePairEncoding;
pub(crate) struct BacktrackEncoder<'a> {
bpe: &'a BytePairEncoding,
text: &'a [u8],
tokens: Vec<u32>,
next_token: Option<u32>,
pos: usize,
bitfield: BitField,
}
impl<'a> BacktrackEncoder<'a> {
pub(crate) fn new(bpe: &'a BytePairEncoding, text: &'a [u8]) -> Self {
Self::with_capacity(bpe, text, text.len() / 3)
}
pub(crate) fn with_capacity(bpe: &'a BytePairEncoding, text: &'a [u8], cap: usize) -> Self {
Self {
bpe,
text,
tokens: Vec::with_capacity(cap),
next_token: bpe.next_match(text),
pos: 0,
bitfield: BitField::new(text.len() + 1),
}
}
pub(crate) fn step(&mut self) -> Option<u32> {
let mut token = self.next_token?;
let last = self.tokens.last().copied();
loop {
let token_len = self.bpe.token_len(token);
let end_pos = self.pos + token_len;
if self.bitfield.is_set(end_pos)
&& last
.map(|last_token| self.bpe.is_valid_token_pair(last_token, token))
.unwrap_or(true)
{
self.tokens.push(token);
self.pos = end_pos;
self.next_token = self.bpe.next_match(&self.text[end_pos..]);
break;
} else if let Some(shorter) = self.bpe.next_prefix(token) {
token = shorter;
} else {
self.bitfield.clear(self.pos);
self.tokens.pop();
self.pos -= last.map(|t| self.bpe.token_len(t)).unwrap_or(0);
self.next_token = last;
break;
}
}
self.next_token
}
pub(crate) fn count(&self) -> usize {
self.tokens.len()
}
pub(crate) fn pos(&self) -> usize {
self.pos
}
pub(crate) fn last_token(&self) -> Option<u32> {
self.tokens.last().copied()
}
pub(crate) fn into_tokens(self) -> Vec<u32> {
self.tokens
}
}