tokenizers/processors/
bert.rs

1use crate::tokenizer::{Encoding, PostProcessor, Result};
2use ahash::AHashMap;
3use serde::{Deserialize, Serialize};
4use std::iter::FromIterator;
5
6#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, Eq)]
7#[serde(tag = "type")]
8pub struct BertProcessing {
9    pub sep: (String, u32),
10    pub cls: (String, u32),
11}
12
13impl Default for BertProcessing {
14    fn default() -> Self {
15        Self {
16            sep: ("[SEP]".into(), 102),
17            cls: ("[CLS]".into(), 101),
18        }
19    }
20}
21
22impl BertProcessing {
23    pub fn new(sep: (String, u32), cls: (String, u32)) -> Self {
24        Self { sep, cls }
25    }
26
27    pub fn get_sep_copy(&self) -> (String, u32) {
28        (self.sep.0.clone(), self.sep.1)
29    }
30
31    pub fn get_cls_copy(&self) -> (String, u32) {
32        (self.cls.0.clone(), self.cls.1)
33    }
34}
35
36#[derive(thiserror::Error, Debug)]
37pub enum BertProcessorError {
38    #[error("encodings vector length must be either 1 or 2")]
39    InvalidEncodingsVecLength,
40}
41
42impl PostProcessor for BertProcessing {
43    fn added_tokens(&self, is_pair: bool) -> usize {
44        if is_pair {
45            3
46        } else {
47            2
48        }
49    }
50
51    fn process_encodings(
52        &self,
53        mut encodings: Vec<Encoding>,
54        add_special_tokens: bool,
55    ) -> Result<Vec<Encoding>> {
56        if !add_special_tokens {
57            return Ok(encodings);
58        }
59
60        let encodings: Vec<Encoding> = encodings
61            .iter_mut()
62            .enumerate()
63            .map(|(i, encoding)| {
64                if i == 0 {
65                    let ids = [&[self.cls.1], encoding.get_ids(), &[self.sep.1]].concat();
66                    let type_ids = [&[0], encoding.get_type_ids(), &[0]].concat();
67                    let tokens = [
68                        std::slice::from_ref(&self.cls.0),
69                        encoding.get_tokens(),
70                        std::slice::from_ref(&self.sep.0),
71                    ]
72                    .concat();
73                    let words = [&[None], encoding.get_word_ids(), &[None]].concat();
74                    let offsets = [&[(0, 0)], encoding.get_offsets(), &[(0, 0)]].concat();
75                    let special_tokens =
76                        [&[1u32], &vec![0; encoding.get_ids().len()][..], &[1]].concat();
77                    let attention_mask = vec![1; ids.len()];
78
79                    // For compatibility with `TemplateProcessing`, the sequence_ranges shouldn't contain
80                    // the special tokens.
81                    let sequence_ranges = AHashMap::from_iter(vec![(0, 1..ids.len() - 1)]);
82                    Encoding::new(
83                        ids,
84                        type_ids,
85                        tokens,
86                        words,
87                        offsets,
88                        special_tokens,
89                        attention_mask,
90                        encoding
91                            .take_overflowing()
92                            .into_iter()
93                            .map(|encoding| {
94                                let ids =
95                                    [&[self.cls.1], encoding.get_ids(), &[self.sep.1]].concat();
96                                let type_ids = [&[0], encoding.get_type_ids(), &[0]].concat();
97                                let tokens = [
98                                    std::slice::from_ref(&self.cls.0),
99                                    encoding.get_tokens(),
100                                    std::slice::from_ref(&self.sep.0),
101                                ]
102                                .concat();
103                                let words = [&[None], encoding.get_word_ids(), &[None]].concat();
104                                let offsets =
105                                    [&[(0, 0)], encoding.get_offsets(), &[(0, 0)]].concat();
106                                let special_tokens =
107                                    [&[1u32], &vec![0; encoding.get_ids().len()][..], &[1]]
108                                        .concat();
109                                let attention_mask = vec![1; ids.len()];
110
111                                // For compatibility with `TemplateProcessing`, the sequence_ranges shouldn't
112                                // contain the special tokens.
113                                let sequence_ranges =
114                                    AHashMap::from_iter(vec![(0, 1..ids.len() - 1)]);
115                                Encoding::new(
116                                    ids,
117                                    type_ids,
118                                    tokens,
119                                    words,
120                                    offsets,
121                                    special_tokens,
122                                    attention_mask,
123                                    vec![],
124                                    sequence_ranges,
125                                )
126                            })
127                            .collect(),
128                        sequence_ranges,
129                    )
130                } else {
131                    let pair_ids = [encoding.get_ids(), &[self.sep.1]].concat();
132                    let pair_type_ids = [encoding.get_type_ids(), &[1]].concat();
133                    let pair_tokens =
134                        [encoding.get_tokens(), std::slice::from_ref(&self.sep.0)].concat();
135                    let pair_words = [encoding.get_word_ids(), &[None]].concat();
136                    let pair_offsets = [encoding.get_offsets(), &[(0, 0)]].concat();
137                    let pair_special_tokens =
138                        [&vec![0u32; encoding.get_type_ids().len()][..], &[1]].concat();
139                    let pair_attention_mask = vec![1; pair_ids.len()];
140
141                    // For compatibility with `TemplateProcessing`, the sequence_ranges shouldn't contain
142                    // the special tokens.
143                    let pair_sequence_ranges =
144                        AHashMap::from_iter(vec![(1, 0..pair_ids.len() - 1)]);
145                    Encoding::new(
146                        pair_ids,
147                        pair_type_ids,
148                        pair_tokens,
149                        pair_words,
150                        pair_offsets,
151                        pair_special_tokens,
152                        pair_attention_mask,
153                        encoding
154                            .take_overflowing()
155                            .into_iter()
156                            .map(|encoding| {
157                                let pair_ids = [encoding.get_ids(), &[self.sep.1]].concat();
158                                let pair_type_ids = [encoding.get_type_ids(), &[1]].concat();
159                                let pair_tokens =
160                                    [encoding.get_tokens(), std::slice::from_ref(&self.sep.0)]
161                                        .concat();
162                                let pair_words = [encoding.get_word_ids(), &[None]].concat();
163                                let pair_offsets = [encoding.get_offsets(), &[(0, 0)]].concat();
164                                let pair_special_tokens =
165                                    [&vec![0u32; encoding.get_type_ids().len()][..], &[1]].concat();
166                                let pair_attention_mask = vec![1; pair_ids.len()];
167
168                                // For compatibility with `TemplateProcessing`, the sequence_ranges
169                                // shouldn't contain the special tokens.
170                                let pair_sequence_ranges =
171                                    AHashMap::from_iter(vec![(1, 0..pair_ids.len() - 1)]);
172                                Encoding::new(
173                                    pair_ids,
174                                    pair_type_ids,
175                                    pair_tokens,
176                                    pair_words,
177                                    pair_offsets,
178                                    pair_special_tokens,
179                                    pair_attention_mask,
180                                    vec![],
181                                    pair_sequence_ranges,
182                                )
183                            })
184                            .collect(),
185                        pair_sequence_ranges,
186                    )
187                }
188            })
189            .collect();
190
191        Ok(encodings)
192    }
193}
194
195#[cfg(test)]
196mod tests {
197    use super::*;
198
199    #[test]
200    fn serde() {
201        let bert = BertProcessing::default();
202        let bert_r = r#"{"type":"BertProcessing","sep":["[SEP]",102],"cls":["[CLS]",101]}"#;
203        assert_eq!(serde_json::to_string(&bert).unwrap(), bert_r);
204        assert_eq!(
205            serde_json::from_str::<BertProcessing>(bert_r).unwrap(),
206            bert
207        );
208    }
209
210    #[test]
211    fn bert_processing() {
212        let processor = BertProcessing::default();
213        assert_eq!(processor.added_tokens(false), 2);
214        assert_eq!(processor.added_tokens(true), 3);
215
216        use crate::Token;
217        let encoding = Encoding::from_tokens(
218            vec![
219                Token::new(12, "Hello".into(), (0, 5)),
220                Token::new(14, "there".into(), (6, 11)),
221            ],
222            0,
223        );
224        let pair = Encoding::from_tokens(vec![Token::new(15, "pair".into(), (0, 4))], 0);
225        let single_encoding = processor.process(encoding.clone(), None, true).unwrap();
226        assert_eq!(
227            single_encoding,
228            Encoding::new(
229                vec![101, 12, 14, 102],
230                vec![0, 0, 0, 0],
231                vec![
232                    "[CLS]".into(),
233                    "Hello".into(),
234                    "there".into(),
235                    "[SEP]".into()
236                ],
237                vec![None, None, None, None],
238                vec![(0, 0), (0, 5), (6, 11), (0, 0)],
239                vec![1, 0, 0, 1],
240                vec![1, 1, 1, 1],
241                vec![],
242                AHashMap::from_iter(vec![(0, 1..3)]),
243            )
244        );
245        assert_eq!(single_encoding.token_to_sequence(2), Some(0));
246        assert_eq!(single_encoding.token_to_sequence(3), None);
247        let pair_encoding = processor
248            .process(encoding.clone(), Some(pair.clone()), true)
249            .unwrap();
250        assert_eq!(
251            pair_encoding,
252            Encoding::new(
253                vec![101, 12, 14, 102, 15, 102],
254                vec![0, 0, 0, 0, 1, 1],
255                vec![
256                    "[CLS]".into(),
257                    "Hello".into(),
258                    "there".into(),
259                    "[SEP]".into(),
260                    "pair".into(),
261                    "[SEP]".into()
262                ],
263                vec![None, None, None, None, None, None],
264                vec![(0, 0), (0, 5), (6, 11), (0, 0), (0, 4), (0, 0)],
265                vec![1, 0, 0, 1, 0, 1],
266                vec![1, 1, 1, 1, 1, 1],
267                vec![],
268                AHashMap::from_iter(vec![(0, 1..3), (1, 4..5)]),
269            )
270        );
271        assert_eq!(pair_encoding.token_to_sequence(2), Some(0));
272        assert_eq!(pair_encoding.token_to_sequence(3), None);
273        assert_eq!(pair_encoding.token_to_sequence(4), Some(1));
274        assert_eq!(pair_encoding.token_to_sequence(5), None);
275
276        // No special tokens
277        let pair_encoding = processor.process(encoding, Some(pair), false).unwrap();
278        assert_eq!(
279            pair_encoding,
280            Encoding::new(
281                vec![12, 14, 15],
282                vec![0, 0, 1],
283                vec!["Hello".into(), "there".into(), "pair".into(),],
284                vec![None, None, None],
285                vec![(0, 5), (6, 11), (0, 4)],
286                vec![0, 0, 0],
287                vec![1, 1, 1],
288                vec![],
289                AHashMap::from_iter(vec![(0, 0..2), (1, 2..3)]),
290            )
291        );
292        assert_eq!(pair_encoding.token_to_sequence(0), Some(0));
293        assert_eq!(pair_encoding.token_to_sequence(1), Some(0));
294        assert_eq!(pair_encoding.token_to_sequence(2), Some(1));
295    }
296}