trustformers_tokenizers/
streaming.rs1use anyhow::Result as AnyhowResult;
2use std::io::{BufRead, BufReader, Read};
3use trustformers_core::errors::Result;
4use trustformers_core::traits::{TokenizedInput, Tokenizer};
5
6pub struct StreamingTokenizer<T: Tokenizer> {
8 tokenizer: T,
9 buffer_size: usize,
10 overlap_size: usize,
11 max_chunk_length: Option<usize>,
12}
13
14impl<T: Tokenizer> StreamingTokenizer<T> {
15 pub fn new(tokenizer: T) -> Self {
17 Self {
18 tokenizer,
19 buffer_size: 8192, overlap_size: 256, max_chunk_length: None,
22 }
23 }
24
25 pub fn with_buffer_size(mut self, buffer_size: usize) -> Self {
27 self.buffer_size = buffer_size;
28 self
29 }
30
31 pub fn with_overlap_size(mut self, overlap_size: usize) -> Self {
33 self.overlap_size = overlap_size;
34 self
35 }
36
37 pub fn with_max_chunk_length(mut self, max_length: usize) -> Self {
39 self.max_chunk_length = Some(max_length);
40 self
41 }
42
43 pub fn process_stream<R: Read>(&self, reader: R) -> Result<Vec<TokenizedInput>> {
45 let mut buf_reader = BufReader::with_capacity(self.buffer_size, reader);
46 let mut chunks = Vec::new();
47 let mut buffer = String::new();
48 let mut previous_overlap = String::new();
49
50 loop {
51 buffer.clear();
52 let bytes_read = buf_reader.read_line(&mut buffer).map_err(|e| {
53 trustformers_core::errors::TrustformersError::other(format!("I/O error: {}", e))
54 })?;
55
56 if bytes_read == 0 {
57 break; }
59
60 let full_text = if previous_overlap.is_empty() {
62 buffer.clone()
63 } else {
64 format!("{}{}", previous_overlap, buffer)
65 };
66
67 let tokenized = self.tokenize_chunk(&full_text)?;
69 chunks.push(tokenized);
70
71 if full_text.len() > self.overlap_size {
73 previous_overlap = full_text[full_text.len() - self.overlap_size..].to_string();
74 } else {
75 previous_overlap.clear();
76 }
77 }
78
79 Ok(chunks)
80 }
81
82 pub fn process_text(&self, text: &str) -> Result<Vec<TokenizedInput>> {
84 let mut chunks = Vec::new();
85 let mut start = 0;
86 let chunk_size = self.buffer_size;
87
88 if text.is_empty() {
90 let empty_chunk = self.tokenize_chunk("")?;
91 chunks.push(empty_chunk);
92 return Ok(chunks);
93 }
94
95 while start < text.len() {
96 let end = std::cmp::min(start + chunk_size, text.len());
97 let mut chunk_end = end;
98
99 if end < text.len() {
101 if let Some(last_space) = text[start..end].rfind(' ') {
102 chunk_end = start + last_space;
103 }
104 }
105
106 if chunk_end <= start {
108 chunk_end = std::cmp::min(start + 1, text.len());
109 }
110
111 let chunk_text = &text[start..chunk_end];
112 let tokenized = self.tokenize_chunk(chunk_text)?;
113 chunks.push(tokenized);
114
115 let next_start = if chunk_end > self.overlap_size {
117 chunk_end - self.overlap_size
118 } else {
119 chunk_end
120 };
121
122 start = std::cmp::max(next_start, start + 1);
124 }
125
126 Ok(chunks)
127 }
128
129 pub fn process_lines<I>(&self, lines: I) -> Result<Vec<TokenizedInput>>
131 where
132 I: Iterator<Item = String>,
133 {
134 let mut chunks = Vec::new();
135 let mut current_chunk = String::new();
136
137 for line in lines {
138 if !current_chunk.is_empty() {
140 current_chunk.push('\n');
141 }
142 current_chunk.push_str(&line);
143
144 if current_chunk.len() >= self.buffer_size {
146 let tokenized = self.tokenize_chunk(¤t_chunk)?;
147 chunks.push(tokenized);
148
149 if current_chunk.len() > self.overlap_size {
151 current_chunk =
152 current_chunk[current_chunk.len() - self.overlap_size..].to_string();
153 } else {
154 current_chunk.clear();
155 }
156 }
157 }
158
159 if !current_chunk.is_empty() {
161 let tokenized = self.tokenize_chunk(¤t_chunk)?;
162 chunks.push(tokenized);
163 }
164
165 Ok(chunks)
166 }
167
168 fn tokenize_chunk(&self, text: &str) -> Result<TokenizedInput> {
170 let mut tokenized = self.tokenizer.encode(text)?;
171
172 if let Some(max_len) = self.max_chunk_length {
174 if tokenized.input_ids.len() > max_len {
175 tokenized.input_ids.truncate(max_len);
176 tokenized.attention_mask.truncate(max_len);
177 if let Some(ref mut token_type_ids) = tokenized.token_type_ids {
178 token_type_ids.truncate(max_len);
179 }
180 }
181 }
182
183 Ok(tokenized)
184 }
185
186 pub fn tokenizer(&self) -> &T {
188 &self.tokenizer
189 }
190
191 pub fn buffer_size(&self) -> usize {
193 self.buffer_size
194 }
195
196 pub fn overlap_size(&self) -> usize {
198 self.overlap_size
199 }
200
201 pub fn max_chunk_length(&self) -> Option<usize> {
203 self.max_chunk_length
204 }
205}
206
207pub struct BatchedStreamingTokenizer<T: Tokenizer> {
209 streaming_tokenizer: StreamingTokenizer<T>,
210 batch_size: usize,
211}
212
213impl<T: Tokenizer> BatchedStreamingTokenizer<T> {
214 pub fn new(tokenizer: T, batch_size: usize) -> Self {
216 Self {
217 streaming_tokenizer: StreamingTokenizer::new(tokenizer),
218 batch_size,
219 }
220 }
221
222 pub fn with_streaming_params(mut self, buffer_size: usize, overlap_size: usize) -> Self {
224 self.streaming_tokenizer = self
225 .streaming_tokenizer
226 .with_buffer_size(buffer_size)
227 .with_overlap_size(overlap_size);
228 self
229 }
230
231 pub fn with_max_chunk_length(mut self, max_length: usize) -> Self {
233 self.streaming_tokenizer = self.streaming_tokenizer.with_max_chunk_length(max_length);
234 self
235 }
236
237 pub fn process_text_batch(&self, texts: &[String]) -> Result<Vec<Vec<TokenizedInput>>> {
239 let mut results = Vec::new();
240
241 for batch in texts.chunks(self.batch_size) {
242 let mut batch_results = Vec::new();
243 for text in batch {
244 let tokenized_chunks = self.streaming_tokenizer.process_text(text)?;
245 batch_results.push(tokenized_chunks);
246 }
247 results.extend(batch_results);
248 }
249
250 Ok(results)
251 }
252
253 pub fn batch_size(&self) -> usize {
255 self.batch_size
256 }
257
258 pub fn streaming_tokenizer(&self) -> &StreamingTokenizer<T> {
260 &self.streaming_tokenizer
261 }
262}
263
264pub struct TextFileIterator<R: BufRead> {
266 reader: R,
267 buffer: String,
268 chunk_size: usize,
269 #[allow(dead_code)]
272 overlap_size: usize,
273 eof: bool,
274}
275
276impl<R: BufRead> TextFileIterator<R> {
277 pub fn new(reader: R, chunk_size: usize, overlap_size: usize) -> Self {
279 Self {
280 reader,
281 buffer: String::new(),
282 chunk_size,
283 overlap_size,
284 eof: false,
285 }
286 }
287
288 pub fn next_chunk(&mut self) -> AnyhowResult<Option<String>> {
290 if self.eof {
291 return Ok(None);
292 }
293
294 self.buffer.clear();
295
296 let mut bytes_read = 0;
298 let mut temp_buf = String::new();
299
300 while bytes_read < self.chunk_size {
301 temp_buf.clear();
302 let n = self.reader.read_line(&mut temp_buf)?;
303 if n == 0 {
304 self.eof = true;
305 break;
306 }
307 self.buffer.push_str(&temp_buf);
308 bytes_read += n;
309 }
310
311 if self.buffer.is_empty() {
312 Ok(None)
313 } else {
314 Ok(Some(self.buffer.clone()))
315 }
316 }
317}
318
319impl<R: BufRead> Iterator for TextFileIterator<R> {
320 type Item = AnyhowResult<String>;
321
322 fn next(&mut self) -> Option<Self::Item> {
323 match self.next_chunk() {
324 Ok(Some(chunk)) => Some(Ok(chunk)),
325 Ok(None) => None,
326 Err(e) => Some(Err(e)),
327 }
328 }
329}
330
331#[cfg(test)]
332mod tests {
333 use super::*;
334 use crate::char::CharTokenizer;
335 use std::io::Cursor;
336
337 fn create_test_tokenizer() -> CharTokenizer {
338 let mut vocab = std::collections::HashMap::new();
339 vocab.insert("a".to_string(), 0);
340 vocab.insert("b".to_string(), 1);
341 vocab.insert("c".to_string(), 2);
342 vocab.insert(" ".to_string(), 3);
343 CharTokenizer::new(vocab)
344 }
345
346 #[test]
347 fn test_streaming_tokenizer_basic() {
348 let tokenizer = create_test_tokenizer();
349 let streaming = StreamingTokenizer::new(tokenizer);
350
351 let text = "Hello world! This is a test of streaming tokenization.";
352 let chunks = streaming.process_text(text).expect("Operation failed in test");
353
354 assert!(!chunks.is_empty());
355 for chunk in chunks {
357 assert!(!chunk.input_ids.is_empty());
358 assert!(!chunk.attention_mask.is_empty());
359 }
360 }
361
362 #[test]
363 fn test_streaming_tokenizer_with_params() {
364 let tokenizer = create_test_tokenizer();
365 let streaming = StreamingTokenizer::new(tokenizer)
366 .with_buffer_size(50)
367 .with_overlap_size(10)
368 .with_max_chunk_length(20);
369
370 let text = "This is a longer text that should be split into multiple chunks based on the buffer size.";
371 let chunks = streaming.process_text(text).expect("Operation failed in test");
372
373 assert!(chunks.len() > 1);
374
375 for chunk in chunks {
377 assert!(chunk.input_ids.len() <= 20);
378 }
379 }
380
381 #[test]
382 fn test_streaming_tokenizer_from_reader() {
383 let tokenizer = create_test_tokenizer();
384 let streaming = StreamingTokenizer::new(tokenizer);
385
386 let text = "Line 1\nLine 2\nLine 3\n";
387 let cursor = Cursor::new(text.as_bytes());
388 let chunks = streaming.process_stream(cursor).expect("Operation failed in test");
389
390 assert!(!chunks.is_empty());
391 for chunk in chunks {
392 assert!(!chunk.input_ids.is_empty());
393 }
394 }
395
396 #[test]
397 fn test_streaming_tokenizer_lines() {
398 let tokenizer = create_test_tokenizer();
399 let streaming = StreamingTokenizer::new(tokenizer).with_buffer_size(20);
400
401 let lines = vec![
402 "First line".to_string(),
403 "Second line".to_string(),
404 "Third line".to_string(),
405 ];
406
407 let chunks = streaming.process_lines(lines.into_iter()).expect("Operation failed in test");
408 assert!(!chunks.is_empty());
409 }
410
411 #[test]
412 fn test_batched_streaming_tokenizer() {
413 let tokenizer = create_test_tokenizer();
414 let batched = BatchedStreamingTokenizer::new(tokenizer, 2).with_streaming_params(50, 10);
415
416 let texts = vec![
417 "First text to tokenize".to_string(),
418 "Second text to tokenize".to_string(),
419 "Third text to tokenize".to_string(),
420 ];
421
422 let results = batched.process_text_batch(&texts).expect("Operation failed in test");
423 assert_eq!(results.len(), 3);
424
425 for result in results {
426 assert!(!result.is_empty());
427 for chunk in result {
428 assert!(!chunk.input_ids.is_empty());
429 }
430 }
431 }
432
433 #[test]
434 fn test_text_file_iterator() {
435 let text = "Line 1\nLine 2\nLine 3\nLine 4\n";
436 let cursor = Cursor::new(text.as_bytes());
437 let buf_reader = BufReader::new(cursor);
438
439 let iterator = TextFileIterator::new(buf_reader, 10, 2);
440
441 let chunks: std::result::Result<Vec<_>, _> = iterator.collect();
442 let chunks = chunks.expect("Operation failed in test");
443
444 assert!(!chunks.is_empty());
445 for chunk in chunks {
446 assert!(!chunk.is_empty());
447 }
448 }
449
450 #[test]
451 fn test_streaming_empty_text() {
452 let tokenizer = create_test_tokenizer();
453 let streaming = StreamingTokenizer::new(tokenizer);
454
455 let chunks = streaming.process_text("").expect("Operation failed in test");
456 assert_eq!(chunks.len(), 1); assert!(chunks[0].input_ids.is_empty() || chunks[0].input_ids.len() == 1);
458 }
460
461 #[test]
462 fn test_streaming_configuration() {
463 let tokenizer = create_test_tokenizer();
464 let streaming = StreamingTokenizer::new(tokenizer)
465 .with_buffer_size(1024)
466 .with_overlap_size(128)
467 .with_max_chunk_length(512);
468
469 assert_eq!(streaming.buffer_size(), 1024);
470 assert_eq!(streaming.overlap_size(), 128);
471 assert_eq!(streaming.max_chunk_length(), Some(512));
472 }
473}