tokenizers/pre_tokenizers/
fixed_length.rs

1use 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}