use crate::core::byte_level::{byte_level_decode, byte_level_decode_bytes};
use crate::core::decode_table::Decoder;
use crate::core::decoder::parse_byte_token;
use crate::core::metaspace::WORD_BOUNDARY;
use rustc_hash::{FxHashMap, FxHashSet};
use std::borrow::Cow;
use std::sync::Arc;
pub(crate) enum Surfaces {
ById(Arc<Decoder>),
ByIndex(Arc<Vec<String>>),
}
pub(crate) enum ByteFallbackRule {
None,
Table(Arc<FxHashMap<u32, u8>>),
ParseSurface,
DeclaredRun,
}
pub(crate) enum Lead {
None,
SpaceUnlessFirst,
}
pub(crate) enum WordSeparator {
None,
EveryToken,
Continuation(String),
}
impl WordSeparator {
#[inline]
fn is_trivial(&self) -> bool {
match self {
Self::None => true,
Self::EveryToken | Self::Continuation(_) => false,
}
}
fn split<'a>(&self, piece: &'a str) -> (Lead, &'a str) {
match self {
Self::None => (Lead::None, piece),
Self::EveryToken => (Lead::SpaceUnlessFirst, piece),
Self::Continuation(marker) => match piece.strip_prefix(marker.as_str()) {
Some(rest) => (Lead::None, rest),
None => (Lead::SpaceUnlessFirst, piece),
},
}
}
}
pub(crate) enum Rendered<'a> {
Skipped,
Bytes { lead: Lead, bytes: Cow<'a, [u8]> },
RunByte(u8),
Unknown,
}
static BYTE_VALUES: [u8; 256] = {
let mut values = [0u8; 256];
let mut b = 0usize;
while b < 256 {
values[b] = b as u8;
b += 1;
}
values
};
pub(crate) struct RenderRules {
surfaces: Surfaces,
special_tokens_decoder: Arc<FxHashMap<u32, String>>,
skip: Arc<FxHashSet<u32>>,
skip_span: Option<(u32, u32)>,
byte_fallback: ByteFallbackRule,
use_byte_level: bool,
surface_replace: Vec<(String, String)>,
unit_cleanup: bool,
word_separator: WordSeparator,
}
#[inline]
fn skip_span(skip: &FxHashSet<u32>) -> Option<(u32, u32)> {
skip.iter().fold(None, |span, &id| match span {
Some((lo, hi)) => Some((lo.min(id), hi.max(id))),
None => Some((id, id)),
})
}
impl RenderRules {
pub(crate) fn new(
surfaces: Surfaces,
special_tokens_decoder: Arc<FxHashMap<u32, String>>,
skip: Arc<FxHashSet<u32>>,
byte_fallback: ByteFallbackRule,
use_byte_level: bool,
use_metaspace: bool,
) -> Self {
let skip_span = skip_span(&skip);
let mut rules = Self {
surfaces,
special_tokens_decoder,
skip,
skip_span,
byte_fallback,
use_byte_level,
surface_replace: Vec::new(),
unit_cleanup: false,
word_separator: WordSeparator::None,
};
if use_metaspace {
rules = rules.with_surface_replace(WORD_BOUNDARY.to_string(), " ".to_string());
}
rules
}
pub(crate) fn declared(byte_fallback: ByteFallbackRule, use_byte_level: bool) -> Self {
Self {
surfaces: Surfaces::ByIndex(Arc::new(Vec::new())),
special_tokens_decoder: Arc::new(FxHashMap::default()),
skip: Arc::new(FxHashSet::default()),
skip_span: None,
byte_fallback,
use_byte_level,
surface_replace: Vec::new(),
unit_cleanup: false,
word_separator: WordSeparator::None,
}
}
pub(crate) fn with_vocabulary(
mut self,
surfaces: Surfaces,
special_tokens_decoder: Arc<FxHashMap<u32, String>>,
skip: Arc<FxHashSet<u32>>,
) -> Self {
self.surfaces = surfaces;
self.special_tokens_decoder = special_tokens_decoder;
self.skip_span = skip_span(&skip);
self.skip = skip;
self
}
pub(crate) fn rendering_specials(mut self) -> Self {
self.skip = Arc::new(FxHashSet::default());
self.skip_span = None;
self
}
pub(crate) fn with_surface_replace(mut self, from: String, to: String) -> Self {
self.surface_replace.push((from, to));
self
}
pub(crate) fn with_unit_cleanup(mut self) -> Self {
self.unit_cleanup = true;
self
}
#[inline]
pub(crate) fn unit_cleanup(&self) -> bool {
self.unit_cleanup
}
pub(crate) fn with_word_separator(mut self, word_separator: WordSeparator) -> Self {
self.word_separator = word_separator;
self
}
#[inline]
fn declared_run_byte(&self, surface: &[u8]) -> Option<u8> {
match self.byte_fallback {
ByteFallbackRule::DeclaredRun => {
std::str::from_utf8(surface).ok().and_then(parse_byte_token)
}
ByteFallbackRule::ParseSurface
| ByteFallbackRule::Table(_)
| ByteFallbackRule::None => None,
}
}
#[inline]
fn surface_replace_applies(&self, piece: &str) -> bool {
self.surface_replace
.iter()
.any(|(from, _)| piece.contains(from.as_str()))
}
pub(crate) fn token_bytes(&self, id: u32) -> Option<Vec<u8>> {
match self.render(id) {
Rendered::Skipped => Some(Vec::new()),
Rendered::Bytes { lead: _, bytes } => Some(bytes.into_owned()),
Rendered::RunByte(byte) => Some(vec![byte]),
Rendered::Unknown => None,
}
}
#[inline]
pub(crate) fn skips(&self, id: u32) -> bool {
match self.skip_span {
Some((lo, hi)) => (lo..=hi).contains(&id) && self.skip.contains(&id),
None => false,
}
}
#[inline]
pub(crate) fn special_surface(&self, id: u32) -> Option<&str> {
self.special_tokens_decoder.get(&id).map(String::as_str)
}
pub(crate) fn plain_by_id(&self) -> Option<&Decoder> {
let map = match &self.surfaces {
Surfaces::ById(map) => map,
Surfaces::ByIndex(_) => return None,
};
let plain = matches!(self.byte_fallback, ByteFallbackRule::None)
&& !self.use_byte_level
&& !self.unit_cleanup
&& self.surface_replace.is_empty()
&& self.word_separator.is_trivial();
plain.then_some(map.as_ref())
}
#[inline]
pub(crate) fn render(&self, id: u32) -> Rendered<'_> {
if self.skips(id) {
return Rendered::Skipped;
}
if let ByteFallbackRule::Table(table) = &self.byte_fallback {
if let Some(&byte) = table.get(&id) {
let b = byte as usize;
return Rendered::Bytes {
lead: Lead::None,
bytes: Cow::Borrowed(&BYTE_VALUES[b..b + 1]),
};
}
}
match &self.surfaces {
Surfaces::ById(map) => {
if let Some(bytes) = map.get(id) {
if let Some(byte) = self.declared_run_byte(bytes) {
return Rendered::RunByte(byte);
}
let bytes = if self.use_byte_level {
match byte_level_decode_bytes(bytes) {
Some(decoded) => Cow::Owned(decoded),
None => Cow::Borrowed(bytes),
}
} else {
Cow::Borrowed(bytes)
};
return Rendered::Bytes {
lead: Lead::None,
bytes,
};
}
}
Surfaces::ByIndex(pieces) => {
if let Some(piece) = pieces.get(id as usize) {
let parsed = match self.byte_fallback {
ByteFallbackRule::ParseSurface => parse_byte_token(piece),
ByteFallbackRule::DeclaredRun
| ByteFallbackRule::Table(_)
| ByteFallbackRule::None => None,
};
if let Some(byte) = parsed {
let b = byte as usize;
return Rendered::Bytes {
lead: Lead::None,
bytes: Cow::Borrowed(&BYTE_VALUES[b..b + 1]),
};
}
if let Some(byte) = self.declared_run_byte(piece.as_bytes()) {
return Rendered::RunByte(byte);
}
let (lead, piece) = self.word_separator.split(piece);
let bytes = if self.use_byte_level {
match byte_level_decode(piece) {
Some(decoded) => Cow::Owned(decoded),
None => Cow::Borrowed(piece.as_bytes()),
}
} else if self.surface_replace_applies(piece) {
let replaced = self
.surface_replace
.iter()
.fold(piece.to_string(), |text, (from, to)| {
text.replace(from.as_str(), to.as_str())
});
Cow::Owned(replaced.into_bytes())
} else {
Cow::Borrowed(piece.as_bytes())
};
return Rendered::Bytes { lead, bytes };
}
}
}
match self.special_surface(id) {
Some(special) => Rendered::Bytes {
lead: Lead::None,
bytes: Cow::Borrowed(special.as_bytes()),
},
None => Rendered::Unknown,
}
}
}
#[cfg(test)]
mod tests {
use super::{ByteFallbackRule, RenderRules, Surfaces, WordSeparator};
use rustc_hash::{FxHashMap, FxHashSet};
use std::sync::Arc;
fn plain(byte_fallback: ByteFallbackRule, use_byte_level: bool) -> RenderRules {
let mut surfaces = crate::core::DecodeTable::default();
surfaces.insert(1u32, b"Hi");
RenderRules::new(
Surfaces::ById(Arc::new(surfaces)),
Arc::new(FxHashMap::default()),
Arc::new(FxHashSet::default()),
byte_fallback,
use_byte_level,
false,
)
}
#[test]
fn plain_by_id_accepts_the_plain_bpe_shape() {
let rules = plain(ByteFallbackRule::None, false);
let map = rules.plain_by_id().expect("the plain BPE shape qualifies");
assert_eq!(map.get(1), Some(b"Hi".as_slice()));
}
#[test]
fn plain_by_id_rejects_by_index_surfaces() {
let rules = RenderRules::new(
Surfaces::ByIndex(Arc::new(vec!["Hi".to_string()])),
Arc::new(FxHashMap::default()),
Arc::new(FxHashSet::default()),
ByteFallbackRule::None,
false,
false,
);
assert!(rules.plain_by_id().is_none());
}
#[test]
fn plain_by_id_rejects_byte_fallback() {
let table = Arc::new(FxHashMap::default());
assert!(plain(ByteFallbackRule::Table(table), false)
.plain_by_id()
.is_none());
assert!(plain(ByteFallbackRule::ParseSurface, false)
.plain_by_id()
.is_none());
assert!(plain(ByteFallbackRule::DeclaredRun, false)
.plain_by_id()
.is_none());
}
#[test]
fn plain_by_id_rejects_byte_level() {
assert!(plain(ByteFallbackRule::None, true).plain_by_id().is_none());
}
#[test]
fn plain_by_id_rejects_surface_replace() {
let rules = plain(ByteFallbackRule::None, false)
.with_surface_replace("a".to_string(), "b".to_string());
assert!(rules.plain_by_id().is_none());
}
#[test]
fn plain_by_id_rejects_unit_cleanup() {
let rules = plain(ByteFallbackRule::None, false).with_unit_cleanup();
assert!(rules.plain_by_id().is_none());
}
#[test]
fn plain_by_id_rejects_word_separator() {
let every =
plain(ByteFallbackRule::None, false).with_word_separator(WordSeparator::EveryToken);
assert!(every.plain_by_id().is_none());
let marked = plain(ByteFallbackRule::None, false)
.with_word_separator(WordSeparator::Continuation("##".to_string()));
assert!(marked.plain_by_id().is_none());
}
#[test]
fn skips_is_answer_preserving_across_a_sparse_span() {
let mut skip = FxHashSet::default();
skip.insert(5u32);
skip.insert(9000u32);
let rules = RenderRules::new(
Surfaces::ById(Arc::new(crate::core::DecodeTable::default())),
Arc::new(FxHashMap::default()),
Arc::new(skip),
ByteFallbackRule::None,
false,
false,
);
assert!(rules.skips(5));
assert!(rules.skips(9000));
assert!(!rules.skips(6));
assert!(!rules.skips(4999));
assert!(!rules.skips(0));
assert!(!rules.skips(4));
assert!(!rules.skips(9001));
assert!(!rules.skips(u32::MAX));
}
#[test]
fn skips_is_false_for_an_empty_skip_set() {
let rules = plain(ByteFallbackRule::None, false);
assert!(!rules.skips(0));
assert!(!rules.skips(1));
assert!(!rules.skips(9000));
assert!(!rules.skips(u32::MAX));
}
#[test]
fn skips_is_correct_after_with_vocabulary_replaces_skip() {
let mut first_skip = FxHashSet::default();
first_skip.insert(5u32);
let rules = RenderRules::declared(ByteFallbackRule::None, false).with_vocabulary(
Surfaces::ById(Arc::new(crate::core::DecodeTable::default())),
Arc::new(FxHashMap::default()),
Arc::new(first_skip),
);
assert!(rules.skips(5));
assert!(!rules.skips(9000));
let mut second_skip = FxHashSet::default();
second_skip.insert(9000u32);
let rules = rules.with_vocabulary(
Surfaces::ById(Arc::new(crate::core::DecodeTable::default())),
Arc::new(FxHashMap::default()),
Arc::new(second_skip),
);
assert!(!rules.skips(5));
assert!(rules.skips(9000));
}
#[test]
fn skips_is_false_after_rendering_specials_clears_skip() {
let mut skip = FxHashSet::default();
skip.insert(5u32);
let rules = RenderRules::new(
Surfaces::ById(Arc::new(crate::core::DecodeTable::default())),
Arc::new(FxHashMap::default()),
Arc::new(skip),
ByteFallbackRule::None,
false,
false,
)
.rendering_specials();
assert!(!rules.skips(5));
assert!(!rules.skips(0));
assert!(!rules.skips(u32::MAX));
}
}