ferrum_testkit/
tokenizer.rs1use ferrum_interfaces::{
4 tokenizer::{TokenizerInfo, TokenizerType},
5 Tokenizer,
6};
7use ferrum_types::{Result, SpecialTokens, TokenId};
8
9pub struct MockTokenizer {
12 vocab_size: usize,
13 special_tokens: SpecialTokens,
14}
15
16impl MockTokenizer {
17 pub fn new(vocab_size: usize) -> Self {
18 let eos = TokenId::new((vocab_size - 1) as u32);
19 let bos = TokenId::new((vocab_size - 2) as u32);
20 Self {
21 vocab_size,
22 special_tokens: SpecialTokens {
23 bos_token: Some(bos),
24 eos_token: Some(eos),
25 unk_token: Some(TokenId::new(0)),
26 pad_token: None,
27 sep_token: None,
28 cls_token: None,
29 mask_token: None,
30 extra_eos_tokens: Vec::new(),
31 },
32 }
33 }
34}
35
36impl Tokenizer for MockTokenizer {
37 fn encode(&self, text: &str, add_special: bool) -> Result<Vec<TokenId>> {
38 let mut tokens = Vec::new();
39 if add_special {
40 if let Some(bos) = self.special_tokens.bos_token {
41 tokens.push(bos);
42 }
43 }
44 for word in text.split_whitespace() {
46 let hash = word
47 .bytes()
48 .fold(0u32, |acc, b| acc.wrapping_mul(31).wrapping_add(b as u32));
49 let id = 1 + (hash % (self.vocab_size as u32 - 3));
50 tokens.push(TokenId::new(id));
51 }
52 if tokens.is_empty() {
53 tokens.push(TokenId::new(1)); }
55 Ok(tokens)
56 }
57
58 fn decode(&self, tokens: &[TokenId], _skip_special: bool) -> Result<String> {
59 Ok(tokens
60 .iter()
61 .map(|t| format!("w{}", t.get()))
62 .collect::<Vec<_>>()
63 .join(" "))
64 }
65
66 fn decode_incremental(&self, _prev: &[TokenId], next: TokenId) -> Result<String> {
67 Ok(format!("w{}", next.get()))
68 }
69
70 fn vocab_size(&self) -> usize {
71 self.vocab_size
72 }
73
74 fn special_tokens(&self) -> &SpecialTokens {
75 &self.special_tokens
76 }
77
78 fn token_id(&self, _text: &str) -> Option<TokenId> {
79 None
80 }
81
82 fn token_text(&self, _token_id: TokenId) -> Option<&str> {
83 None
84 }
85
86 fn info(&self) -> TokenizerInfo {
87 TokenizerInfo {
88 tokenizer_type: TokenizerType::Custom,
89 vocab_size: self.vocab_size,
90 special_tokens: self.special_tokens.clone(),
91 supports_incremental: true,
92 supports_chat_template: false,
93 max_token_length: Some(128),
94 model_name: Some("mock".into()),
95 }
96 }
97}