tokenizers/models/wordpiece/
serialization.rs

1use super::{super::OrderedVocabIter, WordPiece, WordPieceBuilder};
2use ahash::{AHashMap, AHashSet};
3use serde::{
4    de::{MapAccess, Visitor},
5    ser::SerializeStruct,
6    Deserialize, Deserializer, Serialize, Serializer,
7};
8
9impl Serialize for WordPiece {
10    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
11    where
12        S: Serializer,
13    {
14        let mut model = serializer.serialize_struct("WordPiece", 5)?;
15
16        // Small fields first
17        model.serialize_field("type", "WordPiece")?;
18        model.serialize_field("unk_token", &self.unk_token)?;
19        model.serialize_field("continuing_subword_prefix", &self.continuing_subword_prefix)?;
20        model.serialize_field("max_input_chars_per_word", &self.max_input_chars_per_word)?;
21
22        // Then large ones
23        let ordered_vocab = OrderedVocabIter::new(&self.vocab_r);
24        model.serialize_field("vocab", &ordered_vocab)?;
25
26        model.end()
27    }
28}
29
30impl<'de> Deserialize<'de> for WordPiece {
31    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
32    where
33        D: Deserializer<'de>,
34    {
35        deserializer.deserialize_struct(
36            "WordPiece",
37            &[
38                "type",
39                "unk_token",
40                "continuing_subword_prefix",
41                "max_input_chars_per_word",
42                "vocab",
43            ],
44            WordPieceVisitor,
45        )
46    }
47}
48
49struct WordPieceVisitor;
50impl<'de> Visitor<'de> for WordPieceVisitor {
51    type Value = WordPiece;
52
53    fn expecting(&self, fmt: &mut std::fmt::Formatter) -> std::fmt::Result {
54        write!(fmt, "struct WordPiece")
55    }
56
57    fn visit_map<V>(self, mut map: V) -> std::result::Result<Self::Value, V::Error>
58    where
59        V: MapAccess<'de>,
60    {
61        let mut builder = WordPieceBuilder::new();
62        let mut missing_fields = vec![
63            // for retrocompatibility the "type" field is not mandatory
64            "unk_token",
65            "continuing_subword_prefix",
66            "max_input_chars_per_word",
67            "vocab",
68        ]
69        .into_iter()
70        .collect::<AHashSet<_>>();
71
72        while let Some(key) = map.next_key::<String>()? {
73            match key.as_ref() {
74                "unk_token" => builder = builder.unk_token(map.next_value()?),
75                "continuing_subword_prefix" => {
76                    builder = builder.continuing_subword_prefix(map.next_value()?)
77                }
78                "max_input_chars_per_word" => {
79                    builder = builder.max_input_chars_per_word(map.next_value()?)
80                }
81                "vocab" => {
82                    let vocab: AHashMap<String, u32> = map.next_value()?;
83                    builder = builder.vocab(vocab)
84                }
85                "type" => match map.next_value()? {
86                    "WordPiece" => {}
87                    u => {
88                        return Err(serde::de::Error::invalid_value(
89                            serde::de::Unexpected::Str(u),
90                            &"WordPiece",
91                        ))
92                    }
93                },
94                _ => {}
95            }
96            missing_fields.remove::<str>(&key);
97        }
98
99        if !missing_fields.is_empty() {
100            Err(serde::de::Error::missing_field(
101                missing_fields.iter().next().unwrap(),
102            ))
103        } else {
104            Ok(builder.build().map_err(serde::de::Error::custom)?)
105        }
106    }
107}
108
109#[cfg(test)]
110mod tests {
111    use super::*;
112
113    #[test]
114    fn serde() {
115        let wp = WordPiece::default();
116        let wp_s = "{\
117            \"type\":\"WordPiece\",\
118            \"unk_token\":\"[UNK]\",\
119            \"continuing_subword_prefix\":\"##\",\
120            \"max_input_chars_per_word\":100,\
121            \"vocab\":{}\
122        }";
123
124        assert_eq!(serde_json::to_string(&wp).unwrap(), wp_s);
125        assert_eq!(serde_json::from_str::<WordPiece>(wp_s).unwrap(), wp);
126    }
127
128    #[test]
129    fn deserialization_should_fail() {
130        let missing_unk = "{\
131            \"type\":\"WordPiece\",\
132            \"continuing_subword_prefix\":\"##\",\
133            \"max_input_chars_per_word\":100,\
134            \"vocab\":{}\
135        }";
136        assert!(serde_json::from_str::<WordPiece>(missing_unk)
137            .unwrap_err()
138            .to_string()
139            .starts_with("missing field `unk_token`"));
140
141        let wrong_type = "{\
142            \"type\":\"WordLevel\",\
143            \"unk_token\":\"[UNK]\",\
144            \"vocab\":{}\
145        }";
146        assert!(serde_json::from_str::<WordPiece>(wrong_type)
147            .unwrap_err()
148            .to_string()
149            .starts_with("invalid value: string \"WordLevel\", expected WordPiece"));
150    }
151}