use rten_text::models::DecodeError;
use rten_text::{Tokenizer, TokenizerError};
use crate::generator::{GeneratorError, GeneratorItem};
pub struct TextDecoder<'a, G: Iterator<Item = GeneratorItem>> {
generator: G,
tokenizer: &'a Tokenizer,
}
impl<'a, G> TextDecoder<'a, G>
where
G: Iterator<Item = GeneratorItem>,
{
pub fn wrap(generator: G, tokenizer: &'a Tokenizer) -> TextDecoder<'a, G> {
TextDecoder {
generator,
tokenizer,
}
}
}
impl<G: Iterator<Item = GeneratorItem>> Iterator for TextDecoder<'_, G> {
type Item = Result<String, GeneratorError>;
fn next(&mut self) -> Option<Self::Item> {
let mut token_buf = Vec::new();
for token in self.generator.by_ref() {
let token = match token {
Ok(tok) => tok,
Err(err) => return Some(Err(err)),
};
token_buf.push(token);
let text = self.tokenizer.decode(&token_buf);
match text {
Ok(text) => return Some(Ok(text)),
Err(TokenizerError::DecodeError(DecodeError::InvalidUtf8)) => {
continue;
}
Err(err) => {
return Some(Err(GeneratorError::DecodeError(err)));
}
}
}
None
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use rten_text::models::{Bpe, BpeOptions, WordPiece};
use rten_text::pre_tokenizers::Split;
use rten_text::{TokenId, Tokenizer};
use crate::{GeneratorError, GeneratorUtils};
fn create_tokenizer() -> Tokenizer {
let vocab: HashMap<String, TokenId> = [("one", 1), ("two", 2), ("three", 3)]
.into_iter()
.map(|(s, id)| (s.to_string(), id))
.collect();
let model = WordPiece::from_vocab(vocab, Default::default());
Tokenizer::new(model, Default::default())
}
fn create_bpe_tokenizer() -> Tokenizer {
let model = Bpe::new(BpeOptions::default()).unwrap();
Tokenizer::new(model, Default::default()).with_pre_tokenizer(Box::new(Split::gpt2()))
}
#[test]
fn test_decode() {
let tokenizer = create_tokenizer();
let generator = [1, 2, 3].into_iter().map(Ok);
let tokens: Vec<_> = generator
.decode(&tokenizer)
.map(|tok| tok.map_err(|e| e.to_string()))
.collect();
assert_eq!(tokens, ["one", "two", "three"].map(|s| Ok(s.to_string())));
}
#[test]
fn test_decode_partial_utf8() {
let tokenizer = create_bpe_tokenizer();
let token_ids = tokenizer.encode("😊", None).unwrap().into_token_ids();
assert!(token_ids.len() > 1);
let generator = token_ids.into_iter().map(|tok_id| Ok(tok_id as u32));
let tokens: Vec<_> = generator
.decode(&tokenizer)
.map(|tok| tok.map_err(|e| e.to_string()))
.collect();
assert_eq!(tokens, ["😊"].map(|s| Ok(s.to_string())));
}
#[test]
fn test_generate_error() {
let tokenizer = create_tokenizer();
let generator = [
Ok(1),
Err(GeneratorError::GenerateError("oh no".to_string().into())),
Ok(3),
]
.into_iter();
let tokens: Vec<_> = generator
.decode(&tokenizer)
.map(|tok| tok.map_err(|e| e.to_string()))
.collect();
assert_eq!(
tokens,
[
Ok("one".to_string()),
Err("generation error: oh no".to_string()),
Ok("three".to_string())
]
);
}
#[test]
fn test_decode_error() {
let tokenizer = create_tokenizer();
let generator = [1, 5, 3].into_iter().map(Ok);
let tokens: Vec<_> = generator
.decode(&tokenizer)
.map(|tok| tok.map_err(|e| e.to_string()))
.collect();
assert_eq!(
tokens,
[
Ok("one".to_string()),
Err("decode error: decoding failed: cannot decode unknown token ID 5".to_string()),
Ok("three".to_string())
]
);
}
}