tokenizers/pre_tokenizers/
fixed_length.rs1use crate::normalizer::Range;
2use crate::tokenizer::{PreTokenizedString, PreTokenizer, Result};
3use serde::{Deserialize, Serialize};
4
5use crate::utils::macro_rules_attribute;
6
7#[derive(Clone, Debug, PartialEq, Eq)]
8#[macro_rules_attribute(impl_serde_type!)]
9pub struct FixedLength {
10 #[serde(default = "default_length")]
11 pub length: usize,
12}
13
14impl FixedLength {
15 pub fn new(length: usize) -> Self {
16 Self { length }
17 }
18}
19
20fn default_length() -> usize {
21 5
22}
23
24impl PreTokenizer for FixedLength {
25 fn pre_tokenize(&self, pretokenized: &mut PreTokenizedString) -> Result<()> {
26 pretokenized.split(|_, normalized| {
27 let text = normalized.get();
28 if text.is_empty() {
29 return Ok(vec![]);
30 }
31
32 let mut splits = Vec::new();
33 let char_positions: Vec<_> = text.char_indices().collect();
34 for chunk in char_positions.chunks(self.length) {
35 let start = chunk.first().map(|(i, _)| *i).unwrap_or(0);
36 let end = chunk
37 .last()
38 .map(|(i, c)| i + c.len_utf8())
39 .unwrap_or(text.len());
40 splits.push(
41 normalized
42 .slice(Range::Normalized(start..end))
43 .ok_or("Failed to slice normalized text")?,
44 );
45 }
46
47 Ok(splits)
48 })
49 }
50}
51
52#[cfg(test)]
53mod tests {
54 use super::*;
55 use crate::{OffsetReferential, OffsetType, PreTokenizer};
56
57 #[test]
58 fn basic() {
59 let tests = vec![
60 (
61 "Hello world",
62 vec![("Hello", (0, 5)), (" worl", (5, 10)), ("d", (10, 11))],
63 ),
64 ("Short", vec![("Short", (0, 5))]),
65 ("", vec![]),
66 ];
67 let pretok = FixedLength { length: 5 };
68 for (s, res) in tests {
69 let mut pretokenized = PreTokenizedString::from(s);
70 pretok.pre_tokenize(&mut pretokenized).unwrap();
71 assert_eq!(
72 pretokenized
73 .get_splits(OffsetReferential::Original, OffsetType::Byte)
74 .into_iter()
75 .map(|(s, o, _)| (s, o))
76 .collect::<Vec<_>>(),
77 res
78 );
79 }
80 }
81
82 #[test]
83 fn custom_length() {
84 let pretok = FixedLength { length: 3 };
85 let mut pretokenized = PreTokenizedString::from("Hello world");
86 pretok.pre_tokenize(&mut pretokenized).unwrap();
87 assert_eq!(
88 pretokenized
89 .get_splits(OffsetReferential::Original, OffsetType::Byte)
90 .into_iter()
91 .map(|(s, o, _)| (s, o))
92 .collect::<Vec<_>>(),
93 vec![
94 ("Hel", (0, 3)),
95 ("lo ", (3, 6)),
96 ("wor", (6, 9)),
97 ("ld", (9, 11)),
98 ]
99 );
100 }
101
102 #[test]
103 fn utf8_characters() {
104 let pretok = FixedLength { length: 3 };
105 let mut pretokenized = PreTokenizedString::from("Hello 👋 world");
106 pretok.pre_tokenize(&mut pretokenized).unwrap();
107 assert_eq!(
108 pretokenized
109 .get_splits(OffsetReferential::Original, OffsetType::Byte)
110 .into_iter()
111 .map(|(s, o, _)| (s, o))
112 .collect::<Vec<_>>(),
113 vec![
114 ("Hel", (0, 3)),
115 ("lo ", (3, 6)),
116 ("👋 w", (6, 12)),
117 ("orl", (12, 15)),
118 ("d", (15, 16)),
119 ]
120 );
121 }
122}