Skip to main content

icu_segmenter/complex/lstm/
mod.rs

1// This file is part of ICU4X. For terms of use, please see the file
2// called LICENSE at the top level of the ICU4X source tree
3// (online at: https://github.com/unicode-org/icu4x/blob/main/LICENSE ).
4
5use 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
15// A word break iterator using LSTM model. Input string have to be same language.
16
17pub(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    /// Returns `Err` if grapheme data is required but not present
72    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)] // shape (2, 4, hunits)
79        let timew_fw = time_w.submatrix(0).unwrap();
80        #[expect(clippy::unwrap_used)] // shape (2, 4, hunits)
81        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    /// Create an LSTM based break iterator for an `str` (a UTF-8 string).
99    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    /// Create an LSTM based break iterator for a UTF-16 string.
140    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    // input_seq is a sequence of id numbers that represents grapheme clusters or code points in the input line. These ids are used later
204    // in the embedding layer of the model.
205    fn new(segmenter: &'l LstmSegmenter<'data>, input_seq: Vec<u16>) -> Self {
206        let hunits = segmenter.fw_u.dim().1;
207
208        // Backward LSTM
209        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(), /* shape (dict.len() + 1, hunit), g_id is at most dict.len() */
218                h_bw.submatrix_mut(i).unwrap(), // shape (input_seq.len(), hunits)
219                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(), // shape (dict.len() + 1, hunit), g_id is at most dict.len()
254            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)] // shape (input_seq.len(), hunits)
262        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)] // both shape (4)
271        curr_est.add(self.segmenter.time_b).unwrap();
272        // For correct BIES weight calculation we'd now have to apply softmax, however
273        // we're only doing a naive argmax, so a monotonic function doesn't make a difference.
274
275        Some(weights[2] > weights[0] && weights[2] > weights[1] && weights[2] > weights[3])
276    }
277}
278
279/// `compute_hc1` implemens the evaluation of one LSTM layer.
280fn 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)] // first dimension is 4
304    s_t.submatrix_mut::<1>(0).unwrap().sigmoid_transform();
305    #[expect(clippy::unwrap_used)] // first dimension is 4
306    s_t.submatrix_mut::<1>(1).unwrap().sigmoid_transform();
307    #[expect(clippy::unwrap_used)] // first dimension is 4
308    s_t.submatrix_mut::<1>(2).unwrap().tanh_transform();
309    #[expect(clippy::unwrap_used)] // first dimension is 4
310    s_t.submatrix_mut::<1>(3).unwrap().sigmoid_transform();
311
312    #[expect(clippy::unwrap_used)] // first dimension is 4
313    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)] // first dimension is 4
320    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    /// `TestCase` is a struct used to store a single test case.
331    /// Each test case has two attributes: `unseg` which denotes the unsegmented line, and `true_bies` which indicates the Bies
332    /// sequence representing the true segmentation.
333    #[derive(PartialEq, Debug, Deserialize)]
334    struct TestCase {
335        unseg: String,
336        expected_bies: String,
337        true_bies: String,
338    }
339
340    /// `TestTextData` is a struct to store a vector of `TestCase` that represents a test text.
341    #[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        // Importing the test data
366        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        // Testing
377        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}