icu_segmenter/complex/lstm/
mod.rs1use crate::grapheme::GraphemeClusterSegmenterBorrowed;
6use crate::provider::*;
7use alloc::vec::Vec;
8use core::char::{decode_utf16, REPLACEMENT_CHARACTER};
9use potential_utf::PotentialUtf8;
10use zerovec::maps::ZeroMapBorrowed;
11
12mod matrix;
13use matrix::*;
14
15pub(super) struct LstmSegmenterIterator<'s, 'data> {
18 input: &'s str,
19 pos_utf8: usize,
20 bies: BiesIterator<'s, 'data>,
21}
22
23impl Iterator for LstmSegmenterIterator<'_, '_> {
24 type Item = usize;
25
26 fn next(&mut self) -> Option<Self::Item> {
27 loop {
28 let is_e = self.bies.next()?;
29 self.pos_utf8 += self.input[self.pos_utf8..].chars().next()?.len_utf8();
30 if is_e || self.bies.len() == 0 {
31 return Some(self.pos_utf8);
32 }
33 }
34 }
35}
36
37pub(super) struct LstmSegmenterIteratorUtf16<'s, 'data> {
38 bies: BiesIterator<'s, 'data>,
39 pos: usize,
40}
41
42impl Iterator for LstmSegmenterIteratorUtf16<'_, '_> {
43 type Item = usize;
44
45 fn next(&mut self) -> Option<Self::Item> {
46 loop {
47 self.pos += 1;
48 if self.bies.next()? || self.bies.len() == 0 {
49 return Some(self.pos);
50 }
51 }
52 }
53}
54
55pub(super) struct LstmSegmenter<'data> {
56 dic: ZeroMapBorrowed<'data, PotentialUtf8, u16>,
57 embedding: MatrixZero<'data, 2>,
58 fw_w: MatrixZero<'data, 3>,
59 fw_u: MatrixZero<'data, 3>,
60 fw_b: MatrixZero<'data, 2>,
61 bw_w: MatrixZero<'data, 3>,
62 bw_u: MatrixZero<'data, 3>,
63 bw_b: MatrixZero<'data, 2>,
64 timew_fw: MatrixZero<'data, 2>,
65 timew_bw: MatrixZero<'data, 2>,
66 time_b: MatrixZero<'data, 1>,
67 grapheme: Option<GraphemeClusterSegmenterBorrowed<'data>>,
68}
69
70impl<'data> LstmSegmenter<'data> {
71 pub(super) fn new(
73 lstm: &'data LstmData<'data>,
74 grapheme: GraphemeClusterSegmenterBorrowed<'data>,
75 ) -> Self {
76 let LstmData::Float32(lstm) = lstm;
77 let time_w = MatrixZero::from(&lstm.time_w);
78 #[expect(clippy::unwrap_used)] let timew_fw = time_w.submatrix(0).unwrap();
80 #[expect(clippy::unwrap_used)] let timew_bw = time_w.submatrix(1).unwrap();
82 Self {
83 dic: lstm.dic.as_borrowed(),
84 embedding: MatrixZero::from(&lstm.embedding),
85 fw_w: MatrixZero::from(&lstm.fw_w),
86 fw_u: MatrixZero::from(&lstm.fw_u),
87 fw_b: MatrixZero::from(&lstm.fw_b),
88 bw_w: MatrixZero::from(&lstm.bw_w),
89 bw_u: MatrixZero::from(&lstm.bw_u),
90 bw_b: MatrixZero::from(&lstm.bw_b),
91 timew_fw,
92 timew_bw,
93 time_b: MatrixZero::from(&lstm.time_b),
94 grapheme: (lstm.model == ModelType::GraphemeClusters).then_some(grapheme),
95 }
96 }
97
98 pub(super) fn segment_str<'a>(&'a self, input: &'a str) -> LstmSegmenterIterator<'a, 'data> {
100 let input_seq = if let Some(grapheme) = self.grapheme {
101 grapheme
102 .segment_str(input)
103 .collect::<Vec<usize>>()
104 .windows(2)
105 .map(|chunk| {
106 let range = if let [first, second, ..] = chunk {
107 *first..*second
108 } else {
109 unreachable!()
110 };
111 let grapheme_cluster = if let Some(grapheme_cluster) = input.get(range) {
112 grapheme_cluster
113 } else {
114 return self.dic.len() as u16;
115 };
116
117 self.dic
118 .get_copied(PotentialUtf8::from_str(grapheme_cluster))
119 .unwrap_or_else(|| self.dic.len() as u16)
120 })
121 .collect()
122 } else {
123 input
124 .chars()
125 .map(|c| {
126 self.dic
127 .get_copied(PotentialUtf8::from_str(c.encode_utf8(&mut [0; 4])))
128 .unwrap_or_else(|| self.dic.len() as u16)
129 })
130 .collect()
131 };
132 LstmSegmenterIterator {
133 input,
134 pos_utf8: 0,
135 bies: BiesIterator::new(self, input_seq),
136 }
137 }
138
139 pub(super) fn segment_utf16<'a>(
141 &'a self,
142 input: &[u16],
143 ) -> LstmSegmenterIteratorUtf16<'a, 'data> {
144 let input_seq = if let Some(grapheme) = self.grapheme {
145 grapheme
146 .segment_utf16(input)
147 .collect::<Vec<usize>>()
148 .windows(2)
149 .map(|chunk| {
150 let range = if let [first, second, ..] = chunk {
151 *first..*second
152 } else {
153 unreachable!()
154 };
155 let grapheme_cluster = if let Some(grapheme_cluster) = input.get(range) {
156 grapheme_cluster
157 } else {
158 return self.dic.len() as u16;
159 };
160
161 self.dic
162 .get_copied_by(|key| {
163 key.as_bytes().iter().copied().cmp(
164 decode_utf16(grapheme_cluster.iter().copied()).flat_map(|c| {
165 let mut buf = [0; 4];
166 let len = c
167 .unwrap_or(REPLACEMENT_CHARACTER)
168 .encode_utf8(&mut buf)
169 .len();
170 buf.into_iter().take(len)
171 }),
172 )
173 })
174 .unwrap_or_else(|| self.dic.len() as u16)
175 })
176 .collect()
177 } else {
178 decode_utf16(input.iter().copied())
179 .map(|c| c.unwrap_or(REPLACEMENT_CHARACTER))
180 .map(|c| {
181 self.dic
182 .get_copied(PotentialUtf8::from_str(c.encode_utf8(&mut [0; 4])))
183 .unwrap_or_else(|| self.dic.len() as u16)
184 })
185 .collect()
186 };
187 LstmSegmenterIteratorUtf16 {
188 bies: BiesIterator::new(self, input_seq),
189 pos: 0,
190 }
191 }
192}
193
194struct BiesIterator<'l, 'data> {
195 segmenter: &'l LstmSegmenter<'data>,
196 input_seq: core::iter::Enumerate<alloc::vec::IntoIter<u16>>,
197 h_bw: MatrixOwned<2>,
198 curr_fw: MatrixOwned<1>,
199 c_fw: MatrixOwned<1>,
200}
201
202impl<'l, 'data> BiesIterator<'l, 'data> {
203 fn new(segmenter: &'l LstmSegmenter<'data>, input_seq: Vec<u16>) -> Self {
206 let hunits = segmenter.fw_u.dim().1;
207
208 let mut c_bw = MatrixOwned::<1>::new_zero([hunits]);
210 let mut h_bw = MatrixOwned::<2>::new_zero([input_seq.len(), hunits]);
211 for (i, &g_id) in input_seq.iter().enumerate().rev() {
212 if i + 1 < input_seq.len() {
213 h_bw.as_mut().copy_submatrix::<1>(i + 1, i);
214 }
215 #[expect(clippy::unwrap_used)]
216 compute_hc(
217 segmenter.embedding.submatrix::<1>(g_id as usize).unwrap(), h_bw.submatrix_mut(i).unwrap(), c_bw.as_mut(),
220 segmenter.bw_w,
221 segmenter.bw_u,
222 segmenter.bw_b,
223 );
224 }
225
226 Self {
227 input_seq: input_seq.into_iter().enumerate(),
228 h_bw,
229 c_fw: MatrixOwned::<1>::new_zero([hunits]),
230 curr_fw: MatrixOwned::<1>::new_zero([hunits]),
231 segmenter,
232 }
233 }
234}
235
236impl ExactSizeIterator for BiesIterator<'_, '_> {
237 fn len(&self) -> usize {
238 self.input_seq.len()
239 }
240}
241
242impl Iterator for BiesIterator<'_, '_> {
243 type Item = bool;
244
245 fn next(&mut self) -> Option<Self::Item> {
246 let (i, g_id) = self.input_seq.next()?;
247
248 #[expect(clippy::unwrap_used)]
249 compute_hc(
250 self.segmenter
251 .embedding
252 .submatrix::<1>(g_id as usize)
253 .unwrap(), self.curr_fw.as_mut(),
255 self.c_fw.as_mut(),
256 self.segmenter.fw_w,
257 self.segmenter.fw_u,
258 self.segmenter.fw_b,
259 );
260
261 #[expect(clippy::unwrap_used)] let curr_bw = self.h_bw.submatrix::<1>(i).unwrap();
263 let mut weights = [0.0; 4];
264 let mut curr_est = MatrixBorrowedMut {
265 data: &mut weights,
266 dims: [4],
267 };
268 curr_est.add_dot_2d(self.curr_fw.as_borrowed(), self.segmenter.timew_fw);
269 curr_est.add_dot_2d(curr_bw, self.segmenter.timew_bw);
270 #[expect(clippy::unwrap_used)] curr_est.add(self.segmenter.time_b).unwrap();
272 Some(weights[2] > weights[0] && weights[2] > weights[1] && weights[2] > weights[3])
276 }
277}
278
279fn compute_hc<'a>(
281 x_t: MatrixZero<'a, 1>,
282 mut h_tm1: MatrixBorrowedMut<'a, 1>,
283 mut c_tm1: MatrixBorrowedMut<'a, 1>,
284 w: MatrixZero<'a, 3>,
285 u: MatrixZero<'a, 3>,
286 b: MatrixZero<'a, 2>,
287) {
288 #[cfg(debug_assertions)]
289 {
290 let hunits = h_tm1.dim();
291 let embedd_dim = x_t.dim();
292 c_tm1.as_borrowed().debug_assert_dims([hunits]);
293 w.debug_assert_dims([4, hunits, embedd_dim]);
294 u.debug_assert_dims([4, hunits, hunits]);
295 b.debug_assert_dims([4, hunits]);
296 }
297
298 let mut s_t = b.to_owned();
299
300 s_t.as_mut().add_dot_3d_2(x_t, w);
301 s_t.as_mut().add_dot_3d_1(h_tm1.as_borrowed(), u);
302
303 #[expect(clippy::unwrap_used)] s_t.submatrix_mut::<1>(0).unwrap().sigmoid_transform();
305 #[expect(clippy::unwrap_used)] s_t.submatrix_mut::<1>(1).unwrap().sigmoid_transform();
307 #[expect(clippy::unwrap_used)] s_t.submatrix_mut::<1>(2).unwrap().tanh_transform();
309 #[expect(clippy::unwrap_used)] s_t.submatrix_mut::<1>(3).unwrap().sigmoid_transform();
311
312 #[expect(clippy::unwrap_used)] c_tm1.convolve(
314 s_t.as_borrowed().submatrix(0).unwrap(),
315 s_t.as_borrowed().submatrix(2).unwrap(),
316 s_t.as_borrowed().submatrix(1).unwrap(),
317 );
318
319 #[expect(clippy::unwrap_used)] h_tm1.mul_tanh(s_t.as_borrowed().submatrix(3).unwrap(), c_tm1.as_borrowed());
321}
322
323#[cfg(test)]
324mod tests {
325 use super::*;
326 use crate::GraphemeClusterSegmenter;
327 use icu_provider::prelude::*;
328 use serde::Deserialize;
329
330 #[derive(PartialEq, Debug, Deserialize)]
334 struct TestCase {
335 unseg: String,
336 expected_bies: String,
337 true_bies: String,
338 }
339
340 #[derive(PartialEq, Debug, Deserialize)]
342 struct TestTextData {
343 testcases: Vec<TestCase>,
344 }
345
346 #[derive(Debug)]
347 struct TestText {
348 data: TestTextData,
349 }
350
351 #[test]
352 fn segment_file_by_lstm() {
353 let lstm: DataResponse<SegmenterLstmAutoV1> = crate::provider::Baked
354 .load(DataRequest {
355 id: DataIdentifierBorrowed::for_marker_attributes(
356 DataMarkerAttributes::from_str_or_panic(
357 "Thai_codepoints_exclusive_model4_heavy",
358 ),
359 ),
360 ..Default::default()
361 })
362 .unwrap();
363 let lstm = LstmSegmenter::new(lstm.payload.get(), GraphemeClusterSegmenter::new());
364
365 let test_text_data = serde_json::from_str(if lstm.grapheme.is_some() {
367 include_str!("../../../tests/testdata/test_text_graphclust.json")
368 } else {
369 include_str!("../../../tests/testdata/test_text_codepoints.json")
370 })
371 .expect("JSON syntax error");
372 let test_text = TestText {
373 data: test_text_data,
374 };
375
376 for test_case in &test_text.data.testcases {
378 let lstm_output = lstm
379 .segment_str(&test_case.unseg)
380 .bies
381 .map(|is_e| if is_e { 'e' } else { '?' })
382 .collect::<String>();
383 println!("Test case : {}", test_case.unseg);
384 println!("Expected bies : {}", test_case.expected_bies);
385 println!("Estimated bies : {lstm_output}");
386 println!("True bies : {}", test_case.true_bies);
387 println!("****************************************************");
388 assert_eq!(
389 test_case.expected_bies.replace(['b', 'i', 's'], "?"),
390 lstm_output
391 );
392 }
393 }
394}