File size: 6,427 Bytes
f2878d0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
//! Lowercase WordPiece, the same normalisation as `WordPiece.encode` in tinydecide.js.
//!
//! Per code point: drop NUL, U+FFFD and control/format characters (tab, newline and CR stay as
//! spaces); space separators become spaces; CJK ideographs become one-character words; everything
//! else is NFD-decomposed, nonspacing marks are dropped, and each remaining character is lowercased.
//! Words split on whitespace (as JavaScript's `\s` sees it) and on punctuation (the ASCII symbol
//! ranges plus general category P*), then each word is split greedily into vocabulary pieces.

use std::collections::HashMap;

use unicode_normalization::UnicodeNormalization;
use unicode_properties::general_category::{GeneralCategory, GeneralCategoryGroup, UnicodeGeneralCategory};

use crate::error::{Error, Result};

/// Token ids plus the byte range of `text` each token came from.
#[derive(Debug, Clone, Default)]
pub struct Encoded {
    pub ids: Vec<u32>,
    /// `(start, end)` UTF-8 byte offsets into the encoded text.
    pub offsets: Vec<(usize, usize)>,
}

#[derive(Debug, Clone)]
pub struct WordPiece {
    vocab: HashMap<String, u32>,
    unk: u32,
    prefix: String,
    max_chars: usize,
}

fn is_cjk(cp: u32) -> bool {
    (0x4E00..=0x9FFF).contains(&cp)
        || (0x3400..=0x4DBF).contains(&cp)
        || (0x20000..=0x2A6DF).contains(&cp)
        || (0x2A700..=0x2B73F).contains(&cp)
        || (0x2B740..=0x2B81F).contains(&cp)
        || (0x2B820..=0x2CEAF).contains(&cp)
        || (0xF900..=0xFAFF).contains(&cp)
        || (0x2F800..=0x2FA1F).contains(&cp)
}

fn is_punct(ch: char) -> bool {
    let c = ch as u32;
    if (33..=47).contains(&c) || (58..=64).contains(&c) || (91..=96).contains(&c) || (123..=126).contains(&c) {
        return true;
    }
    ch.general_category_group() == GeneralCategoryGroup::Punctuation
}

/// JavaScript's `\s` (WhiteSpace + LineTerminator).
pub(crate) fn is_js_space(ch: char) -> bool {
    matches!(
        ch,
        '\t' | '\n' | '\u{0B}' | '\u{0C}' | '\r' | ' ' | '\u{A0}' | '\u{1680}' | '\u{2000}'..='\u{200A}'
            | '\u{2028}' | '\u{2029}' | '\u{202F}' | '\u{205F}' | '\u{3000}' | '\u{FEFF}'
    )
}

type Ch = (char, usize, usize);

impl WordPiece {
    /// `vocab[i]` is the piece for id `i` (`None` for unused ids). Later duplicates win, as in the JS.
    pub fn new(vocab: &[Option<String>], unk: &str, prefix: &str, max_chars: usize) -> Result<Self> {
        let mut map = HashMap::with_capacity(vocab.len());
        for (i, p) in vocab.iter().enumerate() {
            if let Some(p) = p {
                map.insert(p.clone(), i as u32);
            }
        }
        let unk = *map.get(unk).ok_or_else(|| Error::Model(format!("unknown-token piece {unk:?} is not in the vocabulary")))?;
        Ok(WordPiece { vocab: map, unk, prefix: prefix.to_string(), max_chars })
    }

    pub fn unk(&self) -> u32 {
        self.unk
    }

    pub fn encode(&self, text: &str) -> Encoded {
        let mut chars: Vec<Ch> = Vec::with_capacity(text.len() + 8);
        for (a, ch) in text.char_indices() {
            let b = a + ch.len_utf8();
            let cp = ch as u32;
            let cat = ch.general_category();
            if cp == 0
                || cp == 0xFFFD
                || (matches!(cat, GeneralCategory::Control | GeneralCategory::Format) && !matches!(ch, '\t' | '\n' | '\r'))
            {
                continue;
            }
            if matches!(ch, ' ' | '\t' | '\n' | '\r') || cat == GeneralCategory::SpaceSeparator {
                chars.push((' ', a, b));
                continue;
            }
            if is_cjk(cp) {
                chars.push((' ', a, a));
                chars.push((ch, a, b));
                chars.push((' ', b, b));
                continue;
            }
            for d in std::iter::once(ch).nfd() {
                if d.general_category() == GeneralCategory::NonspacingMark {
                    continue;
                }
                for l in d.to_lowercase() {
                    chars.push((l, a, b));
                }
            }
        }

        let mut words: Vec<Vec<Ch>> = Vec::new();
        let mut cur: Vec<Ch> = Vec::new();
        for c in chars {
            if c.0 == ' ' || is_js_space(c.0) {
                if !cur.is_empty() {
                    words.push(std::mem::take(&mut cur));
                }
                continue;
            }
            if is_punct(c.0) {
                if !cur.is_empty() {
                    words.push(std::mem::take(&mut cur));
                }
                words.push(vec![c]);
                continue;
            }
            cur.push(c);
        }
        if !cur.is_empty() {
            words.push(cur);
        }

        let mut out = Encoded::default();
        let mut s = String::new();
        let mut pieces: Vec<(u32, usize, usize)> = Vec::new();
        for w in &words {
            let whole = (w[0].1, w[w.len() - 1].2);
            if w.len() > self.max_chars {
                out.ids.push(self.unk);
                out.offsets.push(whole);
                continue;
            }
            pieces.clear();
            let mut start = 0;
            let mut bad = false;
            while start < w.len() {
                let mut end = w.len();
                let mut got = None;
                while start < end {
                    s.clear();
                    if start > 0 {
                        s.push_str(&self.prefix);
                    }
                    s.extend(w[start..end].iter().map(|c| c.0));
                    if let Some(&id) = self.vocab.get(s.as_str()) {
                        got = Some(id);
                        break;
                    }
                    end -= 1;
                }
                match got {
                    Some(id) => {
                        pieces.push((id, w[start].1, w[end - 1].2));
                        start = end;
                    }
                    None => {
                        bad = true;
                        break;
                    }
                }
            }
            if bad {
                out.ids.push(self.unk);
                out.offsets.push(whole);
            } else {
                for &(id, a, b) in &pieces {
                    out.ids.push(id);
                    out.offsets.push((a, b));
                }
            }
        }
        out
    }
}