use super::parse::{gpt2_regex, parse};
use super::pipeline::PreTokenizer;
use super::spec::{Behavior, PreTokStage, SplitBehavior, SplitPattern};
use super::split::{split_digits, split_punctuation, split_regex};
use super::stage::SplitMatcher;
use serde_json::Value;
fn split<'p>(piece: &'p str, f: impl Fn(&'p str, &mut dyn FnMut(&'p str))) -> Vec<String> {
let mut out = Vec::new();
f(piece, &mut |p| out.push(p));
out.into_iter().map(str::to_owned).collect()
}
#[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 = SplitMatcher::compile(r"\s+").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 = SplitMatcher::compile(r"\s").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_each_match_when_matches_are_contiguous() {
let split_with = |pattern: &str, behavior, text: &'static str| {
let re = SplitMatcher::compile(pattern).expect("compiles");
let mut out = Vec::new();
split_regex(text, &re, behavior, true, &mut |p| out.push(p));
out.into_iter().map(str::to_owned).collect::<Vec<_>>()
};
assert_eq!(
split_with(r"\w+|\s+", Behavior::Removed, "ab cd"),
vec!["ab", " ", "cd"]
);
assert_eq!(
split_with(r"\w+", Behavior::Removed, "ab cd"),
vec!["ab", "cd"]
);
assert_eq!(
split_with(r"\w+", Behavior::Isolated, "ab cd"),
vec!["ab", " ", "cd"]
);
}
#[test]
fn split_invert_partitions_the_whole_piece() {
let re = gpt2_regex().expect("GPT2_PATTERN compiles");
let mut out = Vec::new();
split_regex(
"def f(x):\n return x",
&re,
Behavior::Isolated,
true,
&mut |p| out.push(p),
);
assert_eq!(out.concat(), "def f(x):\n return x");
assert!(out.len() > 1, "a GPT-2 pattern must split this into pieces");
}
#[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}");
}
}