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