tokenizers/models/wordpiece/
serialization.rs1use 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 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 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 "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}