use std::borrow::Cow;
use regexr::RegexBuilder;
use serde_json::Value;
use super::byte_level::byte_level_encode;
use super::tokenizer::{TokenizerError, GPT2_PATTERN};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum SplitBehavior {
#[default]
Isolated,
Removed,
MergedWithPrevious,
MergedWithNext,
Contiguous,
}
impl SplitBehavior {
fn parse(s: Option<&str>) -> Self {
match s {
Some("Removed") => SplitBehavior::Removed,
Some("MergedWithPrevious") => SplitBehavior::MergedWithPrevious,
Some("MergedWithNext") => SplitBehavior::MergedWithNext,
Some("Contiguous") => SplitBehavior::Contiguous,
_ => SplitBehavior::Isolated,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum SplitPattern {
Literal(String),
Regex(String),
}
#[non_exhaustive]
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum PreTokStage {
Split {
pattern: SplitPattern,
behavior: SplitBehavior,
invert: bool,
},
ByteLevel {
use_regex: bool,
add_prefix_space: bool,
},
Digits { individual: bool },
Punctuation { behavior: SplitBehavior },
WhitespaceSplit,
Whitespace,
}
#[derive(Debug, Clone, Copy, PartialEq)]
enum Behavior {
Isolated,
Removed,
MergedWithPrevious,
MergedWithNext,
Contiguous,
}
impl From<SplitBehavior> for Behavior {
fn from(b: SplitBehavior) -> Self {
match b {
SplitBehavior::Isolated => Behavior::Isolated,
SplitBehavior::Removed => Behavior::Removed,
SplitBehavior::MergedWithPrevious => Behavior::MergedWithPrevious,
SplitBehavior::MergedWithNext => Behavior::MergedWithNext,
SplitBehavior::Contiguous => Behavior::Contiguous,
}
}
}
enum Stage {
Split {
re: Box<regexr::Regex>,
behavior: Behavior,
invert: bool,
},
ByteLevel { re: Option<Box<regexr::Regex>> },
Digits { individual: bool },
Punctuation { behavior: Behavior },
WhitespaceSplit,
Whitespace { re: Box<regexr::Regex> },
}
pub struct PreTokenizer {
spec: Vec<PreTokStage>,
compiled: Vec<Stage>,
add_prefix_space: bool,
byte_level: bool,
}
impl PreTokenizer {
pub fn new(stages: Vec<PreTokStage>) -> Result<Self, TokenizerError> {
let mut compiled = Vec::with_capacity(stages.len());
let mut byte_level = false;
let mut add_prefix_space = false;
for stage in &stages {
compiled.push(match stage {
PreTokStage::Split {
pattern,
behavior,
invert,
} => Stage::Split {
re: Box::new(
RegexBuilder::new(&match pattern {
SplitPattern::Literal(s) => Cow::Owned(regexr::escape(s)),
SplitPattern::Regex(s) => Cow::Borrowed(s.as_str()),
})
.jit(true)
.build()?,
),
behavior: (*behavior).into(),
invert: *invert,
},
PreTokStage::ByteLevel {
use_regex,
add_prefix_space: prefix,
} => {
byte_level = true;
add_prefix_space |= *prefix;
Stage::ByteLevel {
re: match use_regex {
true => Some(Box::new(gpt2_regex()?)),
false => None,
},
}
}
PreTokStage::Digits { individual } => Stage::Digits {
individual: *individual,
},
PreTokStage::Punctuation { behavior } => Stage::Punctuation {
behavior: (*behavior).into(),
},
PreTokStage::WhitespaceSplit => Stage::WhitespaceSplit,
PreTokStage::Whitespace => Stage::Whitespace {
re: Box::new(whitespace_regex()?),
},
});
}
Ok(Self {
spec: stages,
compiled,
add_prefix_space,
byte_level,
})
}
pub fn split(&self, text: &str) -> Vec<String> {
let mut pieces: Vec<String> = if self.add_prefix_space && !text.starts_with(' ') {
vec![format!(" {text}")]
} else {
vec![text.to_string()]
};
for stage in &self.compiled {
let mut next = Vec::with_capacity(pieces.len());
for p in &pieces {
stage.apply(p, &mut next);
}
pieces = next;
}
pieces.retain(|p| !p.is_empty());
pieces
}
pub fn byte_level(&self) -> bool {
self.byte_level
}
pub fn is_empty(&self) -> bool {
self.spec.is_empty()
}
pub fn stages(&self) -> &[PreTokStage] {
&self.spec
}
}
impl std::fmt::Debug for PreTokenizer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PreTokenizer")
.field("stages", &self.spec)
.field("byte_level", &self.byte_level)
.finish()
}
}
impl Stage {
fn apply(&self, piece: &str, out: &mut Vec<String>) {
match self {
Stage::Split {
re,
behavior,
invert,
} => split_regex(piece, re, *behavior, *invert, out),
Stage::ByteLevel { re } => match re {
Some(re) => {
let mut raw = Vec::new();
split_regex(piece, re, Behavior::Isolated, false, &mut raw);
for r in raw {
out.push(byte_level_encode(r.as_bytes()));
}
}
None => out.push(byte_level_encode(piece.as_bytes())),
},
Stage::Digits { individual } => split_digits(piece, *individual, out),
Stage::Punctuation { behavior } => split_punctuation(piece, *behavior, out),
Stage::WhitespaceSplit => {
for w in piece.split_whitespace() {
out.push(w.to_string());
}
}
Stage::Whitespace { re } => split_regex(piece, re, Behavior::Isolated, false, out),
}
}
}
fn split_regex(
piece: &str,
re: ®exr::Regex,
behavior: Behavior,
invert: bool,
out: &mut Vec<String>,
) {
let matches: Vec<(usize, usize)> = re.find_iter(piece).map(|m| (m.start(), m.end())).collect();
let delims: Vec<(usize, usize)> = if invert {
let mut d = Vec::new();
let mut last = 0;
for &(s, e) in &matches {
if s > last {
d.push((last, s));
}
last = e;
}
if last < piece.len() {
d.push((last, piece.len()));
}
d
} else {
matches
};
let mut segs: Vec<(&str, bool)> = Vec::with_capacity(delims.len() * 2 + 1);
let mut last = 0;
for (s, e) in delims {
if s > last {
segs.push((&piece[last..s], false));
}
if e > s {
segs.push((&piece[s..e], true));
}
last = e;
}
if last < piece.len() {
segs.push((&piece[last..], false));
}
emit_segments(&segs, behavior, out);
}
fn emit_segments(segs: &[(&str, bool)], behavior: Behavior, out: &mut Vec<String>) {
match behavior {
Behavior::Isolated => {
for &(t, _) in segs {
out.push(t.to_string());
}
}
Behavior::Removed => {
for &(t, d) in segs {
if !d {
out.push(t.to_string());
}
}
}
Behavior::MergedWithPrevious => {
let mut local: Vec<String> = Vec::new();
for &(t, d) in segs {
if d {
if let Some(prev) = local.last_mut() {
prev.push_str(t);
} else {
local.push(t.to_string());
}
} else {
local.push(t.to_string());
}
}
out.extend(local);
}
Behavior::MergedWithNext => {
let mut pending = String::new();
for &(t, d) in segs {
if d {
pending.push_str(t);
} else {
let mut s = std::mem::take(&mut pending);
s.push_str(t);
out.push(s);
}
}
if !pending.is_empty() {
out.push(pending);
}
}
Behavior::Contiguous => {
let mut run = String::new();
for &(t, d) in segs {
if d {
run.push_str(t);
} else {
if !run.is_empty() {
out.push(std::mem::take(&mut run));
}
out.push(t.to_string());
}
}
if !run.is_empty() {
out.push(run);
}
}
}
}
fn split_digits(piece: &str, individual: bool, out: &mut Vec<String>) {
let mut cur = String::new();
let mut cur_is_digit = false;
for c in piece.chars() {
let d = c.is_numeric();
if !cur.is_empty() && (d != cur_is_digit || (d && individual)) {
out.push(std::mem::take(&mut cur));
}
cur.push(c);
cur_is_digit = d;
if d && individual {
out.push(std::mem::take(&mut cur));
}
}
if !cur.is_empty() {
out.push(cur);
}
}
fn split_punctuation(piece: &str, behavior: Behavior, out: &mut Vec<String>) {
let mut segs: Vec<(&str, bool)> = Vec::new();
let mut content_start = 0;
let mut i = 0;
for c in piece.chars() {
let len = c.len_utf8();
if is_punctuation(c) {
if i > content_start {
segs.push((&piece[content_start..i], false));
}
segs.push((&piece[i..i + len], true));
content_start = i + len;
}
i += len;
}
if i > content_start {
segs.push((&piece[content_start..i], false));
}
emit_segments(&segs, behavior, out);
}
fn is_punctuation(c: char) -> bool {
if c.is_ascii() {
return c.is_ascii_punctuation();
}
use unicode_general_category::{get_general_category, GeneralCategory::*};
matches!(
get_general_category(c),
ConnectorPunctuation
| DashPunctuation
| ClosePunctuation
| FinalPunctuation
| InitialPunctuation
| OtherPunctuation
| OpenPunctuation
)
}
fn gpt2_regex() -> Result<regexr::Regex, TokenizerError> {
Ok(RegexBuilder::new(GPT2_PATTERN).jit(true).build()?)
}
fn whitespace_regex() -> Result<regexr::Regex, TokenizerError> {
Ok(RegexBuilder::new(r"\w+|[^\w\s]+").jit(true).build()?)
}
pub(crate) fn parse(pre: Option<&Value>) -> Result<Option<PreTokenizer>, TokenizerError> {
let Some(pre) = pre else {
return Ok(None);
};
let mut stages = Vec::new();
fn walk(v: &Value, stages: &mut Vec<PreTokStage>) {
match v.get("type").and_then(Value::as_str) {
Some("Sequence") => {
if let Some(list) = v.get("pretokenizers").and_then(Value::as_array) {
for item in list {
walk(item, stages);
}
}
}
Some("ByteLevel") => stages.push(PreTokStage::ByteLevel {
use_regex: v.get("use_regex").and_then(Value::as_bool).unwrap_or(true),
add_prefix_space: v.get("add_prefix_space").and_then(Value::as_bool) == Some(true),
}),
Some("Split") => {
let pat = v.get("pattern").and_then(|p| {
p.get("Regex")
.and_then(Value::as_str)
.map(|s| SplitPattern::Regex(s.to_string()))
.or_else(|| {
p.get("String")
.and_then(Value::as_str)
.map(|s| SplitPattern::Literal(s.to_string()))
})
});
if let Some(pat) = pat {
stages.push(PreTokStage::Split {
pattern: pat,
behavior: SplitBehavior::parse(v.get("behavior").and_then(Value::as_str)),
invert: v.get("invert").and_then(Value::as_bool).unwrap_or(false),
});
}
}
Some("Digits") => stages.push(PreTokStage::Digits {
individual: v
.get("individual_digits")
.and_then(Value::as_bool)
.unwrap_or(false),
}),
Some("Punctuation") => stages.push(PreTokStage::Punctuation {
behavior: SplitBehavior::parse(v.get("behavior").and_then(Value::as_str)),
}),
Some("WhitespaceSplit") => stages.push(PreTokStage::WhitespaceSplit),
Some("Whitespace") => stages.push(PreTokStage::Whitespace),
_ => {}
}
}
walk(pre, &mut stages);
if stages.is_empty() {
return Ok(None);
}
Ok(Some(PreTokenizer::new(stages)?))
}
#[cfg(test)]
mod tests {
use super::*;
fn split(piece: &str, f: impl Fn(&str, &mut Vec<String>)) -> Vec<String> {
let mut out = Vec::new();
f(piece, &mut out);
out
}
#[test]
fn digits_grouped_and_individual() {
assert_eq!(
split("abc123def", |p, o| split_digits(p, false, o)),
vec!["abc", "123", "def"]
);
assert_eq!(
split("abc123", |p, o| split_digits(p, true, o)),
vec!["abc", "1", "2", "3"]
);
assert_eq!(
split("a²½b", |p, o| split_digits(p, true, o)),
vec!["a", "²", "½", "b"]
);
}
#[test]
fn punctuation_isolated_and_contiguous() {
assert_eq!(
split("a,b!", |p, o| split_punctuation(p, Behavior::Isolated, o)),
vec!["a", ",", "b", "!"]
);
assert_eq!(
split("a)=b", |p, o| split_punctuation(p, Behavior::Contiguous, o)),
vec!["a", ")=", "b"]
);
assert_eq!(
split("a,b", |p, o| split_punctuation(p, Behavior::Removed, o)),
vec!["a", "b"]
);
}
#[test]
fn split_merge_behaviors() {
let re = RegexBuilder::new(r"\s+").jit(true).build().unwrap();
let go = |b| split("a b c", |p, o| split_regex(p, &re, b, false, o));
assert_eq!(go(Behavior::Isolated), vec!["a", " ", "b", " ", "c"]);
assert_eq!(go(Behavior::Removed), vec!["a", "b", "c"]);
assert_eq!(go(Behavior::MergedWithPrevious), vec!["a ", "b ", "c"]);
assert_eq!(go(Behavior::MergedWithNext), vec!["a", " b", " c"]);
let re2 = RegexBuilder::new(r"\s").jit(true).build().unwrap();
assert_eq!(
split("a b", |p, o| split_regex(
p,
&re2,
Behavior::Contiguous,
false,
o
)),
vec!["a", " ", "b"]
);
}
#[test]
fn pipeline_digits_then_byte_level() {
let json = serde_json::json!({
"type": "Sequence",
"pretokenizers": [
{"type": "Digits", "individual_digits": true},
{"type": "ByteLevel", "add_prefix_space": false, "use_regex": true}
]
});
let pt = parse(Some(&json)).expect("parses").expect("pipeline");
assert!(pt.byte_level());
let pieces = pt.split("a12");
assert_eq!(pieces, vec!["a", "1", "2"]);
}
#[test]
fn split_invert_keeps_content_between_matches() {
let re = Box::new(gpt2_regex().expect("GPT2_PATTERN compiles"));
let mut out = Vec::new();
super::split_regex("ab", &re, Behavior::Isolated, true, &mut out);
assert_eq!(out.concat(), "ab");
}
#[test]
fn literal_and_regex_split_patterns_are_not_interchangeable() {
let split = |pattern| {
PreTokenizer::new(vec![PreTokStage::Split {
pattern,
behavior: SplitBehavior::Removed,
invert: false,
}])
.expect("pipeline builds")
.split("a.b c")
};
assert_eq!(
split(SplitPattern::Literal(".".to_string())),
vec!["a", "b c"]
);
assert!(split(SplitPattern::Regex(".".to_string())).is_empty());
}
#[test]
fn literal_split_pattern_matches_metacharacters_verbatim() {
let split = |pattern, text: &str| {
PreTokenizer::new(vec![PreTokStage::Split {
pattern,
behavior: SplitBehavior::Removed,
invert: false,
}])
.expect("pipeline builds")
.split(text)
};
assert_eq!(
split(SplitPattern::Literal("a+b".to_string()), "xa+by"),
vec!["x", "y"]
);
assert_eq!(
split(SplitPattern::Literal("|".to_string()), "a|b"),
vec!["a", "b"]
);
}
#[test]
fn split_with_uncompilable_pattern_is_an_error() {
let json = serde_json::json!({
"type": "Split",
"pattern": {"Regex": "("},
"behavior": "Isolated"
});
assert!(parse(Some(&json)).is_err());
}
#[test]
fn parsed_stages_round_trip_through_the_public_builder() {
let probe = "Hello, wörld 42 items!";
let mut cases: Vec<(Value, Vec<PreTokStage>)> = vec![
(
serde_json::json!({
"type": "Sequence",
"pretokenizers": [
{"type": "Sequence", "pretokenizers": [
{"type": "Punctuation", "behavior": "Contiguous"},
{"type": "Digits", "individual_digits": true}
]},
{"type": "ByteLevel", "use_regex": false, "add_prefix_space": true}
]
}),
vec![
PreTokStage::Punctuation {
behavior: SplitBehavior::Contiguous,
},
PreTokStage::Digits { individual: true },
PreTokStage::ByteLevel {
use_regex: false,
add_prefix_space: true,
},
],
),
(
serde_json::json!({"type": "ByteLevel"}),
vec![PreTokStage::ByteLevel {
use_regex: true,
add_prefix_space: false,
}],
),
(
serde_json::json!({"type": "Whitespace"}),
vec![PreTokStage::Whitespace],
),
(
serde_json::json!({"type": "WhitespaceSplit"}),
vec![PreTokStage::WhitespaceSplit],
),
(
serde_json::json!({
"type": "Split",
"pattern": {"String": ","},
"behavior": "Removed"
}),
vec![PreTokStage::Split {
pattern: SplitPattern::Literal(",".to_string()),
behavior: SplitBehavior::Removed,
invert: false,
}],
),
(
serde_json::json!({
"type": "Split",
"pattern": {"String": "."},
"behavior": "Removed"
}),
vec![PreTokStage::Split {
pattern: SplitPattern::Literal(".".to_string()),
behavior: SplitBehavior::Removed,
invert: false,
}],
),
(
serde_json::json!({
"type": "Split",
"pattern": {"Regex": r"\w+"},
"invert": true
}),
vec![PreTokStage::Split {
pattern: SplitPattern::Regex(r"\w+".to_string()),
behavior: SplitBehavior::Isolated,
invert: true,
}],
),
];
for (name, behavior) in [
("Isolated", SplitBehavior::Isolated),
("Removed", SplitBehavior::Removed),
("MergedWithPrevious", SplitBehavior::MergedWithPrevious),
("MergedWithNext", SplitBehavior::MergedWithNext),
("Contiguous", SplitBehavior::Contiguous),
] {
cases.push((
serde_json::json!({
"type": "Split",
"pattern": {"Regex": r"\s+"},
"behavior": name
}),
vec![PreTokStage::Split {
pattern: SplitPattern::Regex(r"\s+".to_string()),
behavior,
invert: false,
}],
));
}
for (json, expected) in cases {
let parsed = parse(Some(&json)).expect("parses").expect("pipeline");
assert_eq!(parsed.stages(), expected.as_slice(), "spec for {json}");
let built = PreTokenizer::new(expected).expect("builds");
assert_eq!(built.byte_level(), parsed.byte_level(), "byte_level {json}");
assert_eq!(built.split(probe), parsed.split(probe), "split for {json}");
}
}
}