@@ -120,12 +120,32 @@ impl VocabPrefixAutomaton {
120120
121121#[ cfg( feature = "pyo3" ) ]
122122pub mod pyo3 {
123- use pyo3:: { Bound , Python , pymethods, types:: PyBytes } ;
123+ use pyo3:: { Bound , FromPyObject , IntoPyObject , Python , pymethods, types:: PyBytes } ;
124124
125125 use crate :: { SortedTokenId , SortedTokenRange , TokenId } ;
126126
127127 use super :: VocabPrefixAutomaton ;
128128
129+ #[ derive( Debug , FromPyObject , IntoPyObject ) ]
130+ enum TokenIdSeq {
131+ TokenId ( TokenId ) ,
132+ Seq ( Vec < TokenId > ) ,
133+ }
134+
135+ impl TokenIdSeq {
136+ fn map < F : FnMut ( TokenId ) -> TokenId > ( self , mut f : F ) -> Self {
137+ match self {
138+ Self :: TokenId ( id) => Self :: TokenId ( f ( id) ) ,
139+ Self :: Seq ( mut items) => {
140+ items. iter_mut ( ) . for_each ( |item| {
141+ * item = f ( * item) ;
142+ } ) ;
143+ Self :: Seq ( items)
144+ }
145+ }
146+ }
147+ }
148+
129149 #[ pymethods]
130150 impl VocabPrefixAutomaton {
131151 #[ new]
@@ -151,10 +171,11 @@ pub mod pyo3 {
151171 #[ pyo3( name = "parse_bytes" ) ]
152172 fn parse_bytes_py (
153173 & self ,
174+ py : Python < ' _ > ,
154175 bytes : & [ u8 ] ,
155176 start_from : usize ,
156177 ) -> Vec < ( usize , SortedTokenRange ) > {
157- self . parse_bytes ( bytes, start_from)
178+ py . allow_threads ( || self . parse_bytes ( bytes, start_from) )
158179 }
159180
160181 #[ pyo3( name = "parse_tokens" ) ]
@@ -163,10 +184,32 @@ pub mod pyo3 {
163184 py : Python < ' py > ,
164185 tokens : Vec < usize > ,
165186 ) -> Vec < ( Bound < ' py , PyBytes > , SortedTokenRange ) > {
166- self . parse_rev_token_id_seq ( tokens. into_iter ( ) . rev ( ) )
167- . into_iter ( )
187+ let res = py . allow_threads ( || self . parse_rev_token_id_seq ( tokens. into_iter ( ) . rev ( ) ) ) ;
188+ res . into_iter ( )
168189 . map ( |( b, c) | ( PyBytes :: new ( py, & b) , c) )
169190 . collect ( )
170191 }
192+
193+ #[ pyo3( name = "parse_tokens_str_suffix" ) ]
194+ fn parse_tokens_str_suffix_py (
195+ & self ,
196+ py : Python < ' _ > ,
197+ tokens : Vec < usize > ,
198+ ) -> Vec < ( String , SortedTokenRange ) > {
199+ py. allow_threads ( || {
200+ self . parse_rev_token_id_seq ( tokens. into_iter ( ) . rev ( ) )
201+ . into_iter ( )
202+ . filter_map ( |( b, c) | String :: from_utf8 ( b. into ( ) ) . ok ( ) . map ( |s| ( s, c) ) )
203+ . collect ( )
204+ } )
205+ }
206+
207+ fn get_original_token_ids ( & self , py : Python < ' _ > , seq : TokenIdSeq ) -> TokenIdSeq {
208+ py. allow_threads ( || seq. map ( |id| self . order . get ( id as usize ) . copied ( ) . unwrap_or ( id) ) )
209+ }
210+
211+ fn get_sorted_token_ids ( & self , py : Python < ' _ > , seq : TokenIdSeq ) -> TokenIdSeq {
212+ py. allow_threads ( || seq. map ( |id| self . rank . get ( id as usize ) . copied ( ) . unwrap_or ( id) ) )
213+ }
171214 }
172215}
0 commit comments